mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-01 22:46:38 +08:00
refactor(auth): modularize auth plugin with physical subpackages and decoupled services
This commit is contained in:
@@ -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)),
|
||||
)
|
||||
}
|
||||
Reference in New Issue
Block a user