package middleware import ( "net/http" "net/http/httptest" "testing" "github.com/gin-gonic/gin" "github.com/stretchr/testify/assert" "github.com/gochat/gochat/internal/config" ) func init() { gin.SetMode(gin.TestMode) } // --- CORS Middleware Tests --- func TestCORS_DevMode_AllOrigins(t *testing.T) { cfg := CORSConfig{DevMode: true} router := gin.New() router.Use(CORS(cfg)) router.GET("/test", func(c *gin.Context) { c.Status(200) }) w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) req.Header.Set("Origin", "https://evil.example.com") router.ServeHTTP(w, req) assert.Equal(t, "*", w.Header().Get("Access-Control-Allow-Origin")) assert.Contains(t, w.Header().Get("Access-Control-Allow-Methods"), "GET") assert.Contains(t, w.Header().Get("Access-Control-Allow-Methods"), "POST") assert.Contains(t, w.Header().Get("Access-Control-Allow-Headers"), "Authorization") assert.Equal(t, "Content-Length", w.Header().Get("Access-Control-Expose-Headers")) assert.Equal(t, "86400", w.Header().Get("Access-Control-Max-Age")) } func TestCORS_DevMode_WithWhitelist(t *testing.T) { cfg := CORSConfig{ AllowedOrigins: []string{"https://app.example.com"}, DevMode: true, } router := gin.New() router.Use(CORS(cfg)) router.GET("/test", func(c *gin.Context) { c.Status(200) }) // Whitelist takes precedence even in dev mode w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) req.Header.Set("Origin", "https://app.example.com") router.ServeHTTP(w, req) assert.Equal(t, "https://app.example.com", w.Header().Get("Access-Control-Allow-Origin")) assert.Equal(t, "Origin", w.Header().Get("Vary")) } func TestCORS_Production_WhitelistMatch(t *testing.T) { cfg := CORSConfig{ AllowedOrigins: []string{"https://app.example.com", "https://admin.example.com"}, AllowCredentials: true, } router := gin.New() router.Use(CORS(cfg)) router.GET("/test", func(c *gin.Context) { c.Status(200) }) w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) req.Header.Set("Origin", "https://app.example.com") router.ServeHTTP(w, req) assert.Equal(t, "https://app.example.com", w.Header().Get("Access-Control-Allow-Origin")) assert.Equal(t, "Origin", w.Header().Get("Vary")) assert.Equal(t, "true", w.Header().Get("Access-Control-Allow-Credentials")) } func TestCORS_Production_WhitelistNoMatch(t *testing.T) { cfg := CORSConfig{ AllowedOrigins: []string{"https://app.example.com"}, } router := gin.New() router.Use(CORS(cfg)) router.GET("/test", func(c *gin.Context) { c.Status(200) }) w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) req.Header.Set("Origin", "https://evil.example.com") router.ServeHTTP(w, req) // No Allow-Origin header for non-whitelisted origin assert.Empty(t, w.Header().Get("Access-Control-Allow-Origin")) // But other CORS headers are still set (methods, headers, etc.) assert.NotEmpty(t, w.Header().Get("Access-Control-Allow-Methods")) } func TestCORS_Production_NoCredentials(t *testing.T) { cfg := CORSConfig{ AllowedOrigins: []string{"https://app.example.com"}, AllowCredentials: false, } router := gin.New() router.Use(CORS(cfg)) router.GET("/test", func(c *gin.Context) { c.Status(200) }) w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) req.Header.Set("Origin", "https://app.example.com") router.ServeHTTP(w, req) assert.Empty(t, w.Header().Get("Access-Control-Allow-Credentials")) } func TestCORS_WildcardSubdomain(t *testing.T) { cfg := CORSConfig{ AllowedOrigins: []string{"*.example.com"}, } router := gin.New() router.Use(CORS(cfg)) router.GET("/test", func(c *gin.Context) { c.Status(200) }) // Subdomain match w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) req.Header.Set("Origin", "https://foo.example.com") router.ServeHTTP(w, req) assert.Equal(t, "https://foo.example.com", w.Header().Get("Access-Control-Allow-Origin")) // Another subdomain w2 := httptest.NewRecorder() req2, _ := http.NewRequest("GET", "/test", nil) req2.Header.Set("Origin", "https://bar.baz.example.com") router.ServeHTTP(w2, req2) assert.Equal(t, "https://bar.baz.example.com", w2.Header().Get("Access-Control-Allow-Origin")) // Base domain itself should NOT match *.example.com w3 := httptest.NewRecorder() req3, _ := http.NewRequest("GET", "/test", nil) req3.Header.Set("Origin", "https://example.com") router.ServeHTTP(w3, req3) assert.Empty(t, w3.Header().Get("Access-Control-Allow-Origin")) } func TestCORS_OptionsPreflight(t *testing.T) { cfg := CORSConfig{ AllowedOrigins: []string{"https://app.example.com"}, AllowCredentials: true, } router := gin.New() router.Use(CORS(cfg)) router.POST("/api/data", func(c *gin.Context) { c.Status(200) }) w := httptest.NewRecorder() req, _ := http.NewRequest("OPTIONS", "/api/data", nil) req.Header.Set("Origin", "https://app.example.com") router.ServeHTTP(w, req) assert.Equal(t, http.StatusNoContent, w.Code) assert.Equal(t, "https://app.example.com", w.Header().Get("Access-Control-Allow-Origin")) assert.Equal(t, "true", w.Header().Get("Access-Control-Allow-Credentials")) assert.NotEmpty(t, w.Header().Get("Access-Control-Allow-Methods")) } func TestCORS_CustomMethodsHeaders(t *testing.T) { cfg := CORSConfig{ AllowedOrigins: []string{"https://app.example.com"}, AllowedMethods: []string{"GET", "POST"}, AllowedHeaders: []string{"X-Custom-Header", "Authorization"}, ExposeHeaders: []string{"X-Total-Count", "Content-Length"}, MaxAge: 3600, } router := gin.New() router.Use(CORS(cfg)) router.GET("/test", func(c *gin.Context) { c.Status(200) }) w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) req.Header.Set("Origin", "https://app.example.com") router.ServeHTTP(w, req) assert.Equal(t, "GET, POST", w.Header().Get("Access-Control-Allow-Methods")) assert.Equal(t, "X-Custom-Header, Authorization", w.Header().Get("Access-Control-Allow-Headers")) assert.Equal(t, "X-Total-Count, Content-Length", w.Header().Get("Access-Control-Expose-Headers")) assert.Equal(t, "3600", w.Header().Get("Access-Control-Max-Age")) } func TestCORS_NoOriginHeader(t *testing.T) { cfg := CORSConfig{ AllowedOrigins: []string{"https://app.example.com"}, } router := gin.New() router.Use(CORS(cfg)) router.GET("/test", func(c *gin.Context) { c.Status(200) }) w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) // No Origin header router.ServeHTTP(w, req) assert.Empty(t, w.Header().Get("Access-Control-Allow-Origin")) assert.NotEmpty(t, w.Header().Get("Access-Control-Allow-Methods")) } func TestCORS_ProductionNoWhitelist_RejectsAll(t *testing.T) { cfg := CORSConfig{DevMode: false} // no whitelist, not dev mode router := gin.New() router.Use(CORS(cfg)) router.GET("/test", func(c *gin.Context) { c.Status(200) }) w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) req.Header.Set("Origin", "https://random.example.com") router.ServeHTTP(w, req) assert.Empty(t, w.Header().Get("Access-Control-Allow-Origin")) } // --- Origin Matching Tests --- func TestMatchOrigin_ExactMatch(t *testing.T) { assert.True(t, matchOrigin("https://app.example.com", "https://app.example.com")) assert.False(t, matchOrigin("https://app.example.com", "https://other.example.com")) } func TestMatchOrigin_WildcardSubdomain(t *testing.T) { assert.True(t, matchOrigin("https://foo.example.com", "*.example.com")) assert.True(t, matchOrigin("http://bar.example.com", "*.example.com")) assert.True(t, matchOrigin("https://deep.sub.example.com", "*.example.com")) assert.False(t, matchOrigin("https://example.com", "*.example.com")) // base domain assert.False(t, matchOrigin("https://foo.other.com", "*.example.com")) } func TestMatchOrigin_WildcardWithPort(t *testing.T) { assert.True(t, matchOrigin("https://foo.example.com:8080", "*.example.com")) assert.False(t, matchOrigin("https://example.com:8080", "*.example.com")) } func TestMatchOrigin_InvalidOrigin(t *testing.T) { assert.False(t, matchOrigin("", "*.example.com")) assert.False(t, matchOrigin("not-a-url", "*.example.com")) } func TestExtractHost(t *testing.T) { assert.Equal(t, "example.com", extractHost("https://example.com")) assert.Equal(t, "example.com", extractHost("https://example.com:8080")) assert.Equal(t, "foo.example.com", extractHost("https://foo.example.com")) assert.Equal(t, "foo.example.com", extractHost("http://foo.example.com:3000/path")) } func TestExtractHost_MalformedURL(t *testing.T) { // URL without scheme — url.Parse fails, fallback extracts host assert.Equal(t, "example.com", extractHost("http://example.com")) // "example.com" alone (no scheme) returns empty via url.Parse, // and the fallback strips after scheme:// which doesn't exist, // then strips path/port — result may be empty or partial // This is expected behavior: malformed origins without scheme // are not valid CORS origins and won't match } // --- CORSConfigFromAppConfig Tests --- func TestCORSConfigFromAppConfig_DebugMode(t *testing.T) { cfg := &config.Config{ Server: config.ServerConfig{ Mode: "debug", CORS: config.CORSConfig{ AllowedOrigins: []string{"https://app.example.com"}, AllowCredentials: true, MaxAge: 3600, }, }, } mwCfg := CORSConfigFromAppConfig(cfg) assert.True(t, mwCfg.DevMode) assert.Equal(t, []string{"https://app.example.com"}, mwCfg.AllowedOrigins) assert.True(t, mwCfg.AllowCredentials) assert.Equal(t, 3600, mwCfg.MaxAge) } func TestCORSConfigFromAppConfig_ProductionMode(t *testing.T) { cfg := &config.Config{ Server: config.ServerConfig{ Mode: "release", CORS: config.CORSConfig{ AllowedOrigins: []string{"https://app.example.com", "*.internal.com"}, AllowedMethods: []string{"GET", "POST", "PUT"}, AllowedHeaders: []string{"Authorization", "Content-Type"}, ExposeHeaders: []string{"X-Total-Count"}, AllowCredentials: true, MaxAge: 7200, }, }, } mwCfg := CORSConfigFromAppConfig(cfg) assert.False(t, mwCfg.DevMode) assert.Equal(t, []string{"https://app.example.com", "*.internal.com"}, mwCfg.AllowedOrigins) assert.Equal(t, []string{"GET", "POST", "PUT"}, mwCfg.AllowedMethods) assert.Equal(t, []string{"Authorization", "Content-Type"}, mwCfg.AllowedHeaders) assert.Equal(t, []string{"X-Total-Count"}, mwCfg.ExposeHeaders) assert.True(t, mwCfg.AllowCredentials) assert.Equal(t, 7200, mwCfg.MaxAge) } func TestCORSConfigFromAppConfig_EmptyCORS(t *testing.T) { cfg := &config.Config{ Server: config.ServerConfig{ Mode: "release", CORS: config.CORSConfig{}, // empty — all defaults }, } mwCfg := CORSConfigFromAppConfig(cfg) assert.False(t, mwCfg.DevMode) assert.Empty(t, mwCfg.AllowedOrigins) assert.Empty(t, mwCfg.AllowedMethods) // defaults applied in CORS() middleware, not here }