diff --git a/internal/middleware/cors.go b/internal/middleware/cors.go new file mode 100644 index 0000000..109a431 --- /dev/null +++ b/internal/middleware/cors.go @@ -0,0 +1,49 @@ +package middleware + +import ( + "net/http" + + "github.com/gin-gonic/gin" + + "gotest/internal/config" +) + +// CORS 跨域中间件,根据配置放行指定来源 +func CORS(cfg *config.CORSConfig) gin.HandlerFunc { + allowed := make(map[string]bool) + wildcard := false + for _, origin := range cfg.AllowedOrigins { + if origin == "*" { + // 配置中包含 * 表示通配,允许任意来源 + wildcard = true + continue + } + allowed[origin] = true + } + + return func(c *gin.Context) { + origin := c.GetHeader("Origin") + switch { + case wildcard: + // 通配模式直接放行所有来源 + c.Header("Access-Control-Allow-Origin", "*") + case origin != "" && allowed[origin]: + // 命中白名单时回写具体来源,并标注 Vary 以便缓存正确 + c.Header("Access-Control-Allow-Origin", origin) + c.Header("Vary", "Origin") + } + + c.Header("Access-Control-Allow-Methods", "GET, POST, PUT, DELETE, OPTIONS, PATCH") + c.Header("Access-Control-Allow-Headers", "Origin, Content-Type, Authorization, X-Requested-With") + c.Header("Access-Control-Allow-Credentials", "true") + c.Header("Access-Control-Max-Age", "86400") + + // 处理 OPTIONS 预检请求,直接返回 204 + if c.Request.Method == http.MethodOptions { + c.AbortWithStatus(http.StatusNoContent) + return + } + + c.Next() + } +}