fix(middleware/cors): 完善CORS中间件的日志和校验逻辑
1. 修复了变量缩进的格式问题 2. 添加初始化时的白名单来源打印 3. 优化origin校验逻辑,兼容大小写的协议前缀 4. 新增各阶段的CORS请求日志 5. 修正日志中的客户端IP字段名 6. 修复文件末尾缺失换行符的问题
This commit is contained in:
+15
-5
@@ -1,4 +1,4 @@
|
|||||||
package middleware
|
package middleware
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"log"
|
"log"
|
||||||
@@ -8,8 +8,8 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
allowedOrigins []string
|
allowedOrigins []string
|
||||||
allowAllOrigins bool
|
allowAllOrigins bool
|
||||||
corsMaxAgeHeader = "86400"
|
corsMaxAgeHeader = "86400"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -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))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user