refactor(auth): modularize auth plugin with physical subpackages and decoupled services

This commit is contained in:
ryan
2026-09-03 09:12:44 +08:00
parent 2124bce7ca
commit 4407589b62
51 changed files with 3859 additions and 2915 deletions
@@ -0,0 +1,39 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package controller provides HTTP handlers and middlewares for the auth plugin.
package controller
import (
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/plugins/domain/auth/model/dto"
"context"
"encoding/json"
"github.com/gin-gonic/gin"
)
// LogForAudit 将登录鉴权审计日志写入 Logger
func LogForAudit(ctx context.Context, user *contracts.UserDTO, c *gin.Context) {
if user == nil || c == nil {
return
}
auditLog := dto.LoginRequiredAuditLog{
UserID: user.ID,
Username: user.Username,
ClientIP: c.ClientIP(),
Method: c.Request.Method,
Path: c.Request.URL.Path,
RequestURI: c.Request.RequestURI,
UserAgent: c.Request.UserAgent(),
Referer: c.Request.Referer(),
}
auditJSON, err := json.Marshal(auditLog)
if err != nil {
logger.ErrorF(ctx, "[LoginRequiredAudit] marshal failed: %v", err)
logger.DebugF(ctx, "[LoginRequiredAudit] %s %d %s", c.ClientIP(), user.ID, user.Username)
} else {
logger.DebugF(ctx, "[LoginRequiredAudit] %s", auditJSON)
}
}
@@ -0,0 +1,102 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package controller provides HTTP handlers and middlewares for the auth plugin.
package controller
import (
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/plugins/domain/auth/consts"
"Wavelet/plugins/domain/auth/model/dto"
"Wavelet/plugins/domain/auth/service"
"net/http"
"github.com/gin-gonic/gin"
)
// CaptchaHandler handles CAPTCHA challenge and redeem endpoints.
type CaptchaHandler struct {
capMgr *service.CaptchaManager
}
// NewCaptchaHandler creates a new CaptchaHandler.
func NewCaptchaHandler(mgr *service.CaptchaManager) *CaptchaHandler {
return &CaptchaHandler{
capMgr: mgr,
}
}
// Challenge 生成 PoW 人机验证难题
// @Summary 生成人机验证难题
// @Description 客户端获取 PoW 难题和签名的 JWT Token,并在后台计算。
// @Tags cap
// @Accept json
// @Produce json
// @Param request body dto.ChallengeRequest false "可选范围限制参数"
// @Success 200 {object} response.Any{data=dto.ChallengeResponse} "成功返回 PoW 难题"
// @Failure 500 {object} response.Any "内部服务错误"
// @Router /api/v1/cap/challenge [get]
// @Router /api/v1/cap/challenge [post]
func (h *CaptchaHandler) Challenge(c *gin.Context) {
var req dto.ChallengeRequest
_ = c.ShouldBind(&req) // 允许不传 body,默认使用 login scope
if req.Scope == "" {
req.Scope = "login"
}
if h.capMgr == nil {
response.AbortInternal(c, consts.ErrCapNotConfigured)
return
}
resp, err := h.capMgr.Generate(c.Request.Context(), req.Scope)
if err != nil {
logger.ErrorF(c.Request.Context(), "Generate cap challenge failed: %v", err)
response.AbortInternal(c, consts.ErrChallengeGenerateFailed)
return
}
c.JSON(http.StatusOK, response.OK(resp))
}
// Redeem 提交 PoW 解答并兑换一次性凭证 Token
// @Summary 校验人机验证解答
// @Description 提交 PoW 解答进行核销,成功后返回一次性 X-Cap-Token 凭证
// @Tags cap
// @Accept json
// @Produce json
// @Param request body dto.RedeemRequest true "难题 Token 与解答 solutions 数组"
// @Success 200 {object} response.Any{data=dto.RedeemResponse} "核销成功,返回 X-Cap-Token"
// @Failure 400 {object} response.Any "参数错误或核销失败"
// @Failure 500 {object} response.Any "内部服务错误"
// @Router /api/v1/cap/redeem [post]
func (h *CaptchaHandler) Redeem(c *gin.Context) {
var req dto.RedeemRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, consts.ErrInvalidRequestParams)
return
}
if req.Scope == "" {
req.Scope = "login"
}
if h.capMgr == nil {
response.AbortInternal(c, consts.ErrCapNotConfigured)
return
}
resp, err := h.capMgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope)
if err != nil {
logger.ErrorF(c.Request.Context(), "Redeem cap solutions failed: %v", err)
response.AbortInternal(c, consts.ErrSolutionVerifyFailed)
return
}
if !resp.Success {
response.AbortBadRequest(c, resp.Error)
return
}
c.JSON(http.StatusOK, response.OK(resp))
}
@@ -0,0 +1,41 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package controller provides HTTP handlers and middlewares for the auth plugin.
package controller
import (
"Wavelet/pkg/response"
"Wavelet/plugins/domain/auth/consts"
"Wavelet/plugins/domain/auth/service"
"github.com/gin-gonic/gin"
)
// VerifyCaptchaMiddleware returns a Gin middleware that checks and consumes the X-Cap-Token header.
func VerifyCaptchaMiddleware(mgr *service.CaptchaManager, settingsMgr *service.CapSettingsManager, scope string) gin.HandlerFunc {
return func(c *gin.Context) {
if settingsMgr != nil && !settingsMgr.CapProtectionEnabled(c.Request.Context()) {
c.Next()
return
}
if mgr == nil {
response.AbortBadRequest(c, consts.ErrCapTokenInvalidOrExpired)
return
}
token := c.GetHeader("X-Cap-Token")
if token == "" {
response.AbortBadRequest(c, consts.ErrCapTokenMissing)
return
}
valid, err := mgr.VerifyToken(c.Request.Context(), token, scope)
if err != nil || !valid {
response.AbortBadRequest(c, consts.ErrCapTokenInvalidOrExpired)
return
}
c.Next()
}
}
@@ -0,0 +1,109 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package controller provides HTTP handlers and middlewares for the auth plugin.
package controller
import (
"Wavelet/core/extpoints"
"Wavelet/plugins/domain/auth/service"
"github.com/gin-gonic/gin"
)
// Controller aggregates all HTTP handlers and middlewares for the auth plugin.
type Controller struct {
svc *service.Service
whitelist *extpoints.PathWhitelist
OAuth *OAuthHandler
UserInfo *UserInfoHandler
Captcha *CaptchaHandler
}
// New creates a new Controller instance.
func New(svc *service.Service) *Controller {
wl := extpoints.NewPathWhitelist()
oauthHandler := NewOAuthHandler(svc.OAuth, svc.Session, svc.DAO)
userInfoHandler := NewUserInfoHandler()
captchaHandler := NewCaptchaHandler(svc.CapManager)
c := &Controller{
svc: svc,
whitelist: wl,
OAuth: oauthHandler,
UserInfo: userInfoHandler,
Captcha: captchaHandler,
}
// Wire middlewares into AuthService
svc.AuthSvc.SetMiddlewareHandlers(
c.LoginRequired(),
c.AdminRequired(),
DisallowTokenAuth(),
CurrentUserIDFromRequestContext,
)
return c
}
// Whitelist returns the whitelist tracker.
func (c *Controller) Whitelist() *extpoints.PathWhitelist {
return c.whitelist
}
// RegisterWhitelist adds path patterns that bypass authentication.
func (c *Controller) RegisterWhitelist(patterns ...string) {
if c.whitelist != nil {
c.whitelist.Add(patterns...)
}
}
// LoginRequired returns the authentication middleware.
func (c *Controller) LoginRequired() gin.HandlerFunc {
return LoginRequiredMiddleware(c.whitelist, c.svc.DAO)
}
// AdminRequired returns the admin authorization middleware.
func (c *Controller) AdminRequired() gin.HandlerFunc {
return AdminRequiredMiddleware(c.svc.DAO)
}
// DisallowTokenAuth returns the token rejection middleware.
func (c *Controller) DisallowTokenAuth() gin.HandlerFunc {
return DisallowTokenAuth()
}
// VerifyCaptcha returns the captcha challenge verification middleware.
func (c *Controller) VerifyCaptcha(scope string) gin.HandlerFunc {
return VerifyCaptchaMiddleware(c.svc.CapManager, c.svc.CapSettings, scope)
}
// RegisterRoutes mounts all auth endpoints onto the router.
func (c *Controller) RegisterRoutes(router extpoints.RouterExtension) {
loginReq := c.LoginRequired()
// 1. OAuth endpoints
oauthGroup := router.Group("/api/v1/oauth")
{
oauthGroup.GET("/sources", c.OAuth.GetLoginSources)
oauthGroup.GET("/login", c.OAuth.GetLoginURL)
oauthGroup.GET("/:source/authorize", c.OAuth.Authorize)
oauthGroup.GET("/logout", c.OAuth.Logout)
oauthGroup.POST("/callback", c.OAuth.Callback)
oauthGroup.GET("/user-info", loginReq, c.UserInfo.UserInfo)
oauthGroup.GET("/external-accounts", loginReq, c.OAuth.ListExternalAccounts)
oauthGroup.POST("/external-accounts/:id/delete", loginReq, c.OAuth.DeleteExternalAccount)
}
// 2. Global user-info route alias
router.GET("/api/v1/user-info", loginReq, c.UserInfo.UserInfo)
// 3. CAPTCHA endpoints
capGroup := router.Group("/api/v1/cap")
{
capGroup.GET("/challenge", c.Captcha.Challenge)
capGroup.POST("/challenge", c.Captcha.Challenge)
capGroup.POST("/redeem", c.Captcha.Redeem)
}
}
@@ -0,0 +1,187 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package controller provides HTTP handlers and middlewares for the auth plugin.
package controller
import (
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/response"
"Wavelet/pkg/trace"
"Wavelet/plugins/domain/auth/consts"
"Wavelet/plugins/domain/auth/dao"
"Wavelet/plugins/domain/auth/model/do"
"Wavelet/plugins/domain/auth/model/dto"
"Wavelet/plugins/domain/auth/service"
"context"
"errors"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
)
// GetUserIDFromSession 从 Session 中提取用户 ID
func GetUserIDFromSession(s sessions.Session) uint64 {
val := s.Get(consts.UserIDKey)
return dto.ParseUserID(val)
}
// GetUserIDFromContext 从 Gin Context 的 Session 中提取用户 ID
func GetUserIDFromContext(c *gin.Context) (uid uint64) {
defer func() {
_ = recover()
}()
session := sessions.Default(c)
return GetUserIDFromSession(session)
}
// CurrentUserIDFromRequestContext 是接入层向 Service 层暴露的登录态桥接。
func CurrentUserIDFromRequestContext(ctx context.Context) (uint64, bool) {
ginCtx, ok := ctx.(*gin.Context)
if !ok {
return 0, false
}
return GetUserIDFromContext(ginCtx), true
}
func getUserByToken(ctx context.Context, d *dao.DAO, tokenStr string) (*contracts.UserDTO, *do.CachedToken, error) {
tokenHash := service.HashToken(tokenStr)
tokenRecord, err := d.GetCachedToken(ctx, tokenHash)
if err != nil || tokenRecord == nil {
tokenRecord, err = d.GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, nil, err
}
d.SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := d.GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || user == nil || !user.IsActive {
user, err = d.GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, nil, err
}
d.SetCachedUser(ctx, tokenRecord.UserID, user)
}
return user, tokenRecord, nil
}
// GetUserFromRequest 从请求中获取当前用户(优先 Access Token,其次 Session)
func GetUserFromRequest(c *gin.Context, d *dao.DAO) (*contracts.UserDTO, error) {
ctx := c.Request.Context()
var tokenStr string
tokenFromQuery := c.Query("token")
if tokenFromQuery != "" {
tokenStr = tokenFromQuery
} else {
authHeader := c.GetHeader("Authorization")
if len(authHeader) > 7 && authHeader[:7] == "Bearer " {
tokenStr = authHeader[7:]
}
}
// 优先使用 Access Token 鉴权
if tokenStr != "" {
if user, tokenRecord, err := getUserByToken(ctx, d, tokenStr); err == nil {
if user.Username == consts.SystemUsername {
return nil, errors.New(consts.ErrSystemUserLoginNotAllowed)
}
ginutil.SetToContext(c, contracts.AuthTokenAuthKey, true)
ginutil.SetToContext(c, contracts.AuthTokenAdminKey, tokenRecord.IsAdmin)
return user, nil
}
}
// 降级使用 Session 鉴权
userID := GetUserIDFromContext(c)
if userID <= 0 {
return nil, errors.New(consts.ErrUnauthorizedInternal)
}
user, err := d.GetCachedUser(ctx, userID)
if err != nil || user == nil || !user.IsActive {
user, err = d.GetActiveUserByID(ctx, userID)
if err != nil {
return nil, err
}
d.SetCachedUser(ctx, userID, user)
}
ginutil.SetToContext(c, contracts.AuthTokenAuthKey, false)
ginutil.SetToContext(c, contracts.AuthTokenAdminKey, false)
if user.Username == consts.SystemUsername {
return nil, errors.New(consts.ErrSystemUserLoginNotAllowed)
}
return user, nil
}
// LoginRequiredMiddleware returns a Gin handler function for authentication check.
func LoginRequiredMiddleware(whitelist *extpoints.PathWhitelist, d *dao.DAO) gin.HandlerFunc {
return func(c *gin.Context) {
if whitelist != nil && whitelist.Match(c.Request.URL.Path) {
c.Next()
return
}
_, span := trace.Start(c.Request.Context(), "LoginRequired")
defer span.End()
user, err := GetUserFromRequest(c, d)
if err != nil {
response.AbortUnauthorized(c, consts.ErrUnAuthorized)
return
}
LogForAudit(c.Request.Context(), user, c)
ginutil.SetToContext(c, contracts.AuthUserObjKey, user)
c.Next()
}
}
// AdminRequiredMiddleware returns a Gin handler function for admin authorization check.
func AdminRequiredMiddleware(d *dao.DAO) gin.HandlerFunc {
return func(c *gin.Context) {
_, span := trace.Start(c.Request.Context(), "AdminRequired")
defer span.End()
user, err := GetUserFromRequest(c, d)
if err != nil {
response.AbortUnauthorized(c, consts.ErrUnAuthorized)
return
}
isTokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey)
isTokenAdmin, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
// Logged-in but lacking admin permission is 403, not 401/404.
if isTokenAuth && !isTokenAdmin && !user.IsAdmin {
response.AbortForbidden(c, consts.ErrInsufficientPermission)
return
}
if !isTokenAuth && !user.IsAdmin {
response.AbortForbidden(c, consts.ErrInsufficientPermission)
return
}
LogForAudit(c.Request.Context(), user, c)
ginutil.SetToContext(c, contracts.AuthUserObjKey, user)
c.Next()
}
}
// DisallowTokenAuth returns a middleware that rejects requests authenticated via access token.
func DisallowTokenAuth() gin.HandlerFunc {
return func(c *gin.Context) {
if tokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
response.AbortForbidden(c, consts.ErrTokenAuthNotAllowed)
return
}
c.Next()
}
}
@@ -0,0 +1,448 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package controller provides HTTP handlers and middlewares for the auth plugin.
package controller
import (
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/plugins/domain/auth/consts"
"Wavelet/plugins/domain/auth/dao"
"Wavelet/plugins/domain/auth/model/do"
"Wavelet/plugins/domain/auth/model/dto"
"Wavelet/plugins/domain/auth/model/entity"
"Wavelet/plugins/domain/auth/service"
"context"
"fmt"
"net/http"
"strconv"
"strings"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
)
// OAuthHandler handles OAuth authentication endpoints.
type OAuthHandler struct {
oauthSvc *service.OAuthService
sessionSvc *service.SessionService
dao *dao.DAO
}
// NewOAuthHandler creates a new OAuthHandler.
func NewOAuthHandler(oauthSvc *service.OAuthService, sessionSvc *service.SessionService, d *dao.DAO) *OAuthHandler {
return &OAuthHandler{
oauthSvc: oauthSvc,
sessionSvc: sessionSvc,
dao: d,
}
}
// GetLoginSources 获取可用登录源列表
// @Summary 获取可用登录源
// @Description 返回当前系统已启用的所有 OAuth 登录源,前端展示登录按钮列表时调用
// @Tags oauth
// @Produce json
// @Success 200 {object} response.Any{data=[]dto.AuthSourceView} "登录源列表"
// @Router /api/v1/oauth/sources [get]
func (h *OAuthHandler) GetLoginSources(c *gin.Context) {
sources, err := h.oauthSvc.ActiveLoginSources(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(sources))
}
// GetLoginURL 获取登录授权地址
// @Summary 获取登录授权地址
// @Description 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用第一个启用的认证源。
// @Tags oauth
// @Produce json
// @Param source query string false "认证源名称,为空使用第一个启用的认证源"
// @Success 200 {object} response.Any{data=dto.OAuthAuthorizeResponse} "授权 URL"
// @Failure 400 {object} response.Any "认证源不存在或未配置"
// @Failure 500 {object} response.Any "构造 URL 失败"
// @Router /api/v1/oauth/login [get]
func (h *OAuthHandler) GetLoginURL(c *gin.Context) {
ctx := c.Request.Context()
if !h.oauthSvc.IsOIDCLoginEnabled(ctx) {
response.AbortBadRequest(c, consts.ErrAuthSourceDisabled)
return
}
source, err := h.oauthSvc.ResolveAuthSource(ctx, c.Query("source"))
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.AbortBadRequest(c, consts.ErrAuthSourceDisabled)
return
}
session := sessions.Default(c)
token, isNew := h.sessionSvc.EnsureSessionToken(session)
if isNew {
if err := session.Save(); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
userID := GetUserIDFromSession(session)
sessionHash := h.sessionSvc.HashSessionToken(token)
if err := h.oauthSvc.ReserveOAuthStateSlot(ctx, sessionHash); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
state := uuid.NewString()
payloadValue, err := (do.OAuthStatePayload{
SourceName: source.Name,
Purpose: consts.OAuthPurposeLogin,
UserID: userID,
SessionHash: sessionHash,
}).Encode()
if err != nil {
response.AbortInternal(c, err.Error())
return
}
stateKey := fmt.Sprintf(consts.OAuthStateCacheKeyFormat, state)
if cache := h.dao.Cache(); cache != nil {
if err := cache.Set(ctx, stateKey, payloadValue, consts.OAuthStateCacheKeyExpiration); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
authorizeURL, err := h.oauthSvc.BuildAuthorizeURL(ctx, source, state)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(dto.OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
}
// Authorize 发起指定认证源授权
// @Summary 发起指定认证源授权
// @Description 根据指定认证源名称发起 OAuth 授权,支持 purpose 参数用于区分登录和账号绑定场景。认证源必须已启用。
// @Tags oauth
// @Produce json
// @Param source path string true "认证源名称"
// @Param purpose query string false "授权目的:login(登录)或 bind(绑定账号),默认 login"
// @Success 200 {object} response.Any{data=dto.OAuthAuthorizeResponse} "授权 URL"
// @Failure 400 {object} response.Any "认证源不存在或未启用"
// @Failure 500 {object} response.Any "构造 URL 失败"
// @Router /api/v1/oauth/{source}/authorize [get]
func (h *OAuthHandler) Authorize(c *gin.Context) {
ctx := c.Request.Context()
if !h.oauthSvc.IsOIDCLoginEnabled(ctx) {
response.AbortBadRequest(c, consts.ErrAuthSourceDisabled)
return
}
source, err := h.oauthSvc.ResolveAuthSource(ctx, c.Param("source"))
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.AbortBadRequest(c, consts.ErrAuthSourceDisabled)
return
}
purpose := strings.ToLower(strings.TrimSpace(c.Query("purpose")))
if purpose != consts.OAuthPurposeBind {
purpose = consts.OAuthPurposeLogin
}
session := sessions.Default(c)
userID := GetUserIDFromSession(session)
if purpose == consts.OAuthPurposeBind && userID == 0 {
response.AbortUnauthorized(c, consts.ErrUnAuthorized)
return
}
token, isNew := h.sessionSvc.EnsureSessionToken(session)
if isNew {
if err := session.Save(); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
sessionHash := h.sessionSvc.HashSessionToken(token)
if err := h.oauthSvc.ReserveOAuthStateSlot(ctx, sessionHash); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
state := uuid.NewString()
payloadValue, err := (do.OAuthStatePayload{
SourceName: source.Name,
Purpose: purpose,
UserID: userID,
SessionHash: sessionHash,
}).Encode()
if err != nil {
response.AbortInternal(c, err.Error())
return
}
stateKey := fmt.Sprintf(consts.OAuthStateCacheKeyFormat, state)
if cache := h.dao.Cache(); cache != nil {
if err := cache.Set(ctx, stateKey, payloadValue, consts.OAuthStateCacheKeyExpiration); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
authorizeURL, err := h.oauthSvc.BuildAuthorizeURL(ctx, source, state)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(dto.OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
}
// Callback OAuth 回调处理
// @Summary OAuth 回调处理
// @Description 接收前端传回的 state 和 code,完成 OAuth/OIDC 认证并建立会话。支持登录(login)和账号绑定(bind)两种场景。
// @Tags oauth
// @Accept json
// @Produce json
// @Param request body dto.CallbackRequest true "回调请求参数"
// @Success 200 {object} response.Any{data=dto.OAuthCallbackResult} "登录或绑定成功"
// @Failure 400 {object} response.Any "state 无效、参数错误或认证源错误"
// @Failure 401 {object} response.Any "绑定场景未登录"
// @Failure 500 {object} response.Any "OAuth 认证失败或内部错误"
// @Router /api/v1/oauth/callback [post]
func (h *OAuthHandler) Callback(c *gin.Context) {
var req dto.CallbackRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
ctx := c.Request.Context()
stateKey := fmt.Sprintf(consts.OAuthStateCacheKeyFormat, req.State)
var payloadRaw string
cache := h.dao.Cache()
if cache == nil {
response.AbortBadRequest(c, consts.ErrInvalidState)
return
}
if err := cache.Get(ctx, stateKey, &payloadRaw); err != nil {
response.AbortBadRequest(c, consts.ErrInvalidState)
return
}
_ = cache.Delete(ctx, stateKey)
payload, err := do.DecodeOAuthStatePayload(payloadRaw)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
session := sessions.Default(c)
currentUserID := GetUserIDFromSession(session)
if payload.Purpose == consts.OAuthPurposeBind && currentUserID == 0 {
response.AbortUnauthorized(c, consts.ErrUnAuthorized)
return
}
token, ok := session.Get(consts.SessionTokenKey).(string)
if !ok || token == "" {
response.AbortBadRequest(c, consts.ErrInvalidSessionContext)
return
}
if h.sessionSvc.HashSessionToken(token) != payload.SessionHash {
response.AbortBadRequest(c, consts.ErrSessionMismatchForOAuth)
return
}
if payload.Purpose == consts.OAuthPurposeBind && currentUserID != payload.UserID {
response.AbortBadRequest(c, consts.ErrUserContextMismatch)
return
}
if !h.oauthSvc.IsOIDCLoginEnabled(ctx) {
response.AbortBadRequest(c, consts.ErrAuthSourceDisabled)
return
}
source, err := h.oauthSvc.ResolveAuthSource(ctx, payload.SourceName)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.AbortBadRequest(c, consts.ErrAuthSourceDisabled)
return
}
redirectURL, err := h.oauthSvc.GetFrontendLoginRedirectURL(ctx)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
userInfo, err := h.oauthSvc.BuildOAuthUserInfo(ctx, source, req.Code, req.State, redirectURL)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := h.oauthSvc.NormalizeOAuthUserInfo(userInfo); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if userInfo.Sub == "" {
userInfo.Sub = userInfo.Username
}
if payload.Purpose == consts.OAuthPurposeBind {
h.handleCallbackBind(ctx, c, source, userInfo)
return
}
h.handleCallbackLogin(ctx, c, source, userInfo)
}
func (h *OAuthHandler) handleCallbackBind(ctx context.Context, c *gin.Context, source *entity.AuthSource, userInfo *contracts.OAuthUserInfoDTO) {
userID := GetUserIDFromContext(c)
if userID == 0 {
response.AbortUnauthorized(c, consts.ErrUnAuthorized)
return
}
user, err := h.dao.GetUserByID(ctx, userID)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := h.oauthSvc.BindExternalAccount(ctx, source.ID, user.ID, userInfo); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(buildCallbackResult(user, "bound")))
}
func (h *OAuthHandler) handleCallbackLogin(ctx context.Context, c *gin.Context, source *entity.AuthSource, userInfo *contracts.OAuthUserInfoDTO) {
user, ok, err := h.oauthSvc.AuthenticateOrRegisterUser(ctx, source, userInfo)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
if !ok || user == nil {
c.JSON(http.StatusOK, response.OK(buildCallbackResult(nil, "need_bind")))
return
}
session := sessions.Default(c)
isSessionCookie, err := h.sessionSvc.ApplyLoginSession(ctx, session, user)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
if isSessionCookie {
h.sessionSvc.StripCookieMaxAgeAndExpires(c.Writer.Header(), h.sessionSvc.Config().SessionCookieName)
}
h.dao.SetCachedUser(ctx, user.ID, user)
logger.InfoF(ctx, "[LoginAudit] successful OAuth login via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
c.JSON(http.StatusOK, response.OK(buildCallbackResult(user, "logged_in")))
}
func buildCallbackResult(user *contracts.UserDTO, status string) dto.OAuthCallbackResult {
result := dto.OAuthCallbackResult{Status: status}
if user != nil {
info := dto.BuildBasicUserInfo(user, false)
result.User = &info
}
return result
}
// Logout 退出登录
// @Summary 退出登录
// @Description 清除当前用户的登录会话,完成退出。清除 Cookie 中的 Session 数据。
// @Tags oauth
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=string} "退出成功"
// @Failure 500 {object} response.Any "Session 清除失败"
// @Router /api/v1/oauth/logout [get]
func (h *OAuthHandler) Logout(c *gin.Context) {
session := sessions.Default(c)
userID := session.Get(consts.UserIDKey)
username := session.Get(consts.UserNameKey)
if userID != nil {
logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP())
if id := dto.ParseUserID(userID); id > 0 {
h.dao.InvalidateCachedUser(c.Request.Context(), id)
}
}
session.Options(h.sessionSvc.GetSessionOptions(-1))
session.Clear()
if err := session.Save(); err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// ListExternalAccounts 获取当前用户的外部帐号绑定列表
// @Summary 获取外部帐号列表
// @Description 返回当前登录用户已绑定的所有外部 OAuth 帐号信息,需要登录
// @Tags oauth
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any "外部帐号列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/oauth/external-accounts [get]
func (h *OAuthHandler) ListExternalAccounts(c *gin.Context) {
userID := GetUserIDFromContext(c)
accounts, err := h.oauthSvc.ListExternalAccounts(c.Request.Context(), userID)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(accounts))
}
// DeleteExternalAccount 解除外部帐号绑定
// @Summary 解除外部帐号绑定
// @Description 解除当前登录用户与指定外部帐号的绑定关系,需要登录
// @Tags oauth
// @Produce json
// @Security SessionCookie
// @Param id path uint64 true "外部帐号绑定记录 ID"
// @Success 200 {object} response.Any{data=string} "解除绑定成功"
// @Failure 400 {object} response.Any "ID 无效或解除失败"
// @Failure 401 {object} response.Any "未登录"
// @Router /api/v1/oauth/external-accounts/{id}/delete [post]
func (h *OAuthHandler) DeleteExternalAccount(c *gin.Context) {
userID := GetUserIDFromContext(c)
if userID == 0 {
response.AbortUnauthorized(c, consts.ErrUnAuthorized)
return
}
rawID := strings.TrimSpace(c.Param("id"))
id, err := strconv.ParseUint(rawID, 10, 64)
if err != nil || id == 0 {
response.AbortBadRequest(c, consts.ErrInvalidExternalAccountBindingID)
return
}
if err := h.oauthSvc.DeleteExternalAccount(c.Request.Context(), id, userID); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,45 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package controller provides HTTP handlers and middlewares for the auth plugin.
package controller
import (
"Wavelet/core/contracts"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/response"
"Wavelet/plugins/domain/auth/model/dto"
"net/http"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
)
// UserInfoHandler handles current user info queries.
type UserInfoHandler struct{}
// NewUserInfoHandler creates a new UserInfoHandler.
func NewUserInfoHandler() *UserInfoHandler {
return &UserInfoHandler{}
}
// UserInfo 获取当前登录用户信息
// @Summary 获取当前登录用户信息
// @Description 返回当前登录用户的基本信息,需要登录。
// @Tags oauth
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=dto.BasicUserInfo} "用户信息"
// @Failure 401 {object} response.Any "未登录"
// @Router /api/v1/oauth/user-info [get]
// @Router /api/v1/user-info [get]
func (h *UserInfoHandler) UserInfo(c *gin.Context) {
user, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
session := sessions.Default(c)
needChange := session.Get("need_change_password") == true || (user != nil && user.NeedChangePassword)
c.JSON(
http.StatusOK,
response.OK(dto.BuildBasicUserInfo(user, needChange)),
)
}