merge: admin auth redesign
This commit is contained in:
@@ -15,8 +15,7 @@ logger:
|
|||||||
addSource: false
|
addSource: false
|
||||||
|
|
||||||
auth:
|
auth:
|
||||||
accessTokenTTLHours: 12
|
tokenTTLHours: 12
|
||||||
refreshTokenTTLDays: 7
|
|
||||||
maxFailedAttempts: 5
|
maxFailedAttempts: 5
|
||||||
credentialLockMinute: 15
|
credentialLockMinute: 15
|
||||||
|
|
||||||
|
|||||||
+1
-1
Submodule docs updated: 480f237ef4...2bf6ba451e
@@ -30,20 +30,6 @@ func (c *AuthController) PostLogin() *web.JsonResult {
|
|||||||
return web.JsonData(ret)
|
return web.JsonData(ret)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *AuthController) PostRefresh_token() *web.JsonResult {
|
|
||||||
cfg := config.Current()
|
|
||||||
req := request.RefreshTokenRequest{}
|
|
||||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
|
||||||
return web.JsonError(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
ret, err := services.AuthService.RefreshToken(req.RefreshToken, cfg.Auth, c.Ctx.RemoteAddr(), c.Ctx.GetHeader("User-Agent"))
|
|
||||||
if err != nil {
|
|
||||||
return web.JsonError(err)
|
|
||||||
}
|
|
||||||
return web.JsonData(ret)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *AuthController) GetWxwork_login() {
|
func (c *AuthController) GetWxwork_login() {
|
||||||
loginURL, err := services.WxWorkLoginService.BuildWxWorkLoginURL(c.Ctx.URLParam("next"))
|
loginURL, err := services.WxWorkLoginService.BuildWxWorkLoginURL(c.Ctx.URLParam("next"))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -91,14 +77,7 @@ func (c *AuthController) PostWxwork_exchange() *web.JsonResult {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *AuthController) PostLogout() *web.JsonResult {
|
func (c *AuthController) PostLogout() *web.JsonResult {
|
||||||
req := request.LogoutRequest{}
|
if err := services.AuthService.Logout(c.Ctx.GetHeader("Authorization")); err != nil {
|
||||||
if c.Ctx.GetContentLength() > 0 {
|
|
||||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
|
||||||
return web.JsonError(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := services.AuthService.Logout(c.Ctx.GetHeader("Authorization"), req.RefreshToken); err != nil {
|
|
||||||
return web.JsonError(err)
|
return web.JsonError(err)
|
||||||
}
|
}
|
||||||
return web.JsonSuccess()
|
return web.JsonSuccess()
|
||||||
|
|||||||
@@ -24,7 +24,6 @@ func (c *SessionController) AnyList() *web.JsonResult {
|
|||||||
|
|
||||||
cnd := params.NewPagedSqlCnd(c.Ctx,
|
cnd := params.NewPagedSqlCnd(c.Ctx,
|
||||||
params.QueryFilter{ParamName: "userId"},
|
params.QueryFilter{ParamName: "userId"},
|
||||||
params.QueryFilter{ParamName: "tokenType"},
|
|
||||||
params.QueryFilter{ParamName: "clientType"},
|
params.QueryFilter{ParamName: "clientType"},
|
||||||
).Desc("id")
|
).Desc("id")
|
||||||
list, paging := services.LoginSessionService.FindPageByCnd(cnd)
|
list, paging := services.LoginSessionService.FindPageByCnd(cnd)
|
||||||
@@ -38,7 +37,6 @@ func (c *SessionController) AnyList() *web.JsonResult {
|
|||||||
ID: item.ID,
|
ID: item.ID,
|
||||||
UserID: item.UserID,
|
UserID: item.UserID,
|
||||||
Username: username,
|
Username: username,
|
||||||
TokenType: item.TokenType,
|
|
||||||
ClientType: item.ClientType,
|
ClientType: item.ClientType,
|
||||||
ClientIP: item.ClientIP,
|
ClientIP: item.ClientIP,
|
||||||
UserAgent: item.UserAgent,
|
UserAgent: item.UserAgent,
|
||||||
|
|||||||
+19
-20
@@ -287,31 +287,30 @@ type UserPermission struct {
|
|||||||
AuditFields
|
AuditFields
|
||||||
}
|
}
|
||||||
|
|
||||||
// LoginSession 登录会话或刷新令牌记录。
|
// LoginSession 表示一次后台登录会话。
|
||||||
type LoginSession struct {
|
type LoginSession struct {
|
||||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
ID int64 `gorm:"primaryKey;autoIncrement"` // ID 为登录会话主键。
|
||||||
UserID int64 `gorm:"type:bigint;not null;index"`
|
UserID int64 `gorm:"type:bigint;not null;index"` // UserID 为登录用户 ID。
|
||||||
TokenID string `gorm:"type:varchar(128);not null;uniqueIndex"`
|
Token string `gorm:"type:varchar(128);not null;uniqueIndex"` // Token 为随机不透明登录凭证,使用 ak_ 前缀。
|
||||||
TokenType string `gorm:"type:varchar(20);not null;default:'';index"`
|
ClientType string `gorm:"type:varchar(50);not null;default:'';index"` // ClientType 为客户端类型,后台 Web 端固定为 admin_web。
|
||||||
ClientType string `gorm:"type:varchar(50);not null;default:'';index"`
|
ClientIP string `gorm:"type:varchar(64);not null;default:''"` // ClientIP 为登录请求来源 IP。
|
||||||
ClientIP string `gorm:"type:varchar(64);not null;default:''"`
|
UserAgent string `gorm:"type:varchar(255);not null;default:''"` // UserAgent 为登录请求浏览器或客户端 UA。
|
||||||
UserAgent string `gorm:"type:varchar(255);not null;default:''"`
|
ExpiredAt time.Time `gorm:"type:datetime;not null;index"` // ExpiredAt 为 token 过期时间。
|
||||||
ExpiredAt time.Time `gorm:"type:datetime;not null;index"`
|
RevokedAt *time.Time `gorm:"type:datetime;index"` // RevokedAt 为主动注销或踢下线时间,非空表示已失效。
|
||||||
RevokedAt *time.Time `gorm:"type:datetime;index"`
|
LastSeenAt *time.Time `gorm:"type:datetime"` // LastSeenAt 为最近一次成功鉴权时间。
|
||||||
LastSeenAt *time.Time `gorm:"type:datetime"`
|
|
||||||
AuditFields
|
AuditFields
|
||||||
}
|
}
|
||||||
|
|
||||||
// LoginCredentialLog 登录凭证校验日志。
|
// LoginCredentialLog 记录一次后台登录凭证校验结果。
|
||||||
type LoginCredentialLog struct {
|
type LoginCredentialLog struct {
|
||||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
ID int64 `gorm:"primaryKey;autoIncrement"` // ID 为登录凭证日志主键。
|
||||||
Principal string `gorm:"type:varchar(100);not null;default:'';index"`
|
Principal string `gorm:"type:varchar(100);not null;default:'';index"` // Principal 为用户输入的登录名。
|
||||||
UserID int64 `gorm:"type:bigint;not null;default:0;index"`
|
UserID int64 `gorm:"type:bigint;not null;default:0;index"` // UserID 为匹配到的用户 ID,未匹配时为 0。
|
||||||
Success bool `gorm:"not null;default:false;index"`
|
Success bool `gorm:"not null;default:false;index"` // Success 表示本次凭证校验是否成功。
|
||||||
ClientIP string `gorm:"type:varchar(64);not null;default:''"`
|
ClientIP string `gorm:"type:varchar(64);not null;default:''"` // ClientIP 为登录请求来源 IP。
|
||||||
UserAgent string `gorm:"type:varchar(255);not null;default:''"`
|
UserAgent string `gorm:"type:varchar(255);not null;default:''"` // UserAgent 为登录请求浏览器或客户端 UA。
|
||||||
Reason string `gorm:"type:varchar(255);not null;default:''"`
|
Reason string `gorm:"type:varchar(255);not null;default:''"` // Reason 为校验结果原因。
|
||||||
CreatedAt time.Time `gorm:"type:datetime;not null;index"`
|
CreatedAt time.Time `gorm:"type:datetime;not null;index"` // CreatedAt 为日志创建时间。
|
||||||
}
|
}
|
||||||
|
|
||||||
// Asset 存储的文件资源,如上传的附件等。
|
// Asset 存储的文件资源,如上传的附件等。
|
||||||
|
|||||||
@@ -55,8 +55,7 @@ type LoggerConfig struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type AuthConfig struct {
|
type AuthConfig struct {
|
||||||
AccessTokenTTLHours int `yaml:"accessTokenTTLHours"`
|
TokenTTLHours int `yaml:"tokenTTLHours"`
|
||||||
RefreshTokenTTLDays int `yaml:"refreshTokenTTLDays"`
|
|
||||||
MaxFailedAttempts int `yaml:"maxFailedAttempts"`
|
MaxFailedAttempts int `yaml:"maxFailedAttempts"`
|
||||||
CredentialLockMinute int `yaml:"credentialLockMinute"`
|
CredentialLockMinute int `yaml:"credentialLockMinute"`
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -8,13 +8,7 @@ const (
|
|||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
AccessTokenPrefix = "atk_"
|
AuthTokenPrefix = "ak_"
|
||||||
RefreshTokenPrefix = "rtk_"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
TokenTypeAccess = "access"
|
|
||||||
TokenTypeRefresh = "refresh"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
|
|||||||
@@ -5,14 +5,6 @@ type LoginRequest struct {
|
|||||||
Password string `json:"password"`
|
Password string `json:"password"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type RefreshTokenRequest struct {
|
|
||||||
RefreshToken string `json:"refreshToken"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type LogoutRequest struct {
|
|
||||||
RefreshToken string `json:"refreshToken"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type WxWorkExchangeRequest struct {
|
type WxWorkExchangeRequest struct {
|
||||||
Ticket string `json:"ticket"`
|
Ticket string `json:"ticket"`
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -47,12 +47,11 @@ type CreateUserResultResponse struct {
|
|||||||
type SessionResponse struct {
|
type SessionResponse struct {
|
||||||
ID int64 `json:"id"`
|
ID int64 `json:"id"`
|
||||||
UserID int64 `json:"userId"`
|
UserID int64 `json:"userId"`
|
||||||
Username string `json:"username,omitempty"`
|
Username string `json:"username"`
|
||||||
TokenType string `json:"tokenType"`
|
|
||||||
ClientType string `json:"clientType"`
|
ClientType string `json:"clientType"`
|
||||||
ClientIP string `json:"clientIp"`
|
ClientIP string `json:"clientIp"`
|
||||||
UserAgent string `json:"userAgent"`
|
UserAgent string `json:"userAgent"`
|
||||||
ExpiredAt string `json:"expiredAt"`
|
ExpiredAt string `json:"expiredAt"`
|
||||||
RevokedAt string `json:"revokedAt,omitempty"`
|
RevokedAt string `json:"revokedAt"`
|
||||||
LastSeenAt string `json:"lastSeenAt,omitempty"`
|
LastSeenAt string `json:"lastSeenAt"`
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -12,10 +12,9 @@ type AuthUserResponse struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type LoginResponse struct {
|
type LoginResponse struct {
|
||||||
AccessToken string `json:"accessToken"`
|
AccessToken string `json:"accessToken"`
|
||||||
RefreshToken string `json:"refreshToken"`
|
ExpiresAt string `json:"expiresAt"`
|
||||||
ExpiresAt string `json:"expiresAt"`
|
User *AuthUserResponse `json:"user"`
|
||||||
User *AuthUserResponse `json:"user"`
|
Permissions []string `json:"permissions"`
|
||||||
Permissions []string `json:"permissions"`
|
Roles []string `json:"roles"`
|
||||||
Roles []string `json:"roles"`
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,12 +3,13 @@ package errorsx
|
|||||||
import "github.com/mlogclub/simple/web"
|
import "github.com/mlogclub/simple/web"
|
||||||
|
|
||||||
const (
|
const (
|
||||||
CodeInvalidParam = 1000
|
CodeInvalidParam = 1000
|
||||||
CodeBusinessError = 2000
|
CodeBusinessError = 2000
|
||||||
CodeAuthUnauthorized = 3000
|
CodeAuthUnauthorized = 3000
|
||||||
CodeAuthForbidden = 3001
|
CodeAuthForbidden = 3001
|
||||||
CodeAuthInvalidToken = 3002
|
CodeAuthInvalidToken = 3002
|
||||||
CodeAuthInvalidAccount = 3003
|
CodeAuthInvalidAccount = 3003
|
||||||
|
CodeAuthCredentialLocked = 3004
|
||||||
)
|
)
|
||||||
|
|
||||||
func InvalidParam(message string) error {
|
func InvalidParam(message string) error {
|
||||||
@@ -34,3 +35,7 @@ func InvalidToken(message string) error {
|
|||||||
func InvalidAccount(message string) error {
|
func InvalidAccount(message string) error {
|
||||||
return web.NewError(CodeAuthInvalidAccount, message)
|
return web.NewError(CodeAuthInvalidAccount, message)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func CredentialLocked(message string) error {
|
||||||
|
return web.NewError(CodeAuthCredentialLocked, message)
|
||||||
|
}
|
||||||
|
|||||||
@@ -81,18 +81,24 @@ func (s *authService) RequirePermission(ctx iris.Context, permission constants.P
|
|||||||
|
|
||||||
func (s *authService) Login(req request.LoginRequest, authCfg config.AuthConfig, clientIP, userAgent string) (*response.LoginResponse, error) {
|
func (s *authService) Login(req request.LoginRequest, authCfg config.AuthConfig, clientIP, userAgent string) (*response.LoginResponse, error) {
|
||||||
username := strings.TrimSpace(req.Username)
|
username := strings.TrimSpace(req.Username)
|
||||||
|
principal := normalizeLoginPrincipal(username)
|
||||||
password := req.Password
|
password := req.Password
|
||||||
if username == "" || strings.TrimSpace(password) == "" {
|
if username == "" || strings.TrimSpace(password) == "" {
|
||||||
return nil, errorsx.InvalidParam("用户名和密码不能为空")
|
return nil, errorsx.InvalidParam("用户名和密码不能为空")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if s.isCredentialLocked(principal, authCfg) {
|
||||||
|
_ = s.createLoginCredentialLog(principal, 0, false, clientIP, userAgent, "credential locked")
|
||||||
|
return nil, errorsx.CredentialLocked("登录失败次数过多,请稍后再试")
|
||||||
|
}
|
||||||
|
|
||||||
user := UserService.GetByUsername(username)
|
user := UserService.GetByUsername(username)
|
||||||
if user == nil || user.Status != enums.StatusOk {
|
if user == nil || user.Status != enums.StatusOk {
|
||||||
_ = s.createLoginCredentialLog(username, 0, false, clientIP, userAgent, "user not found")
|
_ = s.createLoginCredentialLog(principal, 0, false, clientIP, userAgent, "user not found")
|
||||||
return nil, errorsx.InvalidAccount("用户名或密码错误")
|
return nil, errorsx.InvalidAccount("用户名或密码错误")
|
||||||
}
|
}
|
||||||
if strs.IsBlank(user.Password) || bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(password)) != nil {
|
if strs.IsBlank(user.Password) || bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(password)) != nil {
|
||||||
_ = s.createLoginCredentialLog(username, user.ID, false, clientIP, userAgent, "password mismatch")
|
_ = s.createLoginCredentialLog(principal, user.ID, false, clientIP, userAgent, "password mismatch")
|
||||||
return nil, errorsx.InvalidAccount("用户名或密码错误")
|
return nil, errorsx.InvalidAccount("用户名或密码错误")
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -117,60 +123,15 @@ func (s *authService) Login(req request.LoginRequest, authCfg config.AuthConfig,
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
_ = s.createLoginCredentialLog(username, user.ID, true, clientIP, userAgent, "")
|
_ = s.createLoginCredentialLog(principal, user.ID, true, clientIP, userAgent, "")
|
||||||
return ret, nil
|
return ret, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *authService) RefreshToken(refreshToken string, authCfg config.AuthConfig, clientIP, userAgent string) (*response.LoginResponse, error) {
|
func (s *authService) Logout(accessToken string) error {
|
||||||
session, err := s.validateSessionToken(refreshToken, constants.TokenTypeRefresh)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
if session.RevokedAt != nil {
|
|
||||||
return nil, errorsx.InvalidToken("refresh token 已失效")
|
|
||||||
}
|
|
||||||
|
|
||||||
user := UserService.Get(session.UserID)
|
|
||||||
if user == nil || user.Status != enums.StatusOk {
|
|
||||||
return nil, errorsx.Unauthorized("用户不存在或已被禁用")
|
|
||||||
}
|
|
||||||
|
|
||||||
var ret *response.LoginResponse
|
|
||||||
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
|
|
||||||
var dbErr error
|
|
||||||
if dbErr = repositories.LoginSessionRepository.Updates(ctx.Tx, session.ID, map[string]any{
|
|
||||||
"revoked_at": time.Now(),
|
|
||||||
"update_user_id": user.ID,
|
|
||||||
"update_user_name": user.Username,
|
|
||||||
"updated_at": time.Now(),
|
|
||||||
}); dbErr != nil {
|
|
||||||
return dbErr
|
|
||||||
}
|
|
||||||
if ret, dbErr = s.issueTokens(ctx, user, clientIP, userAgent, authCfg); dbErr != nil {
|
|
||||||
return dbErr
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return ret, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *authService) Logout(accessToken, refreshToken string) error {
|
|
||||||
accessToken = s.extractBearerToken(accessToken)
|
accessToken = s.extractBearerToken(accessToken)
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
if accessToken != "" {
|
if accessToken != "" {
|
||||||
if session := LoginSessionService.FindOne(sqls.NewCnd().Eq("token_id", accessToken).Eq("token_type", constants.TokenTypeAccess)); session != nil && session.RevokedAt == nil {
|
if session := LoginSessionService.FindOne(sqls.NewCnd().Eq("token", accessToken)); session != nil && session.RevokedAt == nil {
|
||||||
if err := LoginSessionService.Updates(session.ID, map[string]any{
|
|
||||||
"revoked_at": now,
|
|
||||||
"updated_at": now,
|
|
||||||
}); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if refreshToken != "" {
|
|
||||||
if session := LoginSessionService.FindOne(sqls.NewCnd().Eq("token_id", refreshToken).Eq("token_type", constants.TokenTypeRefresh)); session != nil && session.RevokedAt == nil {
|
|
||||||
if err := LoginSessionService.Updates(session.ID, map[string]any{
|
if err := LoginSessionService.Updates(session.ID, map[string]any{
|
||||||
"revoked_at": now,
|
"revoked_at": now,
|
||||||
"updated_at": now,
|
"updated_at": now,
|
||||||
@@ -195,7 +156,7 @@ func (s *authService) Authenticate(ctx iris.Context) (*dto.AuthPrincipal, error)
|
|||||||
return nil, errorsx.Unauthorized("未登录或登录已过期")
|
return nil, errorsx.Unauthorized("未登录或登录已过期")
|
||||||
}
|
}
|
||||||
|
|
||||||
session, err := s.validateSessionToken(token, constants.TokenTypeAccess)
|
session, err := s.validateSessionToken(token)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -262,48 +223,20 @@ func (s *authService) issueTokens(ctx *sqls.TxContext, user *models.User, client
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
accessTTL, refreshTTL := s.resolveTokenTTL(authCfg)
|
tokenTTL := s.resolveTokenTTL(authCfg)
|
||||||
accessToken, err := randomToken(constants.AccessTokenPrefix)
|
accessToken, err := randomToken(constants.AuthTokenPrefix)
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
refreshToken, err := randomToken(constants.RefreshTokenPrefix)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
// accessSession
|
|
||||||
if err := repositories.LoginSessionRepository.Create(ctx.Tx, &models.LoginSession{
|
if err := repositories.LoginSessionRepository.Create(ctx.Tx, &models.LoginSession{
|
||||||
UserID: user.ID,
|
UserID: user.ID,
|
||||||
TokenID: accessToken,
|
Token: accessToken,
|
||||||
TokenType: constants.TokenTypeAccess,
|
|
||||||
ClientType: constants.ClientTypeAdminWeb,
|
ClientType: constants.ClientTypeAdminWeb,
|
||||||
ClientIP: clientIP,
|
ClientIP: clientIP,
|
||||||
UserAgent: userAgent,
|
UserAgent: userAgent,
|
||||||
ExpiredAt: now.Add(accessTTL),
|
ExpiredAt: now.Add(tokenTTL),
|
||||||
LastSeenAt: &now,
|
|
||||||
AuditFields: models.AuditFields{
|
|
||||||
CreatedAt: now,
|
|
||||||
CreateUserID: user.ID,
|
|
||||||
CreateUserName: user.Username,
|
|
||||||
UpdatedAt: now,
|
|
||||||
UpdateUserID: user.ID,
|
|
||||||
UpdateUserName: user.Username,
|
|
||||||
},
|
|
||||||
}); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// refreshSession
|
|
||||||
if err := repositories.LoginSessionRepository.Create(ctx.Tx, &models.LoginSession{
|
|
||||||
UserID: user.ID,
|
|
||||||
TokenID: refreshToken,
|
|
||||||
TokenType: constants.TokenTypeRefresh,
|
|
||||||
ClientType: constants.ClientTypeAdminWeb,
|
|
||||||
ClientIP: clientIP,
|
|
||||||
UserAgent: userAgent,
|
|
||||||
ExpiredAt: now.Add(refreshTTL),
|
|
||||||
LastSeenAt: &now,
|
LastSeenAt: &now,
|
||||||
AuditFields: models.AuditFields{
|
AuditFields: models.AuditFields{
|
||||||
CreatedAt: now,
|
CreatedAt: now,
|
||||||
@@ -318,9 +251,8 @@ func (s *authService) issueTokens(ctx *sqls.TxContext, user *models.User, client
|
|||||||
}
|
}
|
||||||
|
|
||||||
return &response.LoginResponse{
|
return &response.LoginResponse{
|
||||||
AccessToken: accessToken,
|
AccessToken: accessToken,
|
||||||
RefreshToken: refreshToken,
|
ExpiresAt: now.Add(tokenTTL).Format(time.DateTime),
|
||||||
ExpiresAt: now.Add(accessTTL).Format(time.DateTime),
|
|
||||||
User: &response.AuthUserResponse{
|
User: &response.AuthUserResponse{
|
||||||
ID: user.ID,
|
ID: user.ID,
|
||||||
Username: user.Username,
|
Username: user.Username,
|
||||||
@@ -334,33 +266,27 @@ func (s *authService) issueTokens(ctx *sqls.TxContext, user *models.User, client
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *authService) resolveTokenTTL(authCfg config.AuthConfig) (time.Duration, time.Duration) {
|
func (s *authService) resolveTokenTTL(authCfg config.AuthConfig) time.Duration {
|
||||||
accessTTL := 12 * time.Hour
|
tokenTTL := 12 * time.Hour
|
||||||
refreshTTL := 7 * 24 * time.Hour
|
if authCfg.TokenTTLHours > 0 {
|
||||||
if authCfg.AccessTokenTTLHours > 0 {
|
tokenTTL = time.Duration(authCfg.TokenTTLHours) * time.Hour
|
||||||
accessTTL = time.Duration(authCfg.AccessTokenTTLHours) * time.Hour
|
|
||||||
}
|
}
|
||||||
if authCfg.RefreshTokenTTLDays > 0 {
|
return tokenTTL
|
||||||
refreshTTL = time.Duration(authCfg.RefreshTokenTTLDays) * 24 * time.Hour
|
|
||||||
}
|
|
||||||
return accessTTL, refreshTTL
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *authService) validateSessionToken(token, tokenType string) (*models.LoginSession, error) {
|
func (s *authService) validateSessionToken(token string) (*models.LoginSession, error) {
|
||||||
if strings.TrimSpace(token) == "" {
|
if strings.TrimSpace(token) == "" {
|
||||||
return nil, errorsx.InvalidToken("token 不能为空")
|
return nil, errorsx.Unauthorized("未登录或登录已过期")
|
||||||
}
|
}
|
||||||
session := LoginSessionService.FindOne(sqls.NewCnd().
|
session := LoginSessionService.FindOne(sqls.NewCnd().Eq("token", token))
|
||||||
Eq("token_id", token).
|
|
||||||
Eq("token_type", tokenType))
|
|
||||||
if session == nil {
|
if session == nil {
|
||||||
return nil, errorsx.InvalidToken("token 无效")
|
return nil, errorsx.InvalidToken("登录凭证无效")
|
||||||
}
|
}
|
||||||
if session.RevokedAt != nil {
|
if session.RevokedAt != nil {
|
||||||
return nil, errorsx.InvalidToken("token 已失效")
|
return nil, errorsx.InvalidToken("登录凭证已失效")
|
||||||
}
|
}
|
||||||
if time.Now().After(session.ExpiredAt) {
|
if time.Now().After(session.ExpiredAt) {
|
||||||
return nil, errorsx.InvalidToken("token 已过期")
|
return nil, errorsx.InvalidToken("登录凭证已过期")
|
||||||
}
|
}
|
||||||
return session, nil
|
return session, nil
|
||||||
}
|
}
|
||||||
@@ -480,6 +406,27 @@ func (s *authService) createLoginCredentialLog(principal string, userID int64, s
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *authService) isCredentialLocked(principal string, authCfg config.AuthConfig) bool {
|
||||||
|
maxFailedAttempts := authCfg.MaxFailedAttempts
|
||||||
|
if maxFailedAttempts <= 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
lockMinute := authCfg.CredentialLockMinute
|
||||||
|
if lockMinute <= 0 {
|
||||||
|
lockMinute = 15
|
||||||
|
}
|
||||||
|
since := time.Now().Add(-time.Duration(lockMinute) * time.Minute)
|
||||||
|
return LoginCredentialLogService.Count(sqls.NewCnd().
|
||||||
|
Eq("principal", normalizeLoginPrincipal(principal)).
|
||||||
|
Eq("success", false).
|
||||||
|
NotEq("reason", "credential locked").
|
||||||
|
Where("created_at >= ?", since)) >= int64(maxFailedAttempts)
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeLoginPrincipal(principal string) string {
|
||||||
|
return strings.ToLower(strings.TrimSpace(principal))
|
||||||
|
}
|
||||||
|
|
||||||
func randomToken(prefix string) (string, error) {
|
func randomToken(prefix string) (string, error) {
|
||||||
buf := make([]byte, 24)
|
buf := make([]byte, 24)
|
||||||
if _, err := rand.Read(buf); err != nil {
|
if _, err := rand.Read(buf); err != nil {
|
||||||
|
|||||||
@@ -1,6 +1,24 @@
|
|||||||
package services
|
package services
|
||||||
|
|
||||||
import "testing"
|
import (
|
||||||
|
"errors"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"cs-agent/internal/models"
|
||||||
|
"cs-agent/internal/pkg/config"
|
||||||
|
"cs-agent/internal/pkg/dto/request"
|
||||||
|
"cs-agent/internal/pkg/enums"
|
||||||
|
"cs-agent/internal/pkg/errorsx"
|
||||||
|
|
||||||
|
"github.com/glebarez/sqlite"
|
||||||
|
"github.com/mlogclub/simple/sqls"
|
||||||
|
"github.com/mlogclub/simple/web"
|
||||||
|
"golang.org/x/crypto/bcrypt"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
"gorm.io/gorm/schema"
|
||||||
|
)
|
||||||
|
|
||||||
func TestExtractBearerToken(t *testing.T) {
|
func TestExtractBearerToken(t *testing.T) {
|
||||||
svc := newAuthService()
|
svc := newAuthService()
|
||||||
@@ -13,3 +31,409 @@ func TestExtractBearerToken(t *testing.T) {
|
|||||||
t.Fatalf("expected raw token to be rejected by bearer extractor, got %q", got)
|
t.Fatalf("expected raw token to be rejected by bearer extractor, got %q", got)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestAuthServiceLoginCreatesSingleAccessSession(t *testing.T) {
|
||||||
|
db := setupAuthServiceTestDB(t)
|
||||||
|
user := createAuthTestUser(t, db, "admin", "secret")
|
||||||
|
svc := newAuthService()
|
||||||
|
|
||||||
|
ret, err := svc.Login(request.LoginRequest{
|
||||||
|
Username: " admin ",
|
||||||
|
Password: "secret",
|
||||||
|
}, config.AuthConfig{TokenTTLHours: 2, MaxFailedAttempts: 5, CredentialLockMinute: 15}, "127.0.0.1", "go-test")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("login failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if ret.AccessToken == "" || !strings.HasPrefix(ret.AccessToken, "ak_") {
|
||||||
|
t.Fatalf("expected ak_ access token, got %q", ret.AccessToken)
|
||||||
|
}
|
||||||
|
if ret.ExpiresAt == "" {
|
||||||
|
t.Fatal("expected expiresAt to be returned")
|
||||||
|
}
|
||||||
|
|
||||||
|
var sessions []models.LoginSession
|
||||||
|
if err := db.Find(&sessions).Error; err != nil {
|
||||||
|
t.Fatalf("query login sessions: %v", err)
|
||||||
|
}
|
||||||
|
if len(sessions) != 1 {
|
||||||
|
t.Fatalf("expected exactly one session, got %d", len(sessions))
|
||||||
|
}
|
||||||
|
if sessions[0].Token != ret.AccessToken {
|
||||||
|
t.Fatalf("expected session token %q, got %q", ret.AccessToken, sessions[0].Token)
|
||||||
|
}
|
||||||
|
if sessions[0].UserID != user.ID {
|
||||||
|
t.Fatalf("expected session user %d, got %d", user.ID, sessions[0].UserID)
|
||||||
|
}
|
||||||
|
if sessions[0].ClientType != "admin_web" {
|
||||||
|
t.Fatalf("expected admin_web client type, got %q", sessions[0].ClientType)
|
||||||
|
}
|
||||||
|
|
||||||
|
logs := findCredentialLogs(t, db)
|
||||||
|
if len(logs) != 1 {
|
||||||
|
t.Fatalf("expected one credential log, got %d", len(logs))
|
||||||
|
}
|
||||||
|
if !logs[0].Success || logs[0].Principal != "admin" || logs[0].UserID != user.ID {
|
||||||
|
t.Fatalf("unexpected success credential log: %+v", logs[0])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAuthServiceLoginFailureWritesCredentialLogs(t *testing.T) {
|
||||||
|
db := setupAuthServiceTestDB(t)
|
||||||
|
createAuthTestUser(t, db, "admin", "secret")
|
||||||
|
svc := newAuthService()
|
||||||
|
authCfg := config.AuthConfig{TokenTTLHours: 2, MaxFailedAttempts: 5, CredentialLockMinute: 15}
|
||||||
|
|
||||||
|
if _, err := svc.Login(request.LoginRequest{Username: "missing", Password: "secret"}, authCfg, "127.0.0.1", "go-test"); !hasCode(err, errorsx.CodeAuthInvalidAccount) {
|
||||||
|
t.Fatalf("expected invalid account for missing user, got %v", err)
|
||||||
|
}
|
||||||
|
if _, err := svc.Login(request.LoginRequest{Username: "admin", Password: "wrong"}, authCfg, "127.0.0.1", "go-test"); !hasCode(err, errorsx.CodeAuthInvalidAccount) {
|
||||||
|
t.Fatalf("expected invalid account for password mismatch, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
logs := findCredentialLogs(t, db)
|
||||||
|
if len(logs) != 2 {
|
||||||
|
t.Fatalf("expected two credential logs, got %d", len(logs))
|
||||||
|
}
|
||||||
|
if logs[0].Reason != "user not found" || logs[0].Success {
|
||||||
|
t.Fatalf("unexpected missing-user log: %+v", logs[0])
|
||||||
|
}
|
||||||
|
if logs[1].Reason != "password mismatch" || logs[1].Success {
|
||||||
|
t.Fatalf("unexpected password-mismatch log: %+v", logs[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAuthServiceLoginCredentialLockout(t *testing.T) {
|
||||||
|
db := setupAuthServiceTestDB(t)
|
||||||
|
user := createAuthTestUser(t, db, "admin", "secret")
|
||||||
|
now := time.Now()
|
||||||
|
for i := 0; i < 2; i++ {
|
||||||
|
if err := db.Create(&models.LoginCredentialLog{
|
||||||
|
Principal: "admin",
|
||||||
|
UserID: user.ID,
|
||||||
|
Success: false,
|
||||||
|
Reason: "password mismatch",
|
||||||
|
CreatedAt: now.Add(-time.Duration(i+1) * time.Minute),
|
||||||
|
}).Error; err != nil {
|
||||||
|
t.Fatalf("seed credential log: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := db.Create(&models.LoginCredentialLog{
|
||||||
|
Principal: "admin",
|
||||||
|
UserID: user.ID,
|
||||||
|
Success: false,
|
||||||
|
Reason: "password mismatch",
|
||||||
|
CreatedAt: now.Add(-30 * time.Minute),
|
||||||
|
}).Error; err != nil {
|
||||||
|
t.Fatalf("seed old credential log: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
svc := newAuthService()
|
||||||
|
_, err := svc.Login(request.LoginRequest{Username: "admin", Password: "secret"}, config.AuthConfig{
|
||||||
|
TokenTTLHours: 2,
|
||||||
|
MaxFailedAttempts: 2,
|
||||||
|
CredentialLockMinute: 15,
|
||||||
|
}, "127.0.0.1", "go-test")
|
||||||
|
if !hasCode(err, errorsx.CodeAuthCredentialLocked) {
|
||||||
|
t.Fatalf("expected credential locked error, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var lockedLog models.LoginCredentialLog
|
||||||
|
if err := db.Order("id DESC").Take(&lockedLog).Error; err != nil {
|
||||||
|
t.Fatalf("query latest credential log: %v", err)
|
||||||
|
}
|
||||||
|
if lockedLog.Reason != "credential locked" || lockedLog.Success {
|
||||||
|
t.Fatalf("unexpected locked credential log: %+v", lockedLog)
|
||||||
|
}
|
||||||
|
|
||||||
|
var sessionCount int64
|
||||||
|
if err := db.Model(&models.LoginSession{}).Count(&sessionCount).Error; err != nil {
|
||||||
|
t.Fatalf("count sessions: %v", err)
|
||||||
|
}
|
||||||
|
if sessionCount != 0 {
|
||||||
|
t.Fatalf("expected no session while credential locked, got %d", sessionCount)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAuthServiceCredentialLockoutDoesNotExtendWhileLocked(t *testing.T) {
|
||||||
|
db := setupAuthServiceTestDB(t)
|
||||||
|
createAuthTestUser(t, db, "admin", "secret")
|
||||||
|
now := time.Now()
|
||||||
|
entries := []models.LoginCredentialLog{
|
||||||
|
{
|
||||||
|
Principal: "admin",
|
||||||
|
UserID: 1,
|
||||||
|
Success: false,
|
||||||
|
Reason: "password mismatch",
|
||||||
|
CreatedAt: now.Add(-2 * time.Minute),
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Principal: "admin",
|
||||||
|
UserID: 0,
|
||||||
|
Success: false,
|
||||||
|
Reason: "credential locked",
|
||||||
|
CreatedAt: now.Add(-1 * time.Minute),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if err := db.Create(&entries).Error; err != nil {
|
||||||
|
t.Fatalf("seed credential logs: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ret, err := newAuthService().Login(request.LoginRequest{Username: "admin", Password: "secret"}, config.AuthConfig{
|
||||||
|
TokenTTLHours: 2,
|
||||||
|
MaxFailedAttempts: 2,
|
||||||
|
CredentialLockMinute: 15,
|
||||||
|
}, "127.0.0.1", "go-test")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected locked attempt logs not to extend lockout, got %v", err)
|
||||||
|
}
|
||||||
|
if ret == nil || !strings.HasPrefix(ret.AccessToken, "ak_") {
|
||||||
|
t.Fatalf("expected login response with ak_ token, got %+v", ret)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAuthServiceCredentialLockoutNormalizesPrincipalCase(t *testing.T) {
|
||||||
|
db := setupAuthServiceTestDB(t)
|
||||||
|
createAuthTestUser(t, db, "admin", "secret")
|
||||||
|
if err := db.Create(&models.LoginCredentialLog{
|
||||||
|
Principal: "admin",
|
||||||
|
UserID: 1,
|
||||||
|
Success: false,
|
||||||
|
Reason: "password mismatch",
|
||||||
|
CreatedAt: time.Now().Add(-time.Minute),
|
||||||
|
}).Error; err != nil {
|
||||||
|
t.Fatalf("seed credential log: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := newAuthService().Login(request.LoginRequest{Username: "ADMIN", Password: "secret"}, config.AuthConfig{
|
||||||
|
TokenTTLHours: 2,
|
||||||
|
MaxFailedAttempts: 1,
|
||||||
|
CredentialLockMinute: 15,
|
||||||
|
}, "127.0.0.1", "go-test")
|
||||||
|
if !hasCode(err, errorsx.CodeAuthCredentialLocked) {
|
||||||
|
t.Fatalf("expected normalized principal to be locked, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var lockedLog models.LoginCredentialLog
|
||||||
|
if err := db.Order("id DESC").Take(&lockedLog).Error; err != nil {
|
||||||
|
t.Fatalf("query latest credential log: %v", err)
|
||||||
|
}
|
||||||
|
if lockedLog.Principal != "admin" || lockedLog.Reason != "credential locked" {
|
||||||
|
t.Fatalf("unexpected locked log: %+v", lockedLog)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAuthServiceCredentialLockoutDisabledWhenMaxAttemptsNonPositive(t *testing.T) {
|
||||||
|
db := setupAuthServiceTestDB(t)
|
||||||
|
user := createAuthTestUser(t, db, "admin", "secret")
|
||||||
|
now := time.Now()
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
if err := db.Create(&models.LoginCredentialLog{
|
||||||
|
Principal: "admin",
|
||||||
|
UserID: user.ID,
|
||||||
|
Success: false,
|
||||||
|
Reason: "credential locked",
|
||||||
|
CreatedAt: now.Add(-time.Duration(i+1) * time.Minute),
|
||||||
|
}).Error; err != nil {
|
||||||
|
t.Fatalf("seed credential log: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
ret, err := newAuthService().Login(request.LoginRequest{Username: "admin", Password: "secret"}, config.AuthConfig{
|
||||||
|
TokenTTLHours: 2,
|
||||||
|
MaxFailedAttempts: 0,
|
||||||
|
CredentialLockMinute: 15,
|
||||||
|
}, "127.0.0.1", "go-test")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected lockout to be disabled, got %v", err)
|
||||||
|
}
|
||||||
|
if ret == nil || ret.AccessToken == "" {
|
||||||
|
t.Fatalf("expected login response with access token, got %+v", ret)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidateSessionTokenStates(t *testing.T) {
|
||||||
|
db := setupAuthServiceTestDB(t)
|
||||||
|
svc := newAuthService()
|
||||||
|
now := time.Now()
|
||||||
|
|
||||||
|
if _, err := svc.validateSessionToken(" "); !hasCode(err, errorsx.CodeAuthUnauthorized) {
|
||||||
|
t.Fatalf("expected unauthorized for empty token, got %v", err)
|
||||||
|
}
|
||||||
|
if _, err := svc.validateSessionToken("missing"); !hasCode(err, errorsx.CodeAuthInvalidToken) {
|
||||||
|
t.Fatalf("expected invalid token for missing session, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
revokedAt := now
|
||||||
|
if err := db.Create(&models.LoginSession{
|
||||||
|
UserID: 1,
|
||||||
|
Token: "ak_revoked",
|
||||||
|
ClientType: "admin_web",
|
||||||
|
ExpiredAt: now.Add(time.Hour),
|
||||||
|
RevokedAt: &revokedAt,
|
||||||
|
AuditFields: models.AuditFields{
|
||||||
|
CreatedAt: now,
|
||||||
|
UpdatedAt: now,
|
||||||
|
},
|
||||||
|
}).Error; err != nil {
|
||||||
|
t.Fatalf("seed revoked session: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := svc.validateSessionToken("ak_revoked"); !hasCode(err, errorsx.CodeAuthInvalidToken) {
|
||||||
|
t.Fatalf("expected invalid token for revoked session, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := db.Create(&models.LoginSession{
|
||||||
|
UserID: 1,
|
||||||
|
Token: "ak_expired",
|
||||||
|
ClientType: "admin_web",
|
||||||
|
ExpiredAt: now.Add(-time.Hour),
|
||||||
|
AuditFields: models.AuditFields{
|
||||||
|
CreatedAt: now,
|
||||||
|
UpdatedAt: now,
|
||||||
|
},
|
||||||
|
}).Error; err != nil {
|
||||||
|
t.Fatalf("seed expired session: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := svc.validateSessionToken("ak_expired"); !hasCode(err, errorsx.CodeAuthInvalidToken) {
|
||||||
|
t.Fatalf("expected invalid token for expired session, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := db.Create(&models.LoginSession{
|
||||||
|
UserID: 1,
|
||||||
|
Token: "ak_valid",
|
||||||
|
ClientType: "admin_web",
|
||||||
|
ExpiredAt: now.Add(time.Hour),
|
||||||
|
AuditFields: models.AuditFields{
|
||||||
|
CreatedAt: now,
|
||||||
|
UpdatedAt: now,
|
||||||
|
},
|
||||||
|
}).Error; err != nil {
|
||||||
|
t.Fatalf("seed valid session: %v", err)
|
||||||
|
}
|
||||||
|
session, err := svc.validateSessionToken("ak_valid")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected valid session token, got %v", err)
|
||||||
|
}
|
||||||
|
if session.Token != "ak_valid" {
|
||||||
|
t.Fatalf("expected valid session token ak_valid, got %q", session.Token)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAuthServiceLogoutRevokesCurrentTokenOnly(t *testing.T) {
|
||||||
|
db := setupAuthServiceTestDB(t)
|
||||||
|
now := time.Now()
|
||||||
|
sessions := []models.LoginSession{
|
||||||
|
{
|
||||||
|
UserID: 1,
|
||||||
|
Token: "ak_current",
|
||||||
|
ClientType: "admin_web",
|
||||||
|
ExpiredAt: now.Add(time.Hour),
|
||||||
|
AuditFields: models.AuditFields{
|
||||||
|
CreatedAt: now,
|
||||||
|
UpdatedAt: now,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
UserID: 1,
|
||||||
|
Token: "ak_other",
|
||||||
|
ClientType: "admin_web",
|
||||||
|
ExpiredAt: now.Add(time.Hour),
|
||||||
|
AuditFields: models.AuditFields{
|
||||||
|
CreatedAt: now,
|
||||||
|
UpdatedAt: now,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if err := db.Create(&sessions).Error; err != nil {
|
||||||
|
t.Fatalf("seed sessions: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := newAuthService().Logout("Bearer ak_current"); err != nil {
|
||||||
|
t.Fatalf("logout failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var current models.LoginSession
|
||||||
|
if err := db.Take(¤t, "token = ?", "ak_current").Error; err != nil {
|
||||||
|
t.Fatalf("query current session: %v", err)
|
||||||
|
}
|
||||||
|
if current.RevokedAt == nil {
|
||||||
|
t.Fatal("expected current session to be revoked")
|
||||||
|
}
|
||||||
|
var other models.LoginSession
|
||||||
|
if err := db.Take(&other, "token = ?", "ak_other").Error; err != nil {
|
||||||
|
t.Fatalf("query other session: %v", err)
|
||||||
|
}
|
||||||
|
if other.RevokedAt != nil {
|
||||||
|
t.Fatal("expected other session to remain active")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func setupAuthServiceTestDB(t *testing.T) *gorm.DB {
|
||||||
|
t.Helper()
|
||||||
|
db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{
|
||||||
|
NamingStrategy: schema.NamingStrategy{
|
||||||
|
TablePrefix: "t_",
|
||||||
|
SingularTable: true,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("open sqlite db: %v", err)
|
||||||
|
}
|
||||||
|
if err := db.AutoMigrate(
|
||||||
|
&models.User{},
|
||||||
|
&models.Role{},
|
||||||
|
&models.Permission{},
|
||||||
|
&models.UserRole{},
|
||||||
|
&models.RolePermission{},
|
||||||
|
&models.UserPermission{},
|
||||||
|
&models.LoginSession{},
|
||||||
|
&models.LoginCredentialLog{},
|
||||||
|
); err != nil {
|
||||||
|
t.Fatalf("migrate auth tables: %v", err)
|
||||||
|
}
|
||||||
|
sqls.SetDB(db)
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
func createAuthTestUser(t *testing.T, db *gorm.DB, username, password string) *models.User {
|
||||||
|
t.Helper()
|
||||||
|
passwordHash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("hash password: %v", err)
|
||||||
|
}
|
||||||
|
now := time.Now()
|
||||||
|
user := &models.User{
|
||||||
|
Username: username,
|
||||||
|
Nickname: username,
|
||||||
|
Password: string(passwordHash),
|
||||||
|
Status: enums.StatusOk,
|
||||||
|
AuditFields: models.AuditFields{
|
||||||
|
CreatedAt: now,
|
||||||
|
UpdatedAt: now,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if err := db.Create(user).Error; err != nil {
|
||||||
|
t.Fatalf("create auth test user: %v", err)
|
||||||
|
}
|
||||||
|
return user
|
||||||
|
}
|
||||||
|
|
||||||
|
func findCredentialLogs(t *testing.T, db *gorm.DB) []models.LoginCredentialLog {
|
||||||
|
t.Helper()
|
||||||
|
var logs []models.LoginCredentialLog
|
||||||
|
if err := db.Order("id ASC").Find(&logs).Error; err != nil {
|
||||||
|
t.Fatalf("query credential logs: %v", err)
|
||||||
|
}
|
||||||
|
return logs
|
||||||
|
}
|
||||||
|
|
||||||
|
func hasCode(err error, code int) bool {
|
||||||
|
if err == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
var codeErr *web.CodeError
|
||||||
|
if errors.As(err, &codeErr) {
|
||||||
|
return codeErr.Code == code
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ import { usePathname, useRouter } from "next/navigation"
|
|||||||
|
|
||||||
import { fetchProfile, logout } from "@/lib/api/auth"
|
import { fetchProfile, logout } from "@/lib/api/auth"
|
||||||
import {
|
import {
|
||||||
|
AUTH_SESSION_EXPIRED_EVENT,
|
||||||
clearSession,
|
clearSession,
|
||||||
readSession,
|
readSession,
|
||||||
writeSession,
|
writeSession,
|
||||||
@@ -68,17 +69,35 @@ export function AuthProvider({ children }: { children: ReactNode }) {
|
|||||||
}, [requiresAuth, router])
|
}, [requiresAuth, router])
|
||||||
|
|
||||||
async function signOut() {
|
async function signOut() {
|
||||||
const current = readSession()
|
try {
|
||||||
await logout(current?.refreshToken)
|
await logout()
|
||||||
setSession(null)
|
} finally {
|
||||||
startTransition(() => {
|
setSession(null)
|
||||||
router.replace("/dashboard/login")
|
startTransition(() => {
|
||||||
})
|
router.replace("/dashboard/login")
|
||||||
|
})
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
useEffect(() => {
|
||||||
|
function handleAuthExpired() {
|
||||||
|
setSession(null)
|
||||||
|
if (requiresAuth) {
|
||||||
|
startTransition(() => {
|
||||||
|
router.replace("/dashboard/login")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
window.addEventListener(AUTH_SESSION_EXPIRED_EVENT, handleAuthExpired)
|
||||||
|
return () => {
|
||||||
|
window.removeEventListener(AUTH_SESSION_EXPIRED_EVENT, handleAuthExpired)
|
||||||
|
}
|
||||||
|
}, [requiresAuth, router])
|
||||||
|
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
const stored = readSession()
|
const stored = readSession()
|
||||||
setSession(stored)
|
setSession(stored)
|
||||||
if (stored) {
|
if (stored) {
|
||||||
void refreshProfile()
|
void refreshProfile()
|
||||||
return
|
return
|
||||||
|
|||||||
+1
-4
@@ -30,13 +30,10 @@ export async function fetchProfile() {
|
|||||||
return request<AuthSession>("/api/auth/profile")
|
return request<AuthSession>("/api/auth/profile")
|
||||||
}
|
}
|
||||||
|
|
||||||
export async function logout(refreshToken?: string) {
|
export async function logout() {
|
||||||
try {
|
try {
|
||||||
await request("/api/auth/logout", {
|
await request("/api/auth/logout", {
|
||||||
method: "POST",
|
method: "POST",
|
||||||
body: JSON.stringify({
|
|
||||||
refreshToken,
|
|
||||||
}),
|
|
||||||
})
|
})
|
||||||
} finally {
|
} finally {
|
||||||
clearSession()
|
clearSession()
|
||||||
|
|||||||
+6
-54
@@ -1,4 +1,4 @@
|
|||||||
import { clearSession, readSession, writeSession, type AuthSession } from "@/lib/auth"
|
import { expireSession, readSession } from "@/lib/auth"
|
||||||
|
|
||||||
const API_BASE_URL =
|
const API_BASE_URL =
|
||||||
process.env.NEXT_PUBLIC_API_BASE_URL?.trim() || ""
|
process.env.NEXT_PUBLIC_API_BASE_URL?.trim() || ""
|
||||||
@@ -12,7 +12,6 @@ type JsonResult<T> = {
|
|||||||
|
|
||||||
type RequestOptions = RequestInit & {
|
type RequestOptions = RequestInit & {
|
||||||
skipAuth?: boolean
|
skipAuth?: boolean
|
||||||
retryOnAuthError?: boolean
|
|
||||||
baseUrl?: string
|
baseUrl?: string
|
||||||
onResponse?: (response: Response) => void
|
onResponse?: (response: Response) => void
|
||||||
}
|
}
|
||||||
@@ -20,6 +19,9 @@ type RequestOptions = RequestInit & {
|
|||||||
async function parseResult<T>(response: Response) {
|
async function parseResult<T>(response: Response) {
|
||||||
const payload = (await response.json()) as JsonResult<T>
|
const payload = (await response.json()) as JsonResult<T>
|
||||||
if (!response.ok || !payload.success) {
|
if (!response.ok || !payload.success) {
|
||||||
|
if (payload.errorCode === 3000 || payload.errorCode === 3002) {
|
||||||
|
expireSession()
|
||||||
|
}
|
||||||
const error = new Error(payload.message || "请求失败")
|
const error = new Error(payload.message || "请求失败")
|
||||||
;(error as Error & { errorCode?: number }).errorCode = payload.errorCode
|
;(error as Error & { errorCode?: number }).errorCode = payload.errorCode
|
||||||
throw error
|
throw error
|
||||||
@@ -27,40 +29,11 @@ async function parseResult<T>(response: Response) {
|
|||||||
return payload.data
|
return payload.data
|
||||||
}
|
}
|
||||||
|
|
||||||
async function refreshAccessToken() {
|
|
||||||
const session = readSession()
|
|
||||||
if (!session?.refreshToken) {
|
|
||||||
clearSession()
|
|
||||||
return null
|
|
||||||
}
|
|
||||||
|
|
||||||
const data = await request<AuthSession>(
|
|
||||||
"/api/auth/refresh_token",
|
|
||||||
{
|
|
||||||
method: "POST",
|
|
||||||
body: JSON.stringify({ refreshToken: session.refreshToken }),
|
|
||||||
skipAuth: true,
|
|
||||||
headers: {
|
|
||||||
"Content-Type": "application/json",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
false
|
|
||||||
)
|
|
||||||
const merged = {
|
|
||||||
...data,
|
|
||||||
refreshToken: data.refreshToken || session.refreshToken,
|
|
||||||
}
|
|
||||||
writeSession(merged)
|
|
||||||
return merged
|
|
||||||
}
|
|
||||||
|
|
||||||
export async function request<T>(
|
export async function request<T>(
|
||||||
path: string,
|
path: string,
|
||||||
options: RequestOptions = {},
|
options: RequestOptions = {}
|
||||||
retryOnAuthError = true
|
|
||||||
): Promise<T> {
|
): Promise<T> {
|
||||||
const { headers, skipAuth, baseUrl, onResponse, ...rest } = options
|
const { headers, skipAuth, baseUrl, onResponse, ...rest } = options
|
||||||
delete (rest as RequestOptions).retryOnAuthError
|
|
||||||
delete (rest as RequestOptions).baseUrl
|
delete (rest as RequestOptions).baseUrl
|
||||||
delete (rest as RequestOptions).onResponse
|
delete (rest as RequestOptions).onResponse
|
||||||
const session = readSession()
|
const session = readSession()
|
||||||
@@ -85,26 +58,5 @@ export async function request<T>(
|
|||||||
})
|
})
|
||||||
onResponse?.(response)
|
onResponse?.(response)
|
||||||
|
|
||||||
try {
|
return parseResult<T>(response)
|
||||||
return await parseResult<T>(response)
|
|
||||||
} catch (error) {
|
|
||||||
const errorCode = (error as Error & { errorCode?: number }).errorCode
|
|
||||||
if (
|
|
||||||
!skipAuth &&
|
|
||||||
retryOnAuthError &&
|
|
||||||
(errorCode === 3000 || errorCode === 3002) &&
|
|
||||||
session?.refreshToken
|
|
||||||
) {
|
|
||||||
const refreshed = await refreshAccessToken()
|
|
||||||
if (!refreshed) {
|
|
||||||
throw error
|
|
||||||
}
|
|
||||||
return request<T>(path, options, false)
|
|
||||||
}
|
|
||||||
|
|
||||||
if (errorCode === 3000 || errorCode === 3002) {
|
|
||||||
clearSession()
|
|
||||||
}
|
|
||||||
throw error
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|||||||
+9
-1
@@ -9,7 +9,6 @@ export type AuthUser = {
|
|||||||
|
|
||||||
export type AuthSession = {
|
export type AuthSession = {
|
||||||
accessToken: string
|
accessToken: string
|
||||||
refreshToken: string
|
|
||||||
expiresAt?: string
|
expiresAt?: string
|
||||||
user: AuthUser
|
user: AuthUser
|
||||||
permissions: string[]
|
permissions: string[]
|
||||||
@@ -17,6 +16,7 @@ export type AuthSession = {
|
|||||||
}
|
}
|
||||||
|
|
||||||
const SESSION_STORAGE_KEY = "cs-ai-agent-session"
|
const SESSION_STORAGE_KEY = "cs-ai-agent-session"
|
||||||
|
export const AUTH_SESSION_EXPIRED_EVENT = "cs-ai-agent-auth-expired"
|
||||||
|
|
||||||
function hasWindow() {
|
function hasWindow() {
|
||||||
return typeof window !== "undefined"
|
return typeof window !== "undefined"
|
||||||
@@ -53,3 +53,11 @@ export function clearSession() {
|
|||||||
}
|
}
|
||||||
window.localStorage.removeItem(SESSION_STORAGE_KEY)
|
window.localStorage.removeItem(SESSION_STORAGE_KEY)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export function expireSession() {
|
||||||
|
if (!hasWindow()) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
clearSession()
|
||||||
|
window.dispatchEvent(new Event(AUTH_SESSION_EXPIRED_EVENT))
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user