package middleware import ( "crypto/rand" "encoding/hex" "net/http" "strconv" "github.com/gin-contrib/sessions" "github.com/gin-gonic/gin" ) // CSRF 防护在现有会话存储之上采用同步令牌(synchronizer-token)模式: // - 安全方法(GET/HEAD/OPTIONS):首次使用时按会话生成令牌, // 并通过模板 / JS 暴露,以便嵌入表单。 // - 不安全方法(POST/PUT/PATCH/DELETE):请求必须携带令牌, // 可以是 "_csrf" 表单字段(普通表单、multipart 上传), // 也可以是 "X-CSRF-Token" 请求头(AJAX)。不匹配时以 403 终止请求。 // // 令牌与会话绑定,因此既适用于匿名访客(如评论表单), // 也适用于已登录用户。 const ( // CSRFFieldName 是携带令牌的表单字段名。 CSRFFieldName = "_csrf" // CSRFHeaderName 是携带令牌的 HTTP 请求头名(AJAX)。 CSRFHeaderName = "X-CSRF-Token" // CSRFSessionKey 是服务端存储令牌的会话键。 CSRFSessionKey = "csrf_token" // CSRFContextKey 通过 c.Set 将令牌暴露给处理器/模板。 CSRFContextKey = "csrf_token" ) // newCSRFToken 返回 256 位随机十六进制令牌。crypto/rand 失败不可恢复: // 宁可 panic,也不削弱防御。 func newCSRFToken() string { b := make([]byte, 32) if _, err := rand.Read(b); err != nil { panic("csrf: failed to read random bytes: " + err.Error()) } return hex.EncodeToString(b) } // csrfTokensEqual 以常量时间比较两个令牌。 func csrfTokensEqual(a, b string) bool { if len(a) != len(b) { return false } var v byte for i := 0; i < len(a); i++ { v |= a[i] ^ b[i] } return v == 0 } // CSRFProtect 校验不安全请求是否携带与会话匹配的 CSRF 令牌。 // 必须在 sessions 中间件之后注册。 func CSRFProtect() gin.HandlerFunc { return func(c *gin.Context) { session := sessions.Default(c) token, _ := session.Get(CSRFSessionKey).(string) switch c.Request.Method { case http.MethodGet, http.MethodHead, http.MethodOptions, http.MethodTrace: // 安全方法:确保令牌存在,并交给模板层使用。 if token == "" { token = newCSRFToken() session.Set(CSRFSessionKey, token) _ = session.Save() } c.Set(CSRFContextKey, token) c.Next() return } // 不安全方法:要求携带匹配的令牌。 supplied := c.PostForm(CSRFFieldName) if supplied == "" { supplied = c.GetHeader(CSRFHeaderName) } if token == "" || supplied == "" || !csrfTokensEqual(token, supplied) { c.Header("Cache-Control", "no-store") c.Header("Content-Type", "text/plain; charset=utf-8") c.String(http.StatusForbidden, "403 Forbidden: CSRF token missing or invalid ("+strconv.Quote(c.Request.Method)+" "+c.Request.URL.Path+")") c.Abort() return } c.Set(CSRFContextKey, token) c.Next() } }