package auth import ( "crypto/sha256" "encoding/hex" "errors" "strconv" "strings" "time" "github.com/google/uuid" "server/models" "server/pkg/jwtutil" ) // 认证中心错误定义 var ( ErrSessionLimitExceeded = errors.New("同时在线设备数已达上限") ErrSessionRevoked = errors.New("会话已失效,请重新登录") ErrSessionExpired = errors.New("会话已过期,请重新登录") ErrRefreshTokenInvalid = errors.New("刷新令牌无效或已过期") ErrRefreshTokenReused = errors.New("刷新令牌已被使用,疑似重放攻击,已吊销该登录") ) // TokenPair 签发的令牌对 type TokenPair struct { AccessToken string `json:"access_token"` RefreshToken string `json:"refresh_token"` TokenType string `json:"token_type"` ExpiresIn int `json:"expires_in"` Sid string `json:"sid"` } // TokenIssue 签发令牌的入参 type TokenIssue struct { IdentityID uint64 Tid uint64 ClientID string Sid string Username string UserType string Amr string AccessTTL int // 秒 RefreshTTL int // 秒 } // IssueTokens 签发访问令牌与刷新令牌。 // 刷新令牌明文只在本次返回,库中仅存哈希。 // // 令牌中的 user_id 为认证中心 identity_id(yz_auth_identity.id)。 // 业务表(文件/客户/合同/日程等)中的 uid 已由 scripts/uidmigrate 全量迁移为 // 同一套 ID,因此无需任何兼容换算。 func IssueTokens(opt TokenIssue) (*TokenPair, error) { accessTTL := opt.AccessTTL if accessTTL <= 0 { accessTTL = 1800 } refreshTTL := opt.RefreshTTL if refreshTTL <= 0 { refreshTTL = 2592000 } // 全量迁移后,令牌中的 user_id 统一为认证中心的 identity_id userID := int(opt.IdentityID) // userType 需与现有业务接口的判定保持一致: // 后端接口普遍要求 user_type 为 "backend"(租户后台)或 "app"(移动端), // 若签发 "tenant" 会被判为无权访问。这里按应用编码给出默认值。 userType := opt.UserType if userType == "" { userType = "backend" if opt.ClientID == "yz-uniapp" { userType = "app" } } jti := uuid.NewString() access, err := jwtutil.SignToken(jwtutil.TokenOptions{ Alg: jwtutil.AlgRS256, UserID: userID, Username: opt.Username, TenantID: int(opt.Tid), UserType: userType, ClientID: opt.ClientID, Sid: opt.Sid, Amr: opt.Amr, Subject: strconv.FormatUint(opt.IdentityID, 10), Audience: []string{opt.ClientID}, Jti: jti, TTL: time.Duration(accessTTL) * time.Second, }) if err != nil { return nil, err } plain, err := randomToken(32) if err != nil { return nil, err } family := uuid.NewString() now := time.Now() rt := &models.AuthRefreshToken{ TokenHash: hashToken(plain), IdentityID: opt.IdentityID, Tid: opt.Tid, ClientID: opt.ClientID, Sid: opt.Sid, FamilyID: family, ExpiresAt: now.Add(time.Duration(refreshTTL) * time.Second), } if _, err := models.Orm.Insert(rt); err != nil { return nil, err } return &TokenPair{ AccessToken: access, RefreshToken: plain, TokenType: "Bearer", ExpiresIn: accessTTL, Sid: opt.Sid, }, nil } // RefreshTokens 用刷新令牌换取新的令牌对(轮换 + 重放检测)。 // // 安全要点:检测到已使用的刷新令牌再次出现时,判定为重放, // 吊销整个 family(该次登录的全部令牌)。 func RefreshTokens(plain, clientID string) (*TokenPair, error) { h := hashToken(plain) var rt models.AuthRefreshToken if err := models.Orm.QueryTable(new(models.AuthRefreshToken)).Filter("token_hash", h).One(&rt); err != nil { return nil, ErrRefreshTokenInvalid } if rt.ClientID != "" && clientID != "" && rt.ClientID != clientID { return nil, ErrRefreshTokenInvalid } if rt.Revoked != 0 || rt.ExpiresAt.Before(time.Now()) { return nil, ErrRefreshTokenInvalid } if rt.Used != 0 { // 重放:吊销同族全部令牌与该会话 _, _ = models.Orm.QueryTable(new(models.AuthRefreshToken)). Filter("family_id", rt.FamilyID). Update(map[string]interface{}{"revoked": 1}) _ = RevokeSession(rt.Sid, models.RevokeReasonAdmin) return nil, ErrRefreshTokenReused } // 标记已用(保留 Row 以便审计) _, _ = models.Orm.QueryTable(new(models.AuthRefreshToken)). Filter("id", rt.ID). Update(map[string]interface{}{"used": 1}) // 会话校验:已吊销/过期则拒绝续期 s, err := GetSession(rt.Sid) if err != nil { return nil, err } _ = TouchSession(rt.Sid) pair, err := IssueTokens(TokenIssue{ IdentityID: rt.IdentityID, Tid: s.Tid, // 以会话当前企业为准(支持切换企业后刷新) ClientID: rt.ClientID, Sid: rt.Sid, Amr: derefStr(s.Amr), }) if err != nil { return nil, err } // 新令牌继承同一 family,便于后续溯源与整族吊销 if _, err := models.Orm.QueryTable(new(models.AuthRefreshToken)). Filter("token_hash", hashToken(pair.RefreshToken)). Update(map[string]interface{}{"family_id": rt.FamilyID, "rotated_from": h}); err != nil { return nil, err } return pair, nil } // FindLegacyUID 通过「身份 + 企业」在老表 yz_system_tenant_user 中找到对应 uid。 // // 双轨期(老登录与统一认证并行)下业务接口仍以老 uid 识别用户, // 因此签发令牌、返回用户信息时都要换算回老 uid。 // 匹配顺序:账号 → 手机号 → 邮箱;都匹配不到返回 0。 func FindLegacyUID(tid, identityID uint64) uint64 { if tid == 0 || identityID == 0 { return 0 } var bind models.AuthTenantUser if err := models.Orm.QueryTable(new(models.AuthTenantUser)). Filter("tid", tid). Filter("identity_id", identityID). Filter("delete_time__isnull", true). One(&bind); err != nil { return 0 } qs := models.Orm.QueryTable(new(models.SystemTenantUser)).Filter("tid", tid) match := func(field string, value *string) uint64 { if value == nil || strings.TrimSpace(*value) == "" { return 0 } var row models.SystemTenantUser if err := qs.Filter(field, strings.TrimSpace(*value)). Filter("delete_time__isnull", true). One(&row); err == nil { return row.Uid } return 0 } if uid := match("account", bind.Account); uid > 0 { return uid } if uid := match("phone", bind.Phone); uid > 0 { return uid } return match("email", bind.Email) } // RevokeTokenPair 登出:吊销刷新令牌、会话,并把 access token 的 jti 加入黑名单。 func RevokeTokenPair(plain, accessToken, reason string) error { if plain != "" { _, _ = models.Orm.QueryTable(new(models.AuthRefreshToken)). Filter("token_hash", hashToken(plain)). Update(map[string]interface{}{"revoked": 1}) } if accessToken != "" { if claims, err := jwtutil.ParseTokenRaw(accessToken); err == nil { _ = Blacklist(claims.ID, claims.Sid, reason, time.Unix(claims.ExpiresAt.Unix(), 0)) if claims.Sid != "" { _ = RevokeSession(claims.Sid, reason) } } } return nil } // Blacklist 把 jti 加入吊销表 func Blacklist(jti, sid, reason string, expiresAt time.Time) error { if jti == "" { return nil } item := &models.AuthTokenBlacklist{ Jti: jti, ExpiresAt: expiresAt, } if sid != "" { item.Sid = &sid } if reason != "" { item.Reason = &reason } _, err := models.Orm.InsertOrUpdate(item, "jti") return err } // IsBlacklisted 判断 jti 是否已被吊销 func IsBlacklisted(jti string) bool { if jti == "" { return false } return models.Orm.QueryTable(new(models.AuthTokenBlacklist)).Filter("jti", jti).Exist() } func hashToken(plain string) string { sum := sha256.Sum256([]byte(plain)) return hex.EncodeToString(sum[:]) } func derefStr(p *string) string { if p == nil { return "" } return *p }