56 lines
1.5 KiB
Go
56 lines
1.5 KiB
Go
package middleware
|
||
|
||
import (
|
||
"net/http"
|
||
"strconv"
|
||
"strings"
|
||
|
||
"fashionapi/internal/config"
|
||
|
||
"github.com/gin-gonic/gin"
|
||
)
|
||
|
||
// CORS 跨域中间件,规则来自 yml 配置。
|
||
//
|
||
// allow_origins 为 ["*"] 时直接放行全部来源(与原项目行为一致);
|
||
// 配置为具体域名列表时,仅当请求 Origin 命中白名单才回写该 Origin,
|
||
// 这样生产环境可以收紧到指定域名而无需改代码。
|
||
func CORS(cfg config.CORSConfig) gin.HandlerFunc {
|
||
allowAll := len(cfg.AllowOrigins) == 0
|
||
allowed := make(map[string]struct{}, len(cfg.AllowOrigins))
|
||
for _, o := range cfg.AllowOrigins {
|
||
if o == "*" {
|
||
allowAll = true
|
||
}
|
||
allowed[o] = struct{}{}
|
||
}
|
||
|
||
methods := strings.Join(cfg.AllowMethods, ", ")
|
||
headers := strings.Join(cfg.AllowHeaders, ", ")
|
||
maxAge := strconv.Itoa(cfg.MaxAge)
|
||
|
||
return func(c *gin.Context) {
|
||
origin := c.GetHeader("Origin")
|
||
switch {
|
||
case allowAll:
|
||
c.Header("Access-Control-Allow-Origin", "*")
|
||
case origin != "":
|
||
if _, ok := allowed[origin]; ok {
|
||
c.Header("Access-Control-Allow-Origin", origin)
|
||
// 指定来源时必须声明 Vary,避免 CDN / 代理把响应错误地跨来源复用
|
||
c.Header("Vary", "Origin")
|
||
}
|
||
}
|
||
c.Header("Access-Control-Allow-Methods", methods)
|
||
c.Header("Access-Control-Allow-Headers", headers)
|
||
c.Header("Access-Control-Max-Age", maxAge)
|
||
|
||
// 预检请求直接结束,不进入业务逻辑
|
||
if c.Request.Method == http.MethodOptions {
|
||
c.AbortWithStatus(http.StatusNoContent)
|
||
return
|
||
}
|
||
c.Next()
|
||
}
|
||
}
|