package middleware import ( "crypto/rand" "encoding/hex" "net/http" "sync" "time" "github.com/gin-gonic/gin" ) // ---- 服务端会话存储(内存)---- // 用随机 token 代替固定字符串 cookie,登出/过期即失效。 // 会话绑定到具体用户 ID(多租户)。 // session 会话数据 type session struct { userID uint expiry time.Time } var ( sessions = make(map[string]session) // token -> 会话 sessionsMu sync.RWMutex ) const ( CookieName = "note_token" // 会话 cookie SessionTTL = 7 * 24 * time.Hour CookieMaxAge = 7 * 24 * 3600 sessionCleanT = 10 ) // NewSession 创建并绑定一个用户会话 func NewSession(userID uint) (string, time.Time) { b := make([]byte, 32) if _, err := rand.Read(b); err != nil { b = []byte(time.Now().Format("20060102150405.000000000")) } token := hex.EncodeToString(b) exp := time.Now().Add(SessionTTL) sessionsMu.Lock() sessions[token] = session{userID: userID, expiry: exp} sessionsMu.Unlock() go cleanExpiredSessions() return token, exp } // RevokeSession 登出时删除会话 func RevokeSession(token string) { sessionsMu.Lock() delete(sessions, token) sessionsMu.Unlock() } // IsValidSession 校验 token 是否有效且未过期 func IsValidSession(token string) bool { if token == "" { return false } sessionsMu.RLock() s, ok := sessions[token] sessionsMu.RUnlock() if !ok { return false } if time.Now().After(s.expiry) { RevokeSession(token) return false } return true } // SessionUserID 获取 token 对应的用户 ID(无效返回 0) func SessionUserID(token string) uint { if token == "" { return 0 } sessionsMu.RLock() s, ok := sessions[token] sessionsMu.RUnlock() if !ok { return 0 } if time.Now().After(s.expiry) { RevokeSession(token) return 0 } return s.userID } // cleanExpiredSessions 定期清理过期会话,防止内存泄漏 func cleanExpiredSessions() { sessionsMu.Lock() defer sessionsMu.Unlock() for token, s := range sessions { if time.Now().After(s.expiry) { delete(sessions, token) } } } // CurrentUser 解析当前登录用户 ID 并写入 context(可为 0 表示游客)。 // 供公开路由/混合路由使用。 func CurrentUser() gin.HandlerFunc { return func(c *gin.Context) { uid := SessionUserID(GetAuthToken(c)) c.Set("user_id", uid) c.Next() } } // AuthRequired 认证中间件:必须登录,否则返回 401。 func AuthRequired() gin.HandlerFunc { return func(c *gin.Context) { token, err := c.Cookie(CookieName) if err != nil || !IsValidSession(token) { c.JSON(http.StatusUnauthorized, gin.H{"code": 401, "message": "请先登录"}) c.Abort() return } uid := SessionUserID(token) c.Set("user_id", uid) c.Next() } } // GetUserID 从 context 取当前用户 ID(游客为 0) func GetUserID(c *gin.Context) uint { if v, ok := c.Get("user_id"); ok { if uid, ok2 := v.(uint); ok2 { return uid } } return 0 } // IsLoggedIn 判断当前请求是否已登录 func IsLoggedIn(c *gin.Context) bool { return GetUserID(c) != 0 } // SetAuthCookie 设置认证 cookie(SameSite=Lax 防 CSRF;https 下应启用 Secure) func SetAuthCookie(c *gin.Context, token string, secure bool) { c.SetSameSite(http.SameSiteLaxMode) c.SetCookie(CookieName, token, CookieMaxAge, "/", "", secure, true) } // ClearAuthCookie 清除认证 cookie func ClearAuthCookie(c *gin.Context) { c.SetSameSite(http.SameSiteLaxMode) c.SetCookie(CookieName, "", -1, "/", "", false, true) } // GetAuthToken 从请求读取 token func GetAuthToken(c *gin.Context) string { token, _ := c.Cookie(CookieName) return token }