package router import ( "net/http" "net/http/httptest" "testing" "github.com/gin-gonic/gin" ) // corsEngine 起一个只挂 cors() 的引擎(env 由各用例 t.Setenv 预置)。 func corsEngine() *gin.Engine { gin.SetMode(gin.TestMode) r := gin.New() r.Use(cors()) r.GET("/ping", func(c *gin.Context) { c.String(http.StatusOK, "ok") }) return r } func doCORS(t *testing.T, r *gin.Engine, method, origin string) *httptest.ResponseRecorder { t.Helper() req := httptest.NewRequest(method, "/ping", nil) if origin != "" { req.Header.Set("Origin", origin) } w := httptest.NewRecorder() r.ServeHTTP(w, req) return w } func TestCORS_MultiOrigin(t *testing.T) { t.Setenv("CORS_ALLOW_ORIGIN", "http://a.example,http://b.example") t.Setenv("APP_ENV", "") t.Setenv("GIN_MODE", "") r := corsEngine() // 命中白名单 → 回显请求方 + Vary: Origin(多 origin 时 ACAO 只能发单值)。 w := doCORS(t, r, http.MethodGet, "http://b.example") if got := w.Header().Get("Access-Control-Allow-Origin"); got != "http://b.example" { t.Errorf("命中白名单应回显请求 Origin,得到 %q", got) } if w.Header().Get("Vary") != "Origin" { t.Errorf("多 origin 回显必须带 Vary: Origin,否则缓存串源") } // 未命中 → 不发 ACAO(浏览器按同源策略拦截)。 w = doCORS(t, r, http.MethodGet, "http://evil.example") if got := w.Header().Get("Access-Control-Allow-Origin"); got != "" { t.Errorf("白名单外的 Origin 不应发 ACAO,得到 %q", got) } } func TestCORS_SingleOriginLegacy(t *testing.T) { t.Setenv("CORS_ALLOW_ORIGIN", "http://only.example") t.Setenv("APP_ENV", "") t.Setenv("GIN_MODE", "") r := corsEngine() // 单值配置保持旧行为:不带 Origin 的请求(探活/curl)也能看到头。 w := doCORS(t, r, http.MethodGet, "") if got := w.Header().Get("Access-Control-Allow-Origin"); got != "http://only.example" { t.Errorf("单值配置应无条件直写,得到 %q", got) } } func TestCORS_DevWildcardAndPreflight(t *testing.T) { t.Setenv("CORS_ALLOW_ORIGIN", "") t.Setenv("APP_ENV", "") t.Setenv("GIN_MODE", "") r := corsEngine() w := doCORS(t, r, http.MethodGet, "http://whatever.example") if got := w.Header().Get("Access-Control-Allow-Origin"); got != "*" { t.Errorf("开发缺省应放行 *,得到 %q", got) } // 预检直接 204 短路。 w = doCORS(t, r, http.MethodOptions, "http://whatever.example") if w.Code != http.StatusNoContent { t.Errorf("OPTIONS 预检应 204,得到 %d", w.Code) } } func TestCORS_ProdUnsetDeniesAll(t *testing.T) { t.Setenv("CORS_ALLOW_ORIGIN", "") t.Setenv("APP_ENV", "prod") t.Setenv("GIN_MODE", "") r := corsEngine() w := doCORS(t, r, http.MethodGet, "http://a.example") if got := w.Header().Get("Access-Control-Allow-Origin"); got != "" { t.Errorf("生产未配置不应发 ACAO,得到 %q", got) } }