From fb341478a77e0ba22eec65d2d0d7c1a9d5288878 Mon Sep 17 00:00:00 2001 From: Aaro Koinsaari <89689072+koinsaari@users.noreply.github.com> Date: Sat, 1 Aug 2026 17:43:00 +0300 Subject: [PATCH] feat: add CORS middleware for browser-based clients [skip ci] Restricts cross-origin access to a single configured origin via CORS_ALLOWED_ORIGIN. No-op when unset. --- .env.example | 1 + cmd/api/main.go | 1 + internal/middleware/cors.go | 36 ++++++++ internal/middleware/cors_test.go | 137 +++++++++++++++++++++++++++++++ 4 files changed, 175 insertions(+) create mode 100644 internal/middleware/cors.go create mode 100644 internal/middleware/cors_test.go diff --git a/.env.example b/.env.example index 549b84f..07988ac 100644 --- a/.env.example +++ b/.env.example @@ -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 diff --git a/cmd/api/main.go b/cmd/api/main.go index 2776d5a..cbe2a98 100644 --- a/cmd/api/main.go +++ b/cmd/api/main.go @@ -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) diff --git a/internal/middleware/cors.go b/internal/middleware/cors.go new file mode 100644 index 0000000..5c7802a --- /dev/null +++ b/internal/middleware/cors.go @@ -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) + }) + } +} diff --git a/internal/middleware/cors_test.go b/internal/middleware/cors_test.go new file mode 100644 index 0000000..15b5eb5 --- /dev/null +++ b/internal/middleware/cors_test.go @@ -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") + } +}