package middleware //nolint:testpackage // tests internal CORS behavior import ( "log/slog" "net/http" "net/http/httptest" "testing" "github.com/stretchr/testify/assert" "sneak.berlin/go/upaas/internal/config" ) //nolint:gosec // test credentials func newCORSTestMiddleware(corsOrigins string) *Middleware { return &Middleware{ log: slog.Default(), params: &Params{ Config: &config.Config{ CORSOrigins: corsOrigins, SessionSecret: "test-secret-32-bytes-long-enough", }, }, } } // assertNoCORSHeaders runs a request with the given Origin header through // CORS middleware configured with corsOrigins and asserts that no // Access-Control-Allow-Origin header is set. func assertNoCORSHeaders(t *testing.T, corsOrigins, origin, msg string) { t.Helper() m := newCORSTestMiddleware(corsOrigins) handler := m.CORS()(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) })) req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", nil) req.Header.Set("Origin", origin) rec := httptest.NewRecorder() handler.ServeHTTP(rec, req) assert.Empty(t, rec.Header().Get("Access-Control-Allow-Origin"), msg) } func TestCORS_NoOriginsConfigured_NoCORSHeaders(t *testing.T) { t.Parallel() assertNoCORSHeaders(t, "", "https://evil.com", "expected no CORS headers when no origins configured") } func TestCORS_OriginsConfigured_AllowsMatchingOrigin(t *testing.T) { t.Parallel() m := newCORSTestMiddleware("https://app.example.com,https://other.example.com") handler := m.CORS()(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) })) req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", nil) req.Header.Set("Origin", "https://app.example.com") rec := httptest.NewRecorder() handler.ServeHTTP(rec, req) assert.Equal(t, "https://app.example.com", rec.Header().Get("Access-Control-Allow-Origin")) assert.Equal(t, "true", rec.Header().Get("Access-Control-Allow-Credentials")) } func TestCORS_OriginsConfigured_RejectsNonMatchingOrigin(t *testing.T) { t.Parallel() assertNoCORSHeaders(t, "https://app.example.com", "https://evil.com", "expected no CORS headers for non-matching origin") }