Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .env.example
Original file line number Diff line number Diff line change
Expand Up @@ -4,3 +4,4 @@ DB_NAME=inwheel
DB_SSLMODE=disable
DB_MAX_OPEN_CONNS=5
DB_MAX_IDLE_CONNS=2
# CORS_ALLOWED_ORIGIN=https://example.com
1 change: 1 addition & 0 deletions cmd/api/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -138,6 +138,7 @@ func main() {
srv.validationErrorHandler(w, r, err)
},
})(v1Mux)
v1Handler = middleware.CORS(getEnv("CORS_ALLOWED_ORIGIN", ""))(v1Handler)

mux := http.NewServeMux()
mux.HandleFunc("GET /healthz", srv.handleHealthz)
Expand Down
36 changes: 36 additions & 0 deletions internal/middleware/cors.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,36 @@
/*
* Copyright (C) 2026 InWheel Contributors
* SPDX-License-Identifier: AGPL-3.0-only
*/

package middleware

import "net/http"

func CORS(allowedOrigin string) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
if allowedOrigin == "" {
return next
}
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Vary", "Origin")

origin := r.Header.Get("Origin")
if origin != allowedOrigin {
next.ServeHTTP(w, r)
return
}

w.Header().Set("Access-Control-Allow-Origin", origin)

if r.Method == http.MethodOptions {
w.Header().Set("Access-Control-Allow-Methods", "GET, POST, PATCH, DELETE")
w.Header().Set("Access-Control-Allow-Headers", "X-API-Key, Content-Type")
w.WriteHeader(http.StatusNoContent)
return
}

next.ServeHTTP(w, r)
})
}
}
137 changes: 137 additions & 0 deletions internal/middleware/cors_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,137 @@
/*
* Copyright (C) 2026 InWheel Contributors
* SPDX-License-Identifier: AGPL-3.0-only
*/

package middleware

import (
"net/http"
"net/http/httptest"
"testing"
)

func TestCORS_NoAllowedOrigin_ReturnsNextUnchanged(t *testing.T) {
t.Parallel()
called := false
next := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { called = true })

handler := CORS("")(next)
r := httptest.NewRequest(http.MethodGet, "/", nil)
r.Header.Set("Origin", "https://example.com")
w := httptest.NewRecorder()
handler.ServeHTTP(w, r)

if !called {
t.Fatal("next was not called")
}
if got := w.Header().Get("Access-Control-Allow-Origin"); got != "" {
t.Errorf("Access-Control-Allow-Origin = %q, want empty", got)
}
}

func TestCORS_NoOriginHeader_PassesThrough(t *testing.T) {
t.Parallel()
called := false
next := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { called = true })

handler := CORS("https://example.com")(next)
r := httptest.NewRequest(http.MethodGet, "/", nil)
w := httptest.NewRecorder()
handler.ServeHTTP(w, r)

if !called {
t.Fatal("next was not called")
}
if got := w.Header().Get("Access-Control-Allow-Origin"); got != "" {
t.Errorf("Access-Control-Allow-Origin = %q, want empty", got)
}
}

func TestCORS_MatchingOrigin_SetsHeadersAndCallsNext(t *testing.T) {
t.Parallel()
called := false
next := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { called = true })

handler := CORS("https://example.com")(next)
r := httptest.NewRequest(http.MethodPatch, "/v1/places/1/accessibility", nil)
r.Header.Set("Origin", "https://example.com")
w := httptest.NewRecorder()
handler.ServeHTTP(w, r)

if !called {
t.Fatal("next was not called")
}
if got := w.Header().Get("Access-Control-Allow-Origin"); got != "https://example.com" {
t.Errorf("Access-Control-Allow-Origin = %q, want %q", got, "https://example.com")
}
if got := w.Header().Get("Vary"); got != "Origin" {
t.Errorf("Vary = %q, want %q", got, "Origin")
}
}

func TestCORS_NonMatchingOrigin_NoHeadersButPassesThrough(t *testing.T) {
t.Parallel()
called := false
next := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { called = true })

handler := CORS("https://example.com")(next)
r := httptest.NewRequest(http.MethodGet, "/", nil)
r.Header.Set("Origin", "https://evil.example.com")
w := httptest.NewRecorder()
handler.ServeHTTP(w, r)

if !called {
t.Fatal("next was not called")
}
if got := w.Header().Get("Access-Control-Allow-Origin"); got != "" {
t.Errorf("Access-Control-Allow-Origin = %q, want empty", got)
}
if got := w.Header().Get("Vary"); got != "Origin" {
t.Errorf("Vary = %q, want %q", got, "Origin")
}
}

func TestCORS_PreflightNonMatchingOrigin_PassesThrough(t *testing.T) {
t.Parallel()
called := false
next := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { called = true })

handler := CORS("https://example.com")(next)
r := httptest.NewRequest(http.MethodOptions, "/v1/places/1/accessibility", nil)
r.Header.Set("Origin", "https://evil.example.com")
w := httptest.NewRecorder()
handler.ServeHTTP(w, r)

if !called {
t.Fatal("next was not called for a preflight request from a non-matching origin")
}
if got := w.Header().Get("Access-Control-Allow-Origin"); got != "" {
t.Errorf("Access-Control-Allow-Origin = %q, want empty", got)
}
}

func TestCORS_PreflightMatchingOrigin_ShortCircuits(t *testing.T) {
t.Parallel()
called := false
next := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { called = true })

handler := CORS("https://example.com")(next)
r := httptest.NewRequest(http.MethodOptions, "/v1/places/1/accessibility", nil)
r.Header.Set("Origin", "https://example.com")
w := httptest.NewRecorder()
handler.ServeHTTP(w, r)

if called {
t.Fatal("next was called for a preflight request")
}
if w.Code != http.StatusNoContent {
t.Errorf("status = %d, want %d", w.Code, http.StatusNoContent)
}
if got := w.Header().Get("Access-Control-Allow-Methods"); got != "GET, POST, PATCH, DELETE" {
t.Errorf("Access-Control-Allow-Methods = %q, want %q", got, "GET, POST, PATCH, DELETE")
}
if got := w.Header().Get("Access-Control-Allow-Headers"); got != "X-API-Key, Content-Type" {
t.Errorf("Access-Control-Allow-Headers = %q, want %q", got, "X-API-Key, Content-Type")
}
}