fix(middleware/cors): 完善CORS中间件的日志和校验逻辑

1.  修复了变量缩进的格式问题
2.  添加初始化时的白名单来源打印
3.  优化origin校验逻辑,兼容大小写的协议前缀
4.  新增各阶段的CORS请求日志
5.  修正日志中的客户端IP字段名
6.  修复文件末尾缺失换行符的问题
This commit is contained in:
sun
2026-06-22 13:47:17 +08:00
parent b92973cdc9
commit b92d3dbc00
+13 -3
View File
@@ -1,4 +1,4 @@
package middleware package middleware
import ( import (
"log" "log"
@@ -46,6 +46,9 @@ func InitCORSConfig() {
} }
if len(allowedOrigins) > 0 { if len(allowedOrigins) > 0 {
log.Printf("已配置 %d 个允许的跨域来源白名单", len(allowedOrigins)) 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" { if origin == "" || origin == "null" || origin == "nil" {
return false 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 false
} }
return true return true
@@ -61,17 +66,22 @@ func isValidOrigin(origin string) bool {
func matchOrigin(origin string) (string, bool) { func matchOrigin(origin string) (string, bool) {
if !isValidOrigin(origin) { if !isValidOrigin(origin) {
log.Printf("[CORS] Origin %q 验证失败", origin)
return "", false return "", false
} }
if allowAllOrigins { if allowAllOrigins {
log.Printf("[CORS] Origin %q 匹配 allowAllOrigins", origin)
return "*", true return "*", true
} }
normalized := normalizeOrigin(origin) normalized := normalizeOrigin(origin)
log.Printf("[CORS] 检查 origin %q (normalized: %q) 对比白名单: %v", origin, normalized, allowedOrigins)
for _, allowed := range allowedOrigins { for _, allowed := range allowedOrigins {
if allowed == normalized { if allowed == normalized {
log.Printf("[CORS] Origin %q 匹配白名单 %q", origin, allowed)
return origin, true return origin, true
} }
} }
log.Printf("[CORS] Origin %q 未匹配任何白名单", origin)
return "", false return "", false
} }
@@ -97,7 +107,7 @@ func CORS(next http.Handler) http.Handler {
w.Header().Set("Vary", vary+", Origin") w.Header().Set("Vary", vary+", Origin")
} }
} else { } else {
log.Printf("CORS拦截: 来源=%q 路径=%s 方法=%s 客户端=%s", log.Printf("CORS拦截: 来源=%q 路径=%s 方法=%s 客户端IP=%s",
origin, r.URL.Path, r.Method, GetClientIP(r)) origin, r.URL.Path, r.Method, GetClientIP(r))
} }
} }