From b92d3dbc00e9ae02fe33c40a272bfc6db5b38436 Mon Sep 17 00:00:00 2001 From: sun <3371392206@qq.com> Date: Mon, 22 Jun 2026 13:47:17 +0800 Subject: [PATCH] =?UTF-8?q?fix(middleware/cors):=20=E5=AE=8C=E5=96=84CORS?= =?UTF-8?q?=E4=B8=AD=E9=97=B4=E4=BB=B6=E7=9A=84=E6=97=A5=E5=BF=97=E5=92=8C?= =?UTF-8?q?=E6=A0=A1=E9=AA=8C=E9=80=BB=E8=BE=91?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 1. 修复了变量缩进的格式问题 2. 添加初始化时的白名单来源打印 3. 优化origin校验逻辑,兼容大小写的协议前缀 4. 新增各阶段的CORS请求日志 5. 修正日志中的客户端IP字段名 6. 修复文件末尾缺失换行符的问题 --- middleware/cors.go | 22 ++++++++++++++++------ 1 file changed, 16 insertions(+), 6 deletions(-) diff --git a/middleware/cors.go b/middleware/cors.go index 8b04ff0..c391d13 100644 --- a/middleware/cors.go +++ b/middleware/cors.go @@ -1,4 +1,4 @@ -package middleware +package middleware import ( "log" @@ -8,8 +8,8 @@ import ( ) var ( - allowedOrigins []string - allowAllOrigins bool + allowedOrigins []string + allowAllOrigins bool corsMaxAgeHeader = "86400" ) @@ -46,6 +46,9 @@ func InitCORSConfig() { } if len(allowedOrigins) > 0 { log.Printf("已配置 %d 个允许的跨域来源白名单", len(allowedOrigins)) + for _, o := range allowedOrigins { + log.Printf(" - %s", o) + } } } @@ -53,7 +56,9 @@ func isValidOrigin(origin string) bool { if origin == "" || origin == "null" || origin == "nil" { return false } - if !strings.HasPrefix(origin, "http://") && !strings.HasPrefix(origin, "https://") { + // 使用小写比较,避免大小写问题 + lowerOrigin := strings.ToLower(origin) + if !strings.HasPrefix(lowerOrigin, "http://") && !strings.HasPrefix(lowerOrigin, "https://") { return false } return true @@ -61,17 +66,22 @@ func isValidOrigin(origin string) bool { func matchOrigin(origin string) (string, bool) { if !isValidOrigin(origin) { + log.Printf("[CORS] Origin %q 验证失败", origin) return "", false } if allowAllOrigins { + log.Printf("[CORS] Origin %q 匹配 allowAllOrigins", origin) return "*", true } normalized := normalizeOrigin(origin) + log.Printf("[CORS] 检查 origin %q (normalized: %q) 对比白名单: %v", origin, normalized, allowedOrigins) for _, allowed := range allowedOrigins { if allowed == normalized { + log.Printf("[CORS] Origin %q 匹配白名单 %q", origin, allowed) return origin, true } } + log.Printf("[CORS] Origin %q 未匹配任何白名单", origin) return "", false } @@ -97,7 +107,7 @@ func CORS(next http.Handler) http.Handler { w.Header().Set("Vary", vary+", Origin") } } else { - log.Printf("CORS拦截: 来源=%q 路径=%s 方法=%s 客户端=%s", + log.Printf("CORS拦截: 来源=%q 路径=%s 方法=%s 客户端IP=%s", origin, r.URL.Path, r.Method, GetClientIP(r)) } } @@ -109,4 +119,4 @@ func CORS(next http.Handler) http.Handler { next.ServeHTTP(w, r) }) -} +} \ No newline at end of file