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
+187 -187
View File
@@ -4397,7 +4397,7 @@ const docTemplate = `{
"name": "request",
"in": "body",
"schema": {
"$ref": "#/definitions/auth.challengeRequest"
"$ref": "#/definitions/dto.ChallengeRequest"
}
}
],
@@ -4413,7 +4413,7 @@ const docTemplate = `{
"type": "object",
"properties": {
"data": {
"$ref": "#/definitions/auth.ChallengeResponse"
"$ref": "#/definitions/dto.ChallengeResponse"
}
}
}
@@ -4446,7 +4446,7 @@ const docTemplate = `{
"name": "request",
"in": "body",
"schema": {
"$ref": "#/definitions/auth.challengeRequest"
"$ref": "#/definitions/dto.ChallengeRequest"
}
}
],
@@ -4462,7 +4462,7 @@ const docTemplate = `{
"type": "object",
"properties": {
"data": {
"$ref": "#/definitions/auth.ChallengeResponse"
"$ref": "#/definitions/dto.ChallengeResponse"
}
}
}
@@ -4498,7 +4498,7 @@ const docTemplate = `{
"in": "body",
"required": true,
"schema": {
"$ref": "#/definitions/auth.redeemRequest"
"$ref": "#/definitions/dto.RedeemRequest"
}
}
],
@@ -4514,7 +4514,7 @@ const docTemplate = `{
"type": "object",
"properties": {
"data": {
"$ref": "#/definitions/auth.RedeemResponse"
"$ref": "#/definitions/dto.RedeemResponse"
}
}
}
@@ -4778,7 +4778,7 @@ const docTemplate = `{
"in": "body",
"required": true,
"schema": {
"$ref": "#/definitions/auth.CallbackRequest"
"$ref": "#/definitions/dto.CallbackRequest"
}
}
],
@@ -4794,7 +4794,7 @@ const docTemplate = `{
"type": "object",
"properties": {
"data": {
"$ref": "#/definitions/auth.OAuthCallbackResult"
"$ref": "#/definitions/dto.OAuthCallbackResult"
}
}
}
@@ -4948,7 +4948,7 @@ const docTemplate = `{
"type": "object",
"properties": {
"data": {
"$ref": "#/definitions/auth.OAuthAuthorizeResponse"
"$ref": "#/definitions/dto.OAuthAuthorizeResponse"
}
}
}
@@ -5037,7 +5037,7 @@ const docTemplate = `{
"data": {
"type": "array",
"items": {
"$ref": "#/definitions/auth.AuthSourceView"
"$ref": "#/definitions/dto.AuthSourceView"
}
}
}
@@ -5075,7 +5075,7 @@ const docTemplate = `{
"type": "object",
"properties": {
"data": {
"$ref": "#/definitions/auth.BasicUserInfo"
"$ref": "#/definitions/dto.BasicUserInfo"
}
}
}
@@ -5128,7 +5128,7 @@ const docTemplate = `{
"type": "object",
"properties": {
"data": {
"$ref": "#/definitions/auth.OAuthAuthorizeResponse"
"$ref": "#/definitions/dto.OAuthAuthorizeResponse"
}
}
}
@@ -5445,7 +5445,7 @@ const docTemplate = `{
"type": "object",
"properties": {
"data": {
"$ref": "#/definitions/auth.BasicUserInfo"
"$ref": "#/definitions/dto.BasicUserInfo"
}
}
}
@@ -6051,180 +6051,6 @@ const docTemplate = `{
}
},
"definitions": {
"auth.AuthSourceView": {
"type": "object",
"properties": {
"client_secret_configured": {
"type": "boolean"
},
"display_name": {
"type": "string"
},
"icon_url": {
"type": "string"
},
"id": {
"type": "integer"
},
"is_active": {
"type": "boolean"
},
"name": {
"type": "string"
},
"type": {
"type": "string"
}
}
},
"auth.BasicUserInfo": {
"type": "object",
"properties": {
"avatar_url": {
"type": "string"
},
"bio": {
"type": "string"
},
"email": {
"type": "string"
},
"gender": {
"type": "string"
},
"id": {
"type": "string",
"example": "0"
},
"is_admin": {
"type": "boolean"
},
"location": {
"type": "string"
},
"need_change_password": {
"type": "boolean"
},
"nickname": {
"type": "string"
},
"phone": {
"type": "string"
},
"username": {
"type": "string"
},
"website": {
"type": "string"
}
}
},
"auth.CallbackRequest": {
"type": "object",
"required": [
"code",
"state"
],
"properties": {
"code": {
"type": "string"
},
"state": {
"type": "string"
}
}
},
"auth.ChallengeResponse": {
"type": "object",
"properties": {
"challenge": {
"type": "object",
"properties": {
"c": {
"type": "integer"
},
"d": {
"type": "integer"
},
"s": {
"type": "integer"
}
}
},
"expires": {
"description": "ms timestamp",
"type": "integer"
},
"token": {
"type": "string"
}
}
},
"auth.OAuthAuthorizeResponse": {
"type": "object",
"properties": {
"authorize_url": {
"type": "string"
}
}
},
"auth.OAuthCallbackResult": {
"type": "object",
"properties": {
"status": {
"type": "string"
},
"user": {
"$ref": "#/definitions/auth.BasicUserInfo"
}
}
},
"auth.RedeemResponse": {
"type": "object",
"properties": {
"error": {
"type": "string"
},
"expires": {
"type": "integer"
},
"success": {
"type": "boolean"
},
"token": {
"type": "string"
}
}
},
"auth.challengeRequest": {
"type": "object",
"properties": {
"scope": {
"type": "string"
}
}
},
"auth.redeemRequest": {
"type": "object",
"required": [
"solutions",
"token"
],
"properties": {
"scope": {
"type": "string"
},
"solutions": {
"type": "array",
"items": {
"type": "integer"
}
},
"token": {
"type": "string"
}
}
},
"contracts.AuthSourceDTO": {
"type": "object",
"properties": {
@@ -6678,6 +6504,180 @@ const docTemplate = `{
}
}
},
"dto.AuthSourceView": {
"type": "object",
"properties": {
"client_secret_configured": {
"type": "boolean"
},
"display_name": {
"type": "string"
},
"icon_url": {
"type": "string"
},
"id": {
"type": "integer"
},
"is_active": {
"type": "boolean"
},
"name": {
"type": "string"
},
"type": {
"type": "string"
}
}
},
"dto.BasicUserInfo": {
"type": "object",
"properties": {
"avatar_url": {
"type": "string"
},
"bio": {
"type": "string"
},
"email": {
"type": "string"
},
"gender": {
"type": "string"
},
"id": {
"type": "string",
"example": "0"
},
"is_admin": {
"type": "boolean"
},
"location": {
"type": "string"
},
"need_change_password": {
"type": "boolean"
},
"nickname": {
"type": "string"
},
"phone": {
"type": "string"
},
"username": {
"type": "string"
},
"website": {
"type": "string"
}
}
},
"dto.CallbackRequest": {
"type": "object",
"required": [
"code",
"state"
],
"properties": {
"code": {
"type": "string"
},
"state": {
"type": "string"
}
}
},
"dto.ChallengeRequest": {
"type": "object",
"properties": {
"scope": {
"type": "string"
}
}
},
"dto.ChallengeResponse": {
"type": "object",
"properties": {
"challenge": {
"type": "object",
"properties": {
"c": {
"type": "integer"
},
"d": {
"type": "integer"
},
"s": {
"type": "integer"
}
}
},
"expires": {
"description": "ms timestamp",
"type": "integer"
},
"token": {
"type": "string"
}
}
},
"dto.OAuthAuthorizeResponse": {
"type": "object",
"properties": {
"authorize_url": {
"type": "string"
}
}
},
"dto.OAuthCallbackResult": {
"type": "object",
"properties": {
"status": {
"type": "string"
},
"user": {
"$ref": "#/definitions/dto.BasicUserInfo"
}
}
},
"dto.RedeemRequest": {
"type": "object",
"required": [
"solutions",
"token"
],
"properties": {
"scope": {
"type": "string"
},
"solutions": {
"type": "array",
"items": {
"type": "integer"
}
},
"token": {
"type": "string"
}
}
},
"dto.RedeemResponse": {
"type": "object",
"properties": {
"error": {
"type": "string"
},
"expires": {
"type": "integer"
},
"success": {
"type": "boolean"
},
"token": {
"type": "string"
}
}
},
"entity.PushChannel": {
"type": "object",
"properties": {
+187 -187
View File
@@ -4390,7 +4390,7 @@
"name": "request",
"in": "body",
"schema": {
"$ref": "#/definitions/auth.challengeRequest"
"$ref": "#/definitions/dto.ChallengeRequest"
}
}
],
@@ -4406,7 +4406,7 @@
"type": "object",
"properties": {
"data": {
"$ref": "#/definitions/auth.ChallengeResponse"
"$ref": "#/definitions/dto.ChallengeResponse"
}
}
}
@@ -4439,7 +4439,7 @@
"name": "request",
"in": "body",
"schema": {
"$ref": "#/definitions/auth.challengeRequest"
"$ref": "#/definitions/dto.ChallengeRequest"
}
}
],
@@ -4455,7 +4455,7 @@
"type": "object",
"properties": {
"data": {
"$ref": "#/definitions/auth.ChallengeResponse"
"$ref": "#/definitions/dto.ChallengeResponse"
}
}
}
@@ -4491,7 +4491,7 @@
"in": "body",
"required": true,
"schema": {
"$ref": "#/definitions/auth.redeemRequest"
"$ref": "#/definitions/dto.RedeemRequest"
}
}
],
@@ -4507,7 +4507,7 @@
"type": "object",
"properties": {
"data": {
"$ref": "#/definitions/auth.RedeemResponse"
"$ref": "#/definitions/dto.RedeemResponse"
}
}
}
@@ -4771,7 +4771,7 @@
"in": "body",
"required": true,
"schema": {
"$ref": "#/definitions/auth.CallbackRequest"
"$ref": "#/definitions/dto.CallbackRequest"
}
}
],
@@ -4787,7 +4787,7 @@
"type": "object",
"properties": {
"data": {
"$ref": "#/definitions/auth.OAuthCallbackResult"
"$ref": "#/definitions/dto.OAuthCallbackResult"
}
}
}
@@ -4941,7 +4941,7 @@
"type": "object",
"properties": {
"data": {
"$ref": "#/definitions/auth.OAuthAuthorizeResponse"
"$ref": "#/definitions/dto.OAuthAuthorizeResponse"
}
}
}
@@ -5030,7 +5030,7 @@
"data": {
"type": "array",
"items": {
"$ref": "#/definitions/auth.AuthSourceView"
"$ref": "#/definitions/dto.AuthSourceView"
}
}
}
@@ -5068,7 +5068,7 @@
"type": "object",
"properties": {
"data": {
"$ref": "#/definitions/auth.BasicUserInfo"
"$ref": "#/definitions/dto.BasicUserInfo"
}
}
}
@@ -5121,7 +5121,7 @@
"type": "object",
"properties": {
"data": {
"$ref": "#/definitions/auth.OAuthAuthorizeResponse"
"$ref": "#/definitions/dto.OAuthAuthorizeResponse"
}
}
}
@@ -5438,7 +5438,7 @@
"type": "object",
"properties": {
"data": {
"$ref": "#/definitions/auth.BasicUserInfo"
"$ref": "#/definitions/dto.BasicUserInfo"
}
}
}
@@ -6044,180 +6044,6 @@
}
},
"definitions": {
"auth.AuthSourceView": {
"type": "object",
"properties": {
"client_secret_configured": {
"type": "boolean"
},
"display_name": {
"type": "string"
},
"icon_url": {
"type": "string"
},
"id": {
"type": "integer"
},
"is_active": {
"type": "boolean"
},
"name": {
"type": "string"
},
"type": {
"type": "string"
}
}
},
"auth.BasicUserInfo": {
"type": "object",
"properties": {
"avatar_url": {
"type": "string"
},
"bio": {
"type": "string"
},
"email": {
"type": "string"
},
"gender": {
"type": "string"
},
"id": {
"type": "string",
"example": "0"
},
"is_admin": {
"type": "boolean"
},
"location": {
"type": "string"
},
"need_change_password": {
"type": "boolean"
},
"nickname": {
"type": "string"
},
"phone": {
"type": "string"
},
"username": {
"type": "string"
},
"website": {
"type": "string"
}
}
},
"auth.CallbackRequest": {
"type": "object",
"required": [
"code",
"state"
],
"properties": {
"code": {
"type": "string"
},
"state": {
"type": "string"
}
}
},
"auth.ChallengeResponse": {
"type": "object",
"properties": {
"challenge": {
"type": "object",
"properties": {
"c": {
"type": "integer"
},
"d": {
"type": "integer"
},
"s": {
"type": "integer"
}
}
},
"expires": {
"description": "ms timestamp",
"type": "integer"
},
"token": {
"type": "string"
}
}
},
"auth.OAuthAuthorizeResponse": {
"type": "object",
"properties": {
"authorize_url": {
"type": "string"
}
}
},
"auth.OAuthCallbackResult": {
"type": "object",
"properties": {
"status": {
"type": "string"
},
"user": {
"$ref": "#/definitions/auth.BasicUserInfo"
}
}
},
"auth.RedeemResponse": {
"type": "object",
"properties": {
"error": {
"type": "string"
},
"expires": {
"type": "integer"
},
"success": {
"type": "boolean"
},
"token": {
"type": "string"
}
}
},
"auth.challengeRequest": {
"type": "object",
"properties": {
"scope": {
"type": "string"
}
}
},
"auth.redeemRequest": {
"type": "object",
"required": [
"solutions",
"token"
],
"properties": {
"scope": {
"type": "string"
},
"solutions": {
"type": "array",
"items": {
"type": "integer"
}
},
"token": {
"type": "string"
}
}
},
"contracts.AuthSourceDTO": {
"type": "object",
"properties": {
@@ -6671,6 +6497,180 @@
}
}
},
"dto.AuthSourceView": {
"type": "object",
"properties": {
"client_secret_configured": {
"type": "boolean"
},
"display_name": {
"type": "string"
},
"icon_url": {
"type": "string"
},
"id": {
"type": "integer"
},
"is_active": {
"type": "boolean"
},
"name": {
"type": "string"
},
"type": {
"type": "string"
}
}
},
"dto.BasicUserInfo": {
"type": "object",
"properties": {
"avatar_url": {
"type": "string"
},
"bio": {
"type": "string"
},
"email": {
"type": "string"
},
"gender": {
"type": "string"
},
"id": {
"type": "string",
"example": "0"
},
"is_admin": {
"type": "boolean"
},
"location": {
"type": "string"
},
"need_change_password": {
"type": "boolean"
},
"nickname": {
"type": "string"
},
"phone": {
"type": "string"
},
"username": {
"type": "string"
},
"website": {
"type": "string"
}
}
},
"dto.CallbackRequest": {
"type": "object",
"required": [
"code",
"state"
],
"properties": {
"code": {
"type": "string"
},
"state": {
"type": "string"
}
}
},
"dto.ChallengeRequest": {
"type": "object",
"properties": {
"scope": {
"type": "string"
}
}
},
"dto.ChallengeResponse": {
"type": "object",
"properties": {
"challenge": {
"type": "object",
"properties": {
"c": {
"type": "integer"
},
"d": {
"type": "integer"
},
"s": {
"type": "integer"
}
}
},
"expires": {
"description": "ms timestamp",
"type": "integer"
},
"token": {
"type": "string"
}
}
},
"dto.OAuthAuthorizeResponse": {
"type": "object",
"properties": {
"authorize_url": {
"type": "string"
}
}
},
"dto.OAuthCallbackResult": {
"type": "object",
"properties": {
"status": {
"type": "string"
},
"user": {
"$ref": "#/definitions/dto.BasicUserInfo"
}
}
},
"dto.RedeemRequest": {
"type": "object",
"required": [
"solutions",
"token"
],
"properties": {
"scope": {
"type": "string"
},
"solutions": {
"type": "array",
"items": {
"type": "integer"
}
},
"token": {
"type": "string"
}
}
},
"dto.RedeemResponse": {
"type": "object",
"properties": {
"error": {
"type": "string"
},
"expires": {
"type": "integer"
},
"success": {
"type": "boolean"
},
"token": {
"type": "string"
}
}
},
"entity.PushChannel": {
"type": "object",
"properties": {
+127 -127
View File
@@ -1,119 +1,5 @@
basePath: /
definitions:
auth.AuthSourceView:
properties:
client_secret_configured:
type: boolean
display_name:
type: string
icon_url:
type: string
id:
type: integer
is_active:
type: boolean
name:
type: string
type:
type: string
type: object
auth.BasicUserInfo:
properties:
avatar_url:
type: string
bio:
type: string
email:
type: string
gender:
type: string
id:
example: "0"
type: string
is_admin:
type: boolean
location:
type: string
need_change_password:
type: boolean
nickname:
type: string
phone:
type: string
username:
type: string
website:
type: string
type: object
auth.CallbackRequest:
properties:
code:
type: string
state:
type: string
required:
- code
- state
type: object
auth.ChallengeResponse:
properties:
challenge:
properties:
c:
type: integer
d:
type: integer
s:
type: integer
type: object
expires:
description: ms timestamp
type: integer
token:
type: string
type: object
auth.OAuthAuthorizeResponse:
properties:
authorize_url:
type: string
type: object
auth.OAuthCallbackResult:
properties:
status:
type: string
user:
$ref: '#/definitions/auth.BasicUserInfo'
type: object
auth.RedeemResponse:
properties:
error:
type: string
expires:
type: integer
success:
type: boolean
token:
type: string
type: object
auth.challengeRequest:
properties:
scope:
type: string
type: object
auth.redeemRequest:
properties:
scope:
type: string
solutions:
items:
type: integer
type: array
token:
type: string
required:
- solutions
- token
type: object
contracts.AuthSourceDTO:
properties:
client_id:
@@ -413,6 +299,120 @@ definitions:
required:
- template
type: object
dto.AuthSourceView:
properties:
client_secret_configured:
type: boolean
display_name:
type: string
icon_url:
type: string
id:
type: integer
is_active:
type: boolean
name:
type: string
type:
type: string
type: object
dto.BasicUserInfo:
properties:
avatar_url:
type: string
bio:
type: string
email:
type: string
gender:
type: string
id:
example: "0"
type: string
is_admin:
type: boolean
location:
type: string
need_change_password:
type: boolean
nickname:
type: string
phone:
type: string
username:
type: string
website:
type: string
type: object
dto.CallbackRequest:
properties:
code:
type: string
state:
type: string
required:
- code
- state
type: object
dto.ChallengeRequest:
properties:
scope:
type: string
type: object
dto.ChallengeResponse:
properties:
challenge:
properties:
c:
type: integer
d:
type: integer
s:
type: integer
type: object
expires:
description: ms timestamp
type: integer
token:
type: string
type: object
dto.OAuthAuthorizeResponse:
properties:
authorize_url:
type: string
type: object
dto.OAuthCallbackResult:
properties:
status:
type: string
user:
$ref: '#/definitions/dto.BasicUserInfo'
type: object
dto.RedeemRequest:
properties:
scope:
type: string
solutions:
items:
type: integer
type: array
token:
type: string
required:
- solutions
- token
type: object
dto.RedeemResponse:
properties:
error:
type: string
expires:
type: integer
success:
type: boolean
token:
type: string
type: object
entity.PushChannel:
properties:
created_at:
@@ -4031,7 +4031,7 @@ paths:
in: body
name: request
schema:
$ref: '#/definitions/auth.challengeRequest'
$ref: '#/definitions/dto.ChallengeRequest'
produces:
- application/json
responses:
@@ -4042,7 +4042,7 @@ paths:
- $ref: '#/definitions/response.Any'
- properties:
data:
$ref: '#/definitions/auth.ChallengeResponse'
$ref: '#/definitions/dto.ChallengeResponse'
type: object
"500":
description: 内部服务错误
@@ -4060,7 +4060,7 @@ paths:
in: body
name: request
schema:
$ref: '#/definitions/auth.challengeRequest'
$ref: '#/definitions/dto.ChallengeRequest'
produces:
- application/json
responses:
@@ -4071,7 +4071,7 @@ paths:
- $ref: '#/definitions/response.Any'
- properties:
data:
$ref: '#/definitions/auth.ChallengeResponse'
$ref: '#/definitions/dto.ChallengeResponse'
type: object
"500":
description: 内部服务错误
@@ -4091,7 +4091,7 @@ paths:
name: request
required: true
schema:
$ref: '#/definitions/auth.redeemRequest'
$ref: '#/definitions/dto.RedeemRequest'
produces:
- application/json
responses:
@@ -4102,7 +4102,7 @@ paths:
- $ref: '#/definitions/response.Any'
- properties:
data:
$ref: '#/definitions/auth.RedeemResponse'
$ref: '#/definitions/dto.RedeemResponse'
type: object
"400":
description: 参数错误或核销失败
@@ -4271,7 +4271,7 @@ paths:
- $ref: '#/definitions/response.Any'
- properties:
data:
$ref: '#/definitions/auth.OAuthAuthorizeResponse'
$ref: '#/definitions/dto.OAuthAuthorizeResponse'
type: object
"400":
description: 认证源不存在或未启用
@@ -4295,7 +4295,7 @@ paths:
name: request
required: true
schema:
$ref: '#/definitions/auth.CallbackRequest'
$ref: '#/definitions/dto.CallbackRequest'
produces:
- application/json
responses:
@@ -4306,7 +4306,7 @@ paths:
- $ref: '#/definitions/response.Any'
- properties:
data:
$ref: '#/definitions/auth.OAuthCallbackResult'
$ref: '#/definitions/dto.OAuthCallbackResult'
type: object
"400":
description: state 无效、参数错误或认证源错误
@@ -4399,7 +4399,7 @@ paths:
- $ref: '#/definitions/response.Any'
- properties:
data:
$ref: '#/definitions/auth.OAuthAuthorizeResponse'
$ref: '#/definitions/dto.OAuthAuthorizeResponse'
type: object
"400":
description: 认证源不存在或未配置
@@ -4450,7 +4450,7 @@ paths:
- properties:
data:
items:
$ref: '#/definitions/auth.AuthSourceView'
$ref: '#/definitions/dto.AuthSourceView'
type: array
type: object
summary: 获取可用登录源
@@ -4469,7 +4469,7 @@ paths:
- $ref: '#/definitions/response.Any'
- properties:
data:
$ref: '#/definitions/auth.BasicUserInfo'
$ref: '#/definitions/dto.BasicUserInfo'
type: object
"401":
description: 未登录
@@ -4656,7 +4656,7 @@ paths:
- $ref: '#/definitions/response.Any'
- properties:
data:
$ref: '#/definitions/auth.BasicUserInfo'
$ref: '#/definitions/dto.BasicUserInfo'
type: object
"401":
description: 未登录
@@ -1,217 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"context"
"errors"
"fmt"
"strconv"
"strings"
"github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2"
)
func isOIDCLoginEnabled(ctx context.Context) bool {
val, err := GetSystemConfigValue(ctx, "oidc_login_enabled")
if err != nil || val == "" {
return true
}
b, err := strconv.ParseBool(val)
if err != nil {
return true
}
return b
}
func resolveAuthSource(ctx context.Context, sourceName string) (*AuthSource, error) {
name := strings.TrimSpace(strings.ToLower(sourceName))
if name == "" {
sources, err := GetActiveAuthSourcesCached(ctx)
if err != nil {
return nil, err
}
if len(sources) == 0 {
return nil, errors.New(errNoActiveAuthSource)
}
src, err := GetAuthSourceByNameCached(ctx, sources[0].Name)
if err != nil {
return nil, err
}
return src, nil
}
src, err := GetAuthSourceByNameCached(ctx, name)
if err != nil {
return nil, err
}
return src, nil
}
func activeLoginSources(ctx context.Context) []AuthSourceView {
if !isOIDCLoginEnabled(ctx) {
return nil
}
dbSources, err := GetActiveAuthSourcesCached(ctx)
if err != nil {
return nil
}
sources := make([]AuthSourceView, 0, len(dbSources))
for _, source := range dbSources {
sources = append(sources, AuthSourceView{
ID: source.ID,
Name: source.Name,
Type: source.Type,
DisplayName: source.DisplayName,
IsActive: source.IsActive,
IconURL: source.IconURL,
ClientSecretConfigured: source.ClientSecretConfigured,
})
}
return sources
}
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
val, err := GetSystemConfigValue(ctx, "server_address")
if err != nil || strings.TrimSpace(val) == "" {
return "", errors.New(errServerAddressMissing)
}
return strings.TrimRight(val, "/") + "/login", nil
}
func buildOAuthConfig(ctx context.Context, source *AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
if source == nil {
return nil, nil, errors.New(errAuthSourceRequired)
}
if source.OpenIDDiscoveryURL == "" {
return nil, nil, errors.New(errDiscoveryURLRequired)
}
// Clean the issuer URL
issuer := strings.TrimSuffix(strings.TrimSpace(source.OpenIDDiscoveryURL), "/")
issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration")
issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server")
provider, err := globalOIDCProviderCache.get(ctx, issuer)
if err != nil {
return nil, nil, err
}
verifier := provider.Verifier(&oidc.Config{ClientID: source.ClientID})
scopes := strings.Fields(source.Scopes)
if len(scopes) == 0 {
scopes = []string{oidc.ScopeOpenID, "profile", "email"}
}
if !containsScope(scopes, oidc.ScopeOpenID) {
scopes = append([]string{oidc.ScopeOpenID}, scopes...)
}
return &oauth2.Config{
ClientID: source.ClientID,
ClientSecret: source.ClientSecret,
RedirectURL: redirectURL,
Scopes: scopes,
Endpoint: provider.Endpoint(),
}, verifier, nil
}
func containsScope(scopes []string, scope string) bool {
for _, item := range scopes {
if item == scope {
return true
}
}
return false
}
func buildOAuthUserInfo(ctx context.Context, source *AuthSource, code, nonce, redirectURL string) (*contracts.OAuthUserInfoDTO, error) {
authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
if err != nil {
return nil, err
}
token, err := authConfig.Exchange(ctx, code)
if err != nil {
return nil, err
}
userInfo := &contracts.OAuthUserInfoDTO{Active: true}
if verifier != nil {
if verifyErr := verifyIDToken(ctx, verifier, token, nonce, userInfo); verifyErr != nil {
return nil, verifyErr
}
}
if userInfo.Username == "" && userInfo.PreferredUsername != "" {
userInfo.Username = userInfo.PreferredUsername
}
if userInfo.Username == "" && userInfo.Email != "" {
userInfo.Username = strings.Split(userInfo.Email, "@")[0]
}
if userInfo.Username == "" && userInfo.Sub != "" {
userInfo.Username = userInfo.Sub
}
if userInfo.Name == "" {
userInfo.Name = userInfo.Username
}
return userInfo, nil
}
func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *contracts.OAuthUserInfoDTO) error {
rawIDToken, ok := token.Extra("id_token").(string)
if !ok {
return nil
}
idToken, verifyErr := verifier.Verify(ctx, rawIDToken)
if verifyErr != nil {
return fmt.Errorf(errIDTokenVerifyFailedFormat, errIDTokenVerifyFailed, verifyErr)
}
if nonce != "" && idToken.Nonce != nonce {
return errors.New(errNonceMismatch)
}
if claimsErr := idToken.Claims(userInfo); claimsErr != nil {
return claimsErr
}
return nil
}
func normalizeOAuthUserInfo(userInfo *contracts.OAuthUserInfoDTO) error {
userInfo.Username = strings.TrimSpace(userInfo.Username)
userInfo.PreferredUsername = strings.TrimSpace(userInfo.PreferredUsername)
userInfo.Email = strings.TrimSpace(userInfo.Email)
userInfo.Name = strings.TrimSpace(userInfo.Name)
userInfo.AvatarURL = strings.TrimSpace(userInfo.AvatarURL)
if userInfo.Username == "" && userInfo.PreferredUsername != "" {
userInfo.Username = userInfo.PreferredUsername
}
if userInfo.Username == "" && userInfo.Email != "" {
userInfo.Username = strings.Split(userInfo.Email, "@")[0]
}
if userInfo.Username == "" && userInfo.Sub != "" {
userInfo.Username = userInfo.Sub
}
if userInfo.Username == "" {
return errors.New(errUsernameFromSourceFailed)
}
if userInfo.Name == "" {
userInfo.Name = userInfo.Username
}
if !userInfo.Active {
userInfo.Active = true
}
return nil
}
func buildCallbackResult(user *contracts.UserDTO, status string) OAuthCallbackResult {
result := OAuthCallbackResult{Status: status}
if user != nil {
info := BuildBasicUserInfo(user, false)
result.User = &info
}
return result
}
-23
View File
@@ -1,23 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// HTTP 响应错误文案
const (
errCapTokenMissing = "验证码验证失败,缺少验证码凭证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errCapTokenInvalidOrExpired = "验证码校验失败或已过期,请重试" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errCapNotConfigured = "captcha is not configured"
errChallengeGenerateFailed = "生成验证难题失败,请稍后再试"
errInvalidRequestParams = "无效的参数"
errSolutionVerifyFailed = "校验验证解答失败,请稍后再试"
)
// Redeem 结果码,属于 redeem 响应 JSON 的对外契约取值,禁止改写取值
const (
redeemErrInvalidToken = "invalid_token"
redeemErrNonceStoreFailed = "nonce_store_error"
redeemErrAlreadyRedeemed = "already_redeemed"
redeemErrSettingsLoad = "settings_load_error"
redeemErrTokenStoreFailed = "token_store_error" //nolint:gosec // error code, not hardcoded credentials
)
@@ -1,38 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
// VerifyCaptchaMiddleware returns a Gin middleware that checks and consumes the X-Cap-Token header.
func VerifyCaptchaMiddleware(mgr *CaptchaManager, scope string) gin.HandlerFunc {
return func(c *gin.Context) {
if !CapProtectionEnabled(c.Request.Context()) {
c.Next()
return
}
if mgr == nil {
response.AbortBadRequest(c, errCapTokenInvalidOrExpired)
return
}
token := c.GetHeader("X-Cap-Token")
if token == "" {
response.AbortBadRequest(c, errCapTokenMissing)
return
}
valid, err := mgr.VerifyToken(c.Request.Context(), token, scope)
if err != nil || !valid {
response.AbortBadRequest(c, errCapTokenInvalidOrExpired)
return
}
c.Next()
}
}
@@ -1,197 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
"errors"
"strconv"
"sync/atomic"
"time"
"golang.org/x/sync/singleflight"
)
const (
defaultCapChallengeCount = 1
defaultCapChallengeSize = 32
defaultCapChallengeDifficulty = 4
defaultCapChallengeTTL = 10 * time.Minute
defaultCapTokenTTL = 20 * time.Minute
)
// CapRuntimeSettings is the parsed CAPTCHA runtime configuration loaded from system_configs.
type CapRuntimeSettings struct {
LoginEnabled bool
ChallengeCount int
ChallengeSize int
ChallengeDifficulty int
ChallengeTTL time.Duration
TokenTTL time.Duration
}
// CAP 动态配置键常量
const (
ConfigKeyCapLoginEnabled = "cap_login_enabled"
ConfigKeyCapChallengeCount = "cap_challenge_count"
ConfigKeyCapChallengeSize = "cap_challenge_size"
ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty"
ConfigKeyCapChallengeTTL = "cap_challenge_ttl"
// ConfigKeyCapTokenTTL 验证码 Token 过期时间键
// #nosec G101
ConfigKeyCapTokenTTL = "cap_token_ttl"
)
var capRuntimeConfigKeys = []string{
ConfigKeyCapLoginEnabled,
ConfigKeyCapChallengeCount,
ConfigKeyCapChallengeSize,
ConfigKeyCapChallengeDifficulty,
ConfigKeyCapChallengeTTL,
ConfigKeyCapTokenTTL,
}
var capRuntimeConfigKeySet = func() map[string]struct{} {
set := make(map[string]struct{}, len(capRuntimeConfigKeys))
for _, key := range capRuntimeConfigKeys {
set[key] = struct{}{}
}
return set
}()
type capRuntimeSettingsStore struct {
snapshot atomic.Pointer[CapRuntimeSettings]
loadGroup singleflight.Group
}
var capSettingsStore = &capRuntimeSettingsStore{}
// IsCapRuntimeConfigKey reports whether a system config key affects CAPTCHA runtime settings.
func IsCapRuntimeConfigKey(key string) bool {
_, ok := capRuntimeConfigKeySet[key]
return ok
}
// CurrentCapSettings returns the cached CAPTCHA runtime settings snapshot.
func CurrentCapSettings(ctx context.Context) (CapRuntimeSettings, error) {
return capSettingsStore.current(ctx)
}
// CapProtectionEnabled reports whether CAPTCHA verification is required for protected routes.
func CapProtectionEnabled(ctx context.Context) bool {
settings, err := CurrentCapSettings(ctx)
if err != nil {
return false
}
return settings.LoginEnabled
}
// InvalidateCapRuntimeSettings drops the in-process CAPTCHA settings snapshot.
func InvalidateCapRuntimeSettings() {
capSettingsStore.snapshot.Store(nil)
}
// ResetCapRuntimeSettingsForTest clears the CAPTCHA runtime snapshot.
func ResetCapRuntimeSettingsForTest() {
InvalidateCapRuntimeSettings()
}
// InstallCapTestRuntimeSettings installs a fixed snapshot for unit tests.
func InstallCapTestRuntimeSettings(settings CapRuntimeSettings) func() {
snapshot := settings
capSettingsStore.snapshot.Store(&snapshot)
return InvalidateCapRuntimeSettings
}
func (s *capRuntimeSettingsStore) current(ctx context.Context) (CapRuntimeSettings, error) {
if snapshot := s.snapshot.Load(); snapshot != nil {
return *snapshot, nil
}
loaded, err, _ := s.loadGroup.Do("cap-runtime-settings", func() (any, error) {
if snapshot := s.snapshot.Load(); snapshot != nil {
return *snapshot, nil
}
settings, loadErr := loadCapRuntimeSettings(ctx)
if loadErr != nil {
return CapRuntimeSettings{}, loadErr
}
s.snapshot.Store(&settings)
return settings, nil
})
if err != nil {
return CapRuntimeSettings{}, err
}
settings, ok := loaded.(CapRuntimeSettings)
if !ok {
return CapRuntimeSettings{}, errors.New("cap runtime settings loader returned unexpected type")
}
return settings, nil
}
func loadCapRuntimeSettings(ctx context.Context) (CapRuntimeSettings, error) {
var records []capConfigRecord
db := getDB(ctx)
if db == nil {
return parseCapRuntimeSettings(nil), nil
}
if err := db.Table("w_system_configs").Where("key IN ?", capRuntimeConfigKeys).Find(&records).Error; err != nil {
return CapRuntimeSettings{}, err
}
configs := make(map[string]string, len(records))
for _, r := range records {
configs[r.Key] = r.Value
}
return parseCapRuntimeSettings(configs), nil
}
func parseCapRuntimeSettings(configs map[string]string) CapRuntimeSettings {
settings := CapRuntimeSettings{
ChallengeCount: defaultCapChallengeCount,
ChallengeSize: defaultCapChallengeSize,
ChallengeDifficulty: defaultCapChallengeDifficulty,
ChallengeTTL: defaultCapChallengeTTL,
TokenTTL: defaultCapTokenTTL,
}
if len(configs) == 0 {
return settings
}
if val, ok := configs[ConfigKeyCapLoginEnabled]; ok {
if enabled, err := strconv.ParseBool(val); err == nil {
settings.LoginEnabled = enabled
}
}
if val, ok := configs[ConfigKeyCapChallengeCount]; ok {
if count, err := strconv.Atoi(val); err == nil && count > 0 {
settings.ChallengeCount = count
}
}
if val, ok := configs[ConfigKeyCapChallengeSize]; ok {
if size, err := strconv.Atoi(val); err == nil && size > 0 {
settings.ChallengeSize = size
}
}
if val, ok := configs[ConfigKeyCapChallengeDifficulty]; ok {
if diff, err := strconv.Atoi(val); err == nil && diff > 0 {
settings.ChallengeDifficulty = diff
}
}
if val, ok := configs[ConfigKeyCapChallengeTTL]; ok {
if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 {
settings.ChallengeTTL = time.Duration(ttlSeconds) * time.Second
}
}
if val, ok := configs[ConfigKeyCapTokenTTL]; ok {
if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 {
settings.TokenTTL = time.Duration(ttlSeconds) * time.Second
}
}
return settings
}
+3 -8
View File
@@ -3,12 +3,7 @@
package auth
import "Wavelet/plugins/domain/auth/service"
// SessionConfig defines the session configuration declared by the auth plugin.
type SessionConfig struct {
SessionCookieName string `config:"session_cookie_name" env:"APP_SESSION_COOKIE_NAME" default:"wavelet_session"`
SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"`
SessionDomain string `config:"session_domain" env:"APP_SESSION_DOMAIN"`
SessionAge int `config:"session_age" env:"APP_SESSION_AGE" default:"86400"`
SessionHTTPOnly bool `config:"session_http_only" env:"APP_SESSION_HTTP_ONLY" default:"true"`
SessionSecure bool `config:"session_secure" env:"APP_SESSION_SECURE"`
}
type SessionConfig = service.SessionConfig
+52
View File
@@ -0,0 +1,52 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package consts defines constants, keys, and TTL values for the auth domain plugin.
package consts
import "time"
// CAP 默认参数
const (
DefaultCapChallengeCount = 1
DefaultCapChallengeSize = 32
DefaultCapChallengeDifficulty = 4
DefaultCapChallengeTTL = 10 * time.Minute
DefaultCapTokenTTL = 20 * time.Minute
RedeemTokenIDLength = 8 // 兑换 Token ID 字节长度
RedeemVerTokenLength = 15 // 兑换验证 Token 字节长度
TokenPartsCount = 2 // 兑换 Token 由两部分组成 (id:token)
ValuePartsCount = 2 // 存储值由 scope 和过期时间组成 (expNano|scope)
)
// CAP 动态配置键常量
const (
ConfigKeyCapLoginEnabled = "cap_login_enabled"
ConfigKeyCapChallengeCount = "cap_challenge_count"
ConfigKeyCapChallengeSize = "cap_challenge_size"
ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty"
ConfigKeyCapChallengeTTL = "cap_challenge_ttl"
// ConfigKeyCapTokenTTL 验证码 Token 过期时间键
// #nosec G101
ConfigKeyCapTokenTTL = "cap_token_ttl"
)
// HTTP 响应错误文案
const (
ErrCapTokenMissing = "验证码验证失败,缺少验证码凭证" //nolint:gosec // error message constant
ErrCapTokenInvalidOrExpired = "验证码校验失败或已过期,请重试" //nolint:gosec // error message constant
ErrCapNotConfigured = "captcha is not configured"
ErrChallengeGenerateFailed = "生成验证难题失败,请稍后再试"
ErrInvalidRequestParams = "无效的参数"
ErrSolutionVerifyFailed = "校验验证解答失败,请稍后再试"
)
// Redeem 结果码,属于 redeem 响应 JSON 的对外契约取值,禁止改写取值
const (
RedeemErrInvalidToken = "invalid_token"
RedeemErrNonceStoreFailed = "nonce_store_error"
RedeemErrAlreadyRedeemed = "already_redeemed"
RedeemErrSettingsLoad = "settings_load_error"
RedeemErrTokenStoreFailed = "token_store_error" //nolint:gosec // error code constant
)
@@ -1,11 +1,10 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// Package consts defines constants, keys, and TTL values for the auth domain plugin.
package consts
import (
"time"
)
import "time"
// Session and Context Keys
const (
@@ -14,7 +13,7 @@ const (
UserObjKey = "user_obj"
TokenAuthKey = "token_auth" // 标记当前请求是否通过 Access Token 鉴权
TokenAdminKey = "token_admin" // Access Token 本身是否具有管理员权限
SessionTokenKey = "oauth_session_token" //nolint:gosec // false positive: this is a session key, not hardcoded credentials
SessionTokenKey = "oauth_session_token" //nolint:gosec // false positive: session state key
PasswordHashKey = "password_hash"
SystemUsername = "system"
)
@@ -23,8 +22,8 @@ const (
const (
OAuthStateCacheKeyFormat = "oauth:state:%s"
OAuthStateCacheKeyExpiration = 10 * time.Minute
oauthStateLimitKeyFormat = "oauth:state:limit:%s"
oauthStateLimitMax = 10
OAuthStateLimitKeyFormat = "oauth:state:limit:%s"
OAuthStateLimitMax = 10
)
// OAuth Purpose Constants
@@ -37,3 +36,9 @@ const (
const (
AuthSourceTypeOIDC = "oidc"
)
// Cache TTLs
const (
TokenCacheTTL = 5 * time.Minute
UserCacheTTL = 5 * time.Minute
)
@@ -0,0 +1,53 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package consts defines constants, keys, and TTL values for the auth domain plugin.
package consts
// OAuth and Auth error messages
const (
ErrInvalidState = "非法登录请求"
ErrIDTokenVerifyFailed = "ID Token 验证失败" //nolint:gosec // error message constant
ErrIDTokenVerifyFailedFormat = "%s: %w"
ErrNonceMismatch = "nonce 不匹配,可能存在重放攻击"
ErrNoActiveAuthSource = "未配置可用认证源"
ErrServerAddressMissing = "服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试"
ErrAuthSourceRequired = "认证源不能为空"
ErrDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
ErrUsernameGenerateFailed = "无法生成可用用户名"
ErrUsernameFromSourceFailed = "无法从认证源获取用户名"
ErrAuthSourceDisabled = "认证源未启用"
ErrInvalidExternalAccountBindingID = "绑定记录 ID 无效"
ErrTokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // error message constant
ErrOAuthStateRateLimited = "请求授权过于频繁,请稍后重试"
ErrAuthSourceNameRequired = "认证源名称不能为空"
ErrAuthSourceNameInvalid = "认证源名称格式不正确"
ErrAuthSourceTypeUnsupported = "不支持的认证源类型"
ErrAuthSourceDiscoveryURLRequired = "Discovery URL 不能为空"
//nolint:gosec // error message constant
ErrAuthSourceClientCredentialsRequired = "启用认证源时必须配置 Client ID 和 Client Secret"
ErrAuthSourceIDRequired = "认证源 ID 不能为空"
ErrUserIDRequired = "用户 ID 不能为空"
ErrExternalAccountBindingIncomplete = "外部帐号绑定信息不完整"
ErrExternalAccountAlreadyBoundToAnother = "该外部帐号已被其他用户绑定"
ErrExternalAccountBindingIDRequired = "外部帐号绑定记录 ID 不能为空"
ErrInsufficientPermission = "权限不足"
ErrBannedAccount = "账号已被封禁"
ErrUnAuthorized = "未登录"
)
// Service 层与鉴权中间件内部错误文案
const (
ErrUserNotInContext = "auth: user not found in context"
ErrEmptyToken = "auth: empty token" //nolint:gosec // error message constant
ErrSystemUserTokenNotAllowed = "auth: system user token not allowed" //nolint:gosec // error message constant
ErrUnauthorizedInternal = "unauthorized"
ErrSystemUserLoginNotAllowed = "system user is not allowed to login"
)
// OAuth 回调会话校验错误文案
const (
ErrInvalidSessionContext = "invalid session context"
ErrSessionMismatchForOAuth = "session mismatch for oauth state"
ErrUserContextMismatch = "user context mismatch for oauth binding"
)
@@ -1,11 +1,13 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// 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"
@@ -17,7 +19,7 @@ func LogForAudit(ctx context.Context, user *contracts.UserDTO, c *gin.Context) {
if user == nil || c == nil {
return
}
auditLog := loginRequiredAuditLog{
auditLog := dto.LoginRequiredAuditLog{
UserID: user.ID,
Username: user.Username,
ClientIP: c.ClientIP(),
@@ -1,44 +1,59 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// 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 challengeRequest false "可选范围限制参数"
// @Success 200 {object} response.Any{data=auth.ChallengeResponse} "成功返回 PoW 难题"
// @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 Challenge(c *gin.Context) {
var req challengeRequest
func (h *CaptchaHandler) Challenge(c *gin.Context) {
var req dto.ChallengeRequest
_ = c.ShouldBind(&req) // 允许不传 body,默认使用 login scope
if req.Scope == "" {
req.Scope = "login"
}
mgr := GetDefaultCapManager()
if mgr == nil {
response.AbortInternal(c, errCapNotConfigured)
if h.capMgr == nil {
response.AbortInternal(c, consts.ErrCapNotConfigured)
return
}
resp, err := mgr.Generate(c.Request.Context(), req.Scope)
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, errChallengeGenerateFailed)
response.AbortInternal(c, consts.ErrChallengeGenerateFailed)
return
}
@@ -51,15 +66,15 @@ func Challenge(c *gin.Context) {
// @Tags cap
// @Accept json
// @Produce json
// @Param request body redeemRequest true "难题 Token 与解答 solutions 数组"
// @Success 200 {object} response.Any{data=auth.RedeemResponse} "核销成功,返回 X-Cap-Token"
// @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 Redeem(c *gin.Context) {
var req redeemRequest
func (h *CaptchaHandler) Redeem(c *gin.Context) {
var req dto.RedeemRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, errInvalidRequestParams)
response.AbortBadRequest(c, consts.ErrInvalidRequestParams)
return
}
@@ -67,15 +82,14 @@ func Redeem(c *gin.Context) {
req.Scope = "login"
}
mgr := GetDefaultCapManager()
if mgr == nil {
response.AbortInternal(c, errCapNotConfigured)
if h.capMgr == nil {
response.AbortInternal(c, consts.ErrCapNotConfigured)
return
}
resp, err := mgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope)
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, errSolutionVerifyFailed)
response.AbortInternal(c, consts.ErrSolutionVerifyFailed)
return
}
@@ -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)),
)
}
@@ -0,0 +1,61 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package dao provides data access objects and caching for the auth domain plugin.
package dao
import (
"Wavelet/plugins/domain/auth/model/entity"
"context"
)
// ListAllAuthSources 获取全部认证源(含未启用),按 ID 升序
func (d *DAO) ListAllAuthSources(ctx context.Context) ([]entity.AuthSource, error) {
var sources []entity.AuthSource
if err := d.DB(ctx).Order("id ASC").Find(&sources).Error; err != nil {
return nil, err
}
return sources, nil
}
// GetAuthSourceByID 根据 ID 获取认证源
func (d *DAO) GetAuthSourceByID(ctx context.Context, id uint64) (*entity.AuthSource, error) {
var src entity.AuthSource
if err := d.DB(ctx).First(&src, id).Error; err != nil {
return nil, err
}
return &src, nil
}
// GetAuthSourceByName 根据名称获取认证源
func (d *DAO) GetAuthSourceByName(ctx context.Context, name string) (*entity.AuthSource, error) {
var src entity.AuthSource
if err := d.DB(ctx).Where("name = ?", name).First(&src).Error; err != nil {
return nil, err
}
return &src, nil
}
// ListActiveAuthSources 获取所有启用的认证源
func (d *DAO) ListActiveAuthSources(ctx context.Context) ([]entity.AuthSource, error) {
var sources []entity.AuthSource
if err := d.DB(ctx).Where("is_active = ?", true).Order("id ASC").Find(&sources).Error; err != nil {
return nil, err
}
return sources, nil
}
// CreateAuthSource 新建认证源记录
func (d *DAO) CreateAuthSource(ctx context.Context, source *entity.AuthSource) error {
return d.DB(ctx).Create(source).Error
}
// SaveAuthSource 全量保存认证源记录
func (d *DAO) SaveAuthSource(ctx context.Context, source *entity.AuthSource) error {
return d.DB(ctx).Save(source).Error
}
// DeleteAuthSource 删除认证源记录
func (d *DAO) DeleteAuthSource(ctx context.Context, source *entity.AuthSource) error {
return d.DB(ctx).Delete(source).Error
}
@@ -1,30 +1,20 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// Package dao provides data access objects and caching for the auth domain plugin.
package dao
import (
"Wavelet/core/contracts"
"Wavelet/pkg/cache/ram"
"Wavelet/plugins/domain/auth/consts"
"Wavelet/plugins/domain/auth/model/do"
"context"
"fmt"
"time"
)
const (
tokenCacheTTL = 5 * time.Minute
userCacheTTL = 5 * time.Minute
)
// CachedToken represents the minimal cached representation of an access token.
type CachedToken struct {
ID uint64 `json:"id"`
UserID uint64 `json:"user_id"`
IsAdmin bool `json:"is_admin"`
}
var (
tokenRAM = ram.MustNew[string, *CachedToken](ram.Options{MaximumSize: 2048})
tokenRAM = ram.MustNew[string, *do.CachedToken](ram.Options{MaximumSize: 2048})
userRAM = ram.MustNew[uint64, *contracts.UserDTO](ram.Options{MaximumSize: 2048})
)
@@ -37,13 +27,15 @@ func userCacheKey(userID uint64) string {
}
// GetCachedToken 获取缓存的 Token
func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error) {
//
//nolint:dupl // token and user cache lookup pattern
func (d *DAO) GetCachedToken(ctx context.Context, tokenHash string) (*do.CachedToken, error) {
if val, ok := tokenRAM.GetIfPresent(tokenHash); ok {
return val, nil
}
if cache := getCache(ctx); cache != nil {
var token CachedToken
if cache := d.Cache(); cache != nil {
var token do.CachedToken
key := tokenCacheKey(tokenHash)
if err := cache.Get(ctx, key, &token); err == nil {
tokenRAM.Set(tokenHash, &token)
@@ -54,30 +46,32 @@ func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error)
}
// SetCachedToken 设置 Token 缓存
func SetCachedToken(ctx context.Context, tokenHash string, token *CachedToken) {
func (d *DAO) SetCachedToken(ctx context.Context, tokenHash string, token *do.CachedToken) {
tokenRAM.Set(tokenHash, token)
if cache := getCache(ctx); cache != nil {
if cache := d.Cache(); cache != nil {
key := tokenCacheKey(tokenHash)
_ = cache.Set(ctx, key, token, tokenCacheTTL)
_ = cache.Set(ctx, key, token, consts.TokenCacheTTL)
}
}
// InvalidateCachedToken 吊销/删除 token 缓存
func InvalidateCachedToken(ctx context.Context, tokenHash string) {
func (d *DAO) InvalidateCachedToken(ctx context.Context, tokenHash string) {
tokenRAM.Invalidate(tokenHash)
if cache := getCache(ctx); cache != nil {
if cache := d.Cache(); cache != nil {
key := tokenCacheKey(tokenHash)
_ = cache.Delete(ctx, key)
}
}
// GetCachedUser 获取缓存的 UserDTO
func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
//
//nolint:dupl // token and user cache lookup pattern
func (d *DAO) GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
if val, ok := userRAM.GetIfPresent(userID); ok {
return val, nil
}
if cache := getCache(ctx); cache != nil {
if cache := d.Cache(); cache != nil {
var u contracts.UserDTO
key := userCacheKey(userID)
if err := cache.Get(ctx, key, &u); err == nil {
@@ -89,28 +83,25 @@ func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, erro
}
// SetCachedUser 设置 UserDTO 缓存
func SetCachedUser(ctx context.Context, userID uint64, u *contracts.UserDTO) {
func (d *DAO) SetCachedUser(ctx context.Context, userID uint64, u *contracts.UserDTO) {
userRAM.Set(userID, u)
if cache := getCache(ctx); cache != nil {
if cache := d.Cache(); cache != nil {
key := userCacheKey(userID)
_ = cache.Set(ctx, key, u, userCacheTTL)
_ = cache.Set(ctx, key, u, consts.UserCacheTTL)
}
}
// InvalidateCachedUser 吊销/失效 UserDTO 缓存
func InvalidateCachedUser(ctx context.Context, userID uint64) {
func (d *DAO) InvalidateCachedUser(ctx context.Context, userID uint64) {
userRAM.Invalidate(userID)
if cache := getCache(ctx); cache != nil {
if cache := d.Cache(); cache != nil {
key := userCacheKey(userID)
_ = cache.Delete(ctx, key)
}
}
// StopAuthCacheListener compatibility stub for tests
func StopAuthCacheListener() {}
// ResetAuthRAMCacheForTest clears only the process-local RAM cache.
func ResetAuthRAMCacheForTest() {
// ResetRAMCacheForTest clears only the process-local RAM cache.
func ResetRAMCacheForTest() {
tokenRAM.InvalidateAll()
userRAM.InvalidateAll()
}
+61
View File
@@ -0,0 +1,61 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package dao provides data access objects and caching for the auth domain plugin.
package dao
import (
"Wavelet/core/contracts"
"context"
"gorm.io/gorm"
)
// DAO aggregates all data access objects for the auth domain plugin.
type DAO struct {
dbSvc contracts.DBService
cacheSvc contracts.CacheService
limiterSvc contracts.LimiterService
}
// New creates a new DAO aggregate.
func New(dbSvc contracts.DBService, cacheSvc contracts.CacheService, limiterSvc contracts.LimiterService) *DAO {
return &DAO{
dbSvc: dbSvc,
cacheSvc: cacheSvc,
limiterSvc: limiterSvc,
}
}
// SetDBService updates the DBService reference.
func (d *DAO) SetDBService(db contracts.DBService) {
d.dbSvc = db
}
// SetCacheService updates the CacheService reference.
func (d *DAO) SetCacheService(cache contracts.CacheService) {
d.cacheSvc = cache
}
// SetLimiterService updates the LimiterService reference.
func (d *DAO) SetLimiterService(limiter contracts.LimiterService) {
d.limiterSvc = limiter
}
// DB returns the GORM DB instance associated with the request context.
func (d *DAO) DB(ctx context.Context) *gorm.DB {
if d.dbSvc != nil {
return d.dbSvc.DB(ctx)
}
return nil
}
// Cache returns the CacheService instance.
func (d *DAO) Cache() contracts.CacheService {
return d.cacheSvc
}
// Limiter returns the LimiterService instance.
func (d *DAO) Limiter() contracts.LimiterService {
return d.limiterSvc
}
@@ -0,0 +1,38 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package dao provides data access objects and caching for the auth domain plugin.
package dao
import (
"Wavelet/plugins/domain/auth/model/entity"
"context"
)
// FindExternalAccount 查询指定认证源的外部账号绑定
func (d *DAO) FindExternalAccount(ctx context.Context, authSourceID uint64, externalID string) (*entity.ExternalAccount, error) {
var account entity.ExternalAccount
if err := d.DB(ctx).Where("auth_source_id = ? AND external_id = ?", authSourceID, externalID).First(&account).Error; err != nil {
return nil, err
}
return &account, nil
}
// BindExternalAccount 绑定外部账号
func (d *DAO) BindExternalAccount(ctx context.Context, account *entity.ExternalAccount) error {
return d.DB(ctx).Create(account).Error
}
// ListExternalAccountsByUserID 获取用户绑定的所有外部账号
func (d *DAO) ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]entity.ExternalAccount, error) {
var accounts []entity.ExternalAccount
if err := d.DB(ctx).Where("user_id = ?", userID).Find(&accounts).Error; err != nil {
return nil, err
}
return accounts, nil
}
// UnbindExternalAccount 解绑外部账号
func (d *DAO) UnbindExternalAccount(ctx context.Context, id, userID uint64) error {
return d.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&entity.ExternalAccount{}).Error
}
@@ -0,0 +1,91 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package dao provides data access objects and caching for the auth domain plugin.
package dao
import (
"Wavelet/core/contracts"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/auth/model/do"
"context"
"time"
)
// GetAccessTokenByHash 按令牌哈希读取访问令牌记录(仅取鉴权所需字段)
func (d *DAO) GetAccessTokenByHash(ctx context.Context, tokenHash string) (*do.CachedToken, error) {
var row struct {
ID uint64
UserID uint64
IsAdmin bool
}
if err := d.DB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&row).Error; err != nil {
return nil, err
}
return &do.CachedToken{
ID: row.ID,
UserID: row.UserID,
IsAdmin: row.IsAdmin,
}, nil
}
// GetActiveUserByID 读取仍处于启用状态的用户
func (d *DAO) GetActiveUserByID(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
var user contracts.UserDTO
if err := d.DB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&user).Error; err != nil {
return nil, err
}
return &user, nil
}
// GetUserByID 按 ID 读取用户(不限制启用状态)
func (d *DAO) GetUserByID(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
var user contracts.UserDTO
if err := d.DB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil {
return nil, err
}
return &user, nil
}
// InsertUser 新建用户记录
func (d *DAO) InsertUser(ctx context.Context, user *contracts.UserDTO) error {
return d.DB(ctx).Table("w_users").Create(user).Error
}
// TouchUserLastLogin 刷新用户最后登录时间
func (d *DAO) TouchUserLastLogin(ctx context.Context, userID uint64, at time.Time) error {
return d.DB(ctx).Table("w_users").Where("id = ?", userID).Update("last_login_at", at).Error
}
// ListSimilarUsernames 查询与基础用户名相同或带 `-序号` 后缀的用户名(用于用户名去重)
func (d *DAO) ListSimilarUsernames(ctx context.Context, base string) ([]string, error) {
var existingUsernames []string
if err := d.DB(ctx).Table("w_users").
Where("username = ? OR username LIKE ? ESCAPE '\\'", base, util.EscapeLike(base)+"-%").
Pluck("username", &existingUsernames).Error; err != nil {
return nil, err
}
return existingUsernames, nil
}
// GetSystemConfigValue 读取系统配置项原始值
func (d *DAO) GetSystemConfigValue(ctx context.Context, key string) (string, error) {
var val string
if err := d.DB(ctx).Table("w_system_configs").Where("key = ?", key).Pluck("value", &val).Error; err != nil {
return "", err
}
return val, nil
}
// ListSystemConfigsByKeys 按键批量读取系统配置项
func (d *DAO) ListSystemConfigsByKeys(ctx context.Context, keys []string) ([]do.CapConfigRecord, error) {
var records []do.CapConfigRecord
db := d.DB(ctx)
if db == nil {
return nil, nil
}
if err := db.Table("w_system_configs").Where("key IN ?", keys).Find(&records).Error; err != nil {
return nil, err
}
return records, nil
}
-52
View File
@@ -1,52 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// OAuth and Auth error messages
const (
errInvalidState = "非法登录请求"
errIDTokenVerifyFailed = "ID Token 验证失败" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errIDTokenVerifyFailedFormat = "%s: %w"
errNonceMismatch = "nonce 不匹配,可能存在重放攻击"
errNoActiveAuthSource = "未配置可用认证源"
errServerAddressMissing = "服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试"
errAuthSourceRequired = "认证源不能为空"
errDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
errUsernameGenerateFailed = "无法生成可用用户名"
errUsernameFromSourceFailed = "无法从认证源获取用户名"
errAuthSourceDisabled = "认证源未启用"
errInvalidExternalAccountBindingID = "绑定记录 ID 无效"
ErrTokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errOAuthStateRateLimited = "请求授权过于频繁,请稍后重试"
errAuthSourceNameRequired = "认证源名称不能为空"
errAuthSourceNameInvalid = "认证源名称格式不正确"
errAuthSourceTypeUnsupported = "不支持的认证源类型"
errAuthSourceDiscoveryURLRequired = "Discovery URL 不能为空"
//nolint:gosec // error message, not hardcoded credentials
errAuthSourceClientCredentialsRequired = "启用认证源时必须配置 Client ID 和 Client Secret"
errAuthSourceIDRequired = "认证源 ID 不能为空"
errUserIDRequired = "用户 ID 不能为空"
errExternalAccountBindingIncomplete = "外部帐号绑定信息不完整"
errExternalAccountAlreadyBoundToAnother = "该外部帐号已被其他用户绑定"
errExternalAccountBindingIDRequired = "外部帐号绑定记录 ID 不能为空"
errInsufficientPermission = "权限不足"
errBannedAccount = "账号已被封禁"
errUnAuthorized = "未登录"
)
// Service 层与鉴权中间件内部错误文案(保持与重构前逐字一致)
const (
errUserNotInContext = "auth: user not found in context"
errEmptyToken = "auth: empty token" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errSystemUserTokenNotAllowed = "auth: system user token not allowed" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errUnauthorizedInternal = "unauthorized"
errSystemUserLoginNotAllowed = "system user is not allowed to login"
)
// OAuth 回调会话校验错误文案(保持与重构前逐字一致)
const (
errInvalidSessionContext = "invalid session context"
errSessionMismatchForOAuth = "session mismatch for oauth state"
errUserContextMismatch = "user context mismatch for oauth binding"
)
+338
View File
@@ -0,0 +1,338 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"Wavelet/plugins/domain/auth/consts"
"Wavelet/plugins/domain/auth/controller"
"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"
"net/http"
"sync"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
)
// Exported Type Aliases for Backward Compatibility
//
//nolint:revive // backward compatibility type aliases with legacy naming
type (
AuthSource = entity.AuthSource
ExternalAccount = entity.ExternalAccount
CachedToken = do.CachedToken
CapRuntimeSettings = do.CapRuntimeSettings
AuthSourceView = dto.AuthSourceView
BasicUserInfo = dto.BasicUserInfo
OAuthAuthorizeResponse = dto.OAuthAuthorizeResponse
OAuthCallbackResult = dto.OAuthCallbackResult
CallbackRequest = dto.CallbackRequest
ChallengeResponse = dto.ChallengeResponse
RedeemResponse = dto.RedeemResponse
CaptchaManager = service.CaptchaManager
)
// Exported Constant Aliases for Backward Compatibility
const (
UserNameKey = consts.UserNameKey
UserIDKey = consts.UserIDKey
UserObjKey = consts.UserObjKey
TokenAuthKey = consts.TokenAuthKey
TokenAdminKey = consts.TokenAdminKey
SessionTokenKey = consts.SessionTokenKey
PasswordHashKey = consts.PasswordHashKey
SystemUsername = consts.SystemUsername
OAuthStateCacheKeyFormat = consts.OAuthStateCacheKeyFormat
OAuthStateCacheKeyExpiration = consts.OAuthStateCacheKeyExpiration
OAuthPurposeLogin = consts.OAuthPurposeLogin
OAuthPurposeBind = consts.OAuthPurposeBind
AuthSourceTypeOIDC = consts.AuthSourceTypeOIDC
ErrTokenAuthNotAllowed = consts.ErrTokenAuthNotAllowed
)
var (
defaultMu sync.RWMutex
defaultDAO = dao.New(nil, nil, nil)
defaultService = service.New(defaultDAO, SessionConfig{SessionCookieName: "wavelet_session", SessionAge: 86400, SessionHTTPOnly: true}, nil)
defaultCtrl = controller.New(defaultService)
)
func setDefaultRuntime(d *dao.DAO, s *service.Service, c *controller.Controller) {
defaultMu.Lock()
defer defaultMu.Unlock()
defaultDAO = d
defaultService = s
defaultCtrl = c
}
func getDefaultRuntime() (*dao.DAO, *service.Service, *controller.Controller) {
defaultMu.RLock()
defer defaultMu.RUnlock()
return defaultDAO, defaultService, defaultCtrl
}
// ParseUserID parses a string, int, or float64 user ID representation.
func ParseUserID(v any) uint64 {
return dto.ParseUserID(v)
}
// BuildBasicUserInfo converts UserDTO to BasicUserInfo.
func BuildBasicUserInfo(user *contracts.UserDTO, needChange bool) BasicUserInfo {
return dto.BuildBasicUserInfo(user, needChange)
}
// SetSessionConfig updates the active session configuration.
func SetSessionConfig(cfg SessionConfig) {
_, s, _ := getDefaultRuntime()
s.Session.SetConfig(cfg)
}
// GetSessionConfig returns the active session configuration.
func GetSessionConfig() SessionConfig {
_, s, _ := getDefaultRuntime()
return s.Session.Config()
}
// GetSessionOptions builds session cookie options based on config and maxAge.
func GetSessionOptions(maxAge int) sessions.Options {
_, s, _ := getDefaultRuntime()
return s.Session.GetSessionOptions(maxAge)
}
// StripCookieMaxAgeAndExpires removes max-age and expires from cookie header.
func StripCookieMaxAgeAndExpires(header http.Header, cookieName string) {
_, s, _ := getDefaultRuntime()
s.Session.StripCookieMaxAgeAndExpires(header, cookieName)
}
// GetUserIDFromSession extracts user ID from session.
func GetUserIDFromSession(s sessions.Session) uint64 {
return controller.GetUserIDFromSession(s)
}
// GetUserIDFromContext extracts user ID from Gin context.
func GetUserIDFromContext(c *gin.Context) uint64 {
return controller.GetUserIDFromContext(c)
}
// SetLoginSession sets the login session for the authenticated user.
func SetLoginSession(ctx context.Context, c *gin.Context, user *contracts.UserDTO, extras ...map[string]any) error {
_, s, _ := getDefaultRuntime()
session := sessions.Default(c)
isSessionCookie, err := s.Session.ApplyLoginSession(ctx, session, user, extras...)
if err != nil {
return err
}
if isSessionCookie {
s.Session.StripCookieMaxAgeAndExpires(c.Writer.Header(), s.Session.Config().SessionCookieName)
}
return nil
}
// RegisterWhitelist adds whitelist path patterns.
func RegisterWhitelist(patterns ...string) {
_, _, c := getDefaultRuntime()
c.RegisterWhitelist(patterns...)
}
// IsWhitelisted checks if the path matches the auth whitelist.
func IsWhitelisted(path string) bool {
_, _, c := getDefaultRuntime()
if wl := c.Whitelist(); wl != nil {
return wl.Match(path)
}
return false
}
// GetUserFromRequest extracts user from Request (Token or Session).
func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
d, _, _ := getDefaultRuntime()
return controller.GetUserFromRequest(c, d)
}
// LoginRequired returns authentication required middleware.
func LoginRequired() gin.HandlerFunc {
_, _, c := getDefaultRuntime()
return c.LoginRequired()
}
// AdminRequired returns admin authorization middleware.
func AdminRequired() gin.HandlerFunc {
_, _, c := getDefaultRuntime()
return c.AdminRequired()
}
// LoginAdminRequired alias for AdminRequired.
func LoginAdminRequired() gin.HandlerFunc {
return AdminRequired()
}
// DisallowTokenAuth returns middleware rejecting access token requests.
func DisallowTokenAuth() gin.HandlerFunc {
return controller.DisallowTokenAuth()
}
// GetCachedToken reads cached access token.
func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error) {
d, _, _ := getDefaultRuntime()
return d.GetCachedToken(ctx, tokenHash)
}
// SetCachedToken stores access token into cache.
func SetCachedToken(ctx context.Context, tokenHash string, token *CachedToken) {
d, _, _ := getDefaultRuntime()
d.SetCachedToken(ctx, tokenHash, token)
}
// InvalidateCachedToken invalidates access token cache.
func InvalidateCachedToken(ctx context.Context, tokenHash string) {
d, _, _ := getDefaultRuntime()
d.InvalidateCachedToken(ctx, tokenHash)
}
// GetCachedUser reads cached user.
func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
d, _, _ := getDefaultRuntime()
return d.GetCachedUser(ctx, userID)
}
// SetCachedUser stores user into cache.
func SetCachedUser(ctx context.Context, userID uint64, u *contracts.UserDTO) {
d, _, _ := getDefaultRuntime()
d.SetCachedUser(ctx, userID, u)
}
// InvalidateCachedUser invalidates user cache.
func InvalidateCachedUser(ctx context.Context, userID uint64) {
d, _, _ := getDefaultRuntime()
d.InvalidateCachedUser(ctx, userID)
}
// ResetAuthRAMCacheForTest clears RAM caches.
func ResetAuthRAMCacheForTest() {
dao.ResetRAMCacheForTest()
}
// StopAuthCacheListener compatibility stub.
func StopAuthCacheListener() {}
// SetCapSecret sets CAPTCHA secret.
func SetCapSecret(secret []byte) {
_, s, _ := getDefaultRuntime()
s.CapManager.SetSecret(secret)
}
// GetDefaultCapManager returns the singleton CAPTCHA manager.
func GetDefaultCapManager() *CaptchaManager {
_, s, _ := getDefaultRuntime()
return s.CapManager
}
// CurrentCapSettings returns current CAPTCHA runtime settings.
func CurrentCapSettings(ctx context.Context) (CapRuntimeSettings, error) {
_, s, _ := getDefaultRuntime()
return s.CapSettings.Current(ctx)
}
// CapProtectionEnabled checks if CAPTCHA is enabled.
func CapProtectionEnabled(ctx context.Context) bool {
_, s, _ := getDefaultRuntime()
return s.CapSettings.CapProtectionEnabled(ctx)
}
// InvalidateCapRuntimeSettings invalidates runtime CAPTCHA settings cache.
func InvalidateCapRuntimeSettings() {
_, s, _ := getDefaultRuntime()
s.CapSettings.Invalidate()
}
// ResetCapRuntimeSettingsForTest clears test CAPTCHA settings.
func ResetCapRuntimeSettingsForTest() {
InvalidateCapRuntimeSettings()
}
// InstallCapTestRuntimeSettings installs a test snapshot.
func InstallCapTestRuntimeSettings(settings CapRuntimeSettings) func() {
_, s, _ := getDefaultRuntime()
return s.CapSettings.InstallTestSnapshot(settings)
}
// VerifyCaptchaMiddleware returns captcha verification middleware.
func VerifyCaptchaMiddleware(mgr *service.CaptchaManager, scope string) gin.HandlerFunc {
_, s, _ := getDefaultRuntime()
return controller.VerifyCaptchaMiddleware(mgr, s.CapSettings, scope)
}
// Challenge HTTP handler.
func Challenge(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.Captcha.Challenge(c)
}
// Redeem HTTP handler.
func Redeem(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.Captcha.Redeem(c)
}
// GetLoginSources HTTP handler.
func GetLoginSources(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.OAuth.GetLoginSources(c)
}
// GetLoginURL HTTP handler.
func GetLoginURL(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.OAuth.GetLoginURL(c)
}
// Authorize HTTP handler.
func Authorize(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.OAuth.Authorize(c)
}
// Callback HTTP handler.
func Callback(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.OAuth.Callback(c)
}
// Logout HTTP handler.
func Logout(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.OAuth.Logout(c)
}
// UserInfo HTTP handler.
func UserInfo(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.UserInfo.UserInfo(c)
}
// ListExternalAccounts HTTP handler.
func ListExternalAccounts(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.OAuth.ListExternalAccounts(c)
}
// DeleteExternalAccount HTTP handler.
func DeleteExternalAccount(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.OAuth.DeleteExternalAccount(c)
}
// InvalidateOIDCProviderCache invalidates OIDC provider cache entry.
func InvalidateOIDCProviderCache(issuer string) {
_, s, _ := getDefaultRuntime()
s.OIDCProviderCache.Invalidate(issuer)
}
-587
View File
@@ -1,587 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/idgen"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"context"
"errors"
"fmt"
"net/http"
"strconv"
"strings"
"time"
"github.com/coreos/go-oidc/v3/oidc"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"gorm.io/gorm"
)
// GetLoginSources 获取可用登录源列表
// @Summary 获取可用登录源
// @Description 返回当前系统已启用的所有 OAuth 登录源,前端展示登录按钮列表时调用
// @Tags oauth
// @Produce json
// @Success 200 {object} response.Any{data=[]auth.AuthSourceView} "登录源列表"
// @Router /api/v1/oauth/sources [get]
func GetLoginSources(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(activeLoginSources(c.Request.Context())))
}
// GetLoginURL 获取登录授权地址
// @Summary 获取登录授权地址
// @Description 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用第一个启用的认证源。
// @Tags oauth
// @Produce json
// @Param source query string false "认证源名称,为空使用第一个启用的认证源"
// @Success 200 {object} response.Any{data=auth.OAuthAuthorizeResponse} "授权 URL"
// @Failure 400 {object} response.Any "认证源不存在或未配置"
// @Failure 500 {object} response.Any "构造 URL 失败"
// @Router /api/v1/oauth/login [get]
func GetLoginURL(c *gin.Context) {
ctx := c.Request.Context()
if !isOIDCLoginEnabled(ctx) {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
source, err := resolveAuthSource(ctx, c.Query("source"))
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
session := sessions.Default(c)
token, isNew := ensureSessionToken(session)
if isNew {
if err := session.Save(); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
userID := GetUserIDFromSession(session)
sessionHash := hashSessionToken(token)
if err := reserveOAuthStateSlot(ctx, sessionHash); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
state := uuid.NewString()
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
SourceName: source.Name,
Purpose: OAuthPurposeLogin,
UserID: userID,
SessionHash: sessionHash,
})
if err != nil {
response.AbortInternal(c, err.Error())
return
}
stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, state)
if cache := getCache(ctx); cache != nil {
if err := cache.Set(ctx, stateKey, payloadValue, OAuthStateCacheKeyExpiration); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
}
func buildAuthorizeURL(ctx context.Context, source *AuthSource, state string) (string, error) {
redirectURL, err := getFrontendLoginRedirectURL(ctx)
if err != nil {
return "", err
}
authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
if err != nil {
return "", err
}
if verifier != nil {
return authConfig.AuthCodeURL(state, oidc.Nonce(state)), nil
}
return authConfig.AuthCodeURL(state), nil
}
func reserveOAuthStateSlot(ctx context.Context, sessionHash string) error {
if sessionHash == "" {
return nil
}
if limiter := getLimiter(ctx); limiter != nil {
key := fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash)
res, err := limiter.Allow(ctx, key, contracts.Rate{
Limit: oauthStateLimitMax,
Period: OAuthStateCacheKeyExpiration,
})
if err != nil {
return err
}
if !res.Allowed {
return errors.New(errOAuthStateRateLimited)
}
return nil
}
cache := getCache(ctx)
if cache == nil {
return nil
}
key := fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash)
var count int
_ = cache.Get(ctx, key, &count)
count++
_ = cache.Set(ctx, key, count, OAuthStateCacheKeyExpiration)
if count > oauthStateLimitMax {
return errors.New(errOAuthStateRateLimited)
}
return nil
}
// 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=auth.OAuthAuthorizeResponse} "授权 URL"
// @Failure 400 {object} response.Any "认证源不存在或未启用"
// @Failure 500 {object} response.Any "构造 URL 失败"
// @Router /api/v1/oauth/{source}/authorize [get]
func Authorize(c *gin.Context) {
ctx := c.Request.Context()
if !isOIDCLoginEnabled(ctx) {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
source, err := resolveAuthSource(ctx, c.Param("source"))
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
purpose := strings.ToLower(strings.TrimSpace(c.Query("purpose")))
if purpose != OAuthPurposeBind {
purpose = OAuthPurposeLogin
}
session := sessions.Default(c)
userID := GetUserIDFromSession(session)
if purpose == OAuthPurposeBind && userID == 0 {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
token, isNew := ensureSessionToken(session)
if isNew {
if err := session.Save(); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
sessionHash := hashSessionToken(token)
if err := reserveOAuthStateSlot(ctx, sessionHash); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
state := uuid.NewString()
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
SourceName: source.Name,
Purpose: purpose,
UserID: userID,
SessionHash: sessionHash,
})
if err != nil {
response.AbortInternal(c, err.Error())
return
}
stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, state)
if cache := getCache(ctx); cache != nil {
if err := cache.Set(ctx, stateKey, payloadValue, OAuthStateCacheKeyExpiration); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
}
// Callback OAuth 回调处理
// @Summary OAuth 回调处理
// @Description 接收前端传回的 state 和 code,完成 OAuth/OIDC 认证并建立会话。支持登录(login)和账号绑定(bind)两种场景。
// @Tags oauth
// @Accept json
// @Produce json
// @Param request body auth.CallbackRequest true "回调请求参数"
// @Success 200 {object} response.Any{data=auth.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 Callback(c *gin.Context) {
var req CallbackRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
ctx := c.Request.Context()
stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, req.State)
var payloadRaw string
cache := getCache(ctx)
if cache == nil {
response.AbortBadRequest(c, errInvalidState)
return
}
if err := cache.Get(ctx, stateKey, &payloadRaw); err != nil {
response.AbortBadRequest(c, errInvalidState)
return
}
_ = cache.Delete(ctx, stateKey)
payload, err := decodeOAuthStatePayload(payloadRaw)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
session := sessions.Default(c)
currentUserID := GetUserIDFromSession(session)
if payload.Purpose == OAuthPurposeBind && currentUserID == 0 {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
token, ok := session.Get(SessionTokenKey).(string)
if !ok || token == "" {
response.AbortBadRequest(c, errInvalidSessionContext)
return
}
if hashSessionToken(token) != payload.SessionHash {
response.AbortBadRequest(c, errSessionMismatchForOAuth)
return
}
if payload.Purpose == OAuthPurposeBind && currentUserID != payload.UserID {
response.AbortBadRequest(c, errUserContextMismatch)
return
}
if !isOIDCLoginEnabled(ctx) {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
source, err := resolveAuthSource(ctx, payload.SourceName)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
redirectURL, err := getFrontendLoginRedirectURL(ctx)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
userInfo, err := buildOAuthUserInfo(ctx, source, req.Code, req.State, redirectURL)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := normalizeOAuthUserInfo(userInfo); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if userInfo.Sub == "" {
userInfo.Sub = userInfo.Username
}
if payload.Purpose == OAuthPurposeBind {
handleCallbackBind(ctx, c, source, userInfo)
return
}
handleCallbackLogin(ctx, c, source, userInfo)
}
func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) {
userID := GetUserIDFromContext(c)
if userID == 0 {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
user, err := GetUserByID(ctx, userID)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := BindExternalAccount(ctx, &ExternalAccount{
AuthSourceID: source.ID,
UserID: user.ID,
ExternalID: userInfo.Sub,
ExternalUsername: userInfo.Username,
Email: userInfo.Email,
}); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
user.LastLoginAt = time.Now()
_ = TouchUserLastLogin(ctx, user.ID, user.LastLoginAt)
c.JSON(http.StatusOK, response.OK(buildCallbackResult(user, "bound")))
}
func handleCallbackLogin(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) {
var user *contracts.UserDTO
account, err := FindExternalAccount(ctx, source.ID, userInfo.Sub)
switch {
case err == nil:
loaded, loadErr := GetUserByID(ctx, account.UserID)
if loadErr != nil {
response.AbortInternal(c, loadErr.Error())
return
}
user = loaded
case errors.Is(err, gorm.ErrRecordNotFound):
newUser, ok := handleCallbackRegister(ctx, c, source, userInfo)
if !ok {
return
}
user = &newUser
default:
response.AbortInternal(c, err.Error())
return
}
user.LastLoginAt = time.Now()
_ = TouchUserLastLogin(ctx, user.ID, user.LastLoginAt)
if err := SetLoginSession(ctx, c, user); err != nil {
response.AbortInternal(c, err.Error())
return
}
SetCachedUser(ctx, user.ID, user)
c.JSON(http.StatusOK, response.OK(buildCallbackResult(user, "logged_in")))
}
func uniqueUsername(ctx context.Context, base string) (string, error) {
base = strings.TrimSpace(base)
if base == "" {
base = "user"
}
existingUsernames, err := ListSimilarUsernames(ctx, base)
if err != nil {
return "", err
}
exists := make(map[string]bool, len(existingUsernames))
for _, u := range existingUsernames {
exists[strings.ToLower(u)] = true
}
if !exists[strings.ToLower(base)] {
return base, nil
}
for i := 1; i <= 1000; i++ {
candidate := fmt.Sprintf("%s-%d", base, i)
if !exists[strings.ToLower(candidate)] {
return candidate, nil
}
}
return "", errors.New(errUsernameGenerateFailed)
}
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) (contracts.UserDTO, bool) {
registrationEnabled := true
val, cfgErr := GetSystemConfigValue(ctx, "registration_enabled")
if cfgErr == nil && val != "" {
if b, err := strconv.ParseBool(val); err == nil {
registrationEnabled = b
}
}
if !registrationEnabled {
c.JSON(http.StatusOK, response.OK(buildCallbackResult(nil, "need_bind")))
return contracts.UserDTO{}, false
}
username, uniqueErr := uniqueUsername(ctx, userInfo.Username)
if uniqueErr != nil {
response.AbortInternal(c, uniqueErr.Error())
return contracts.UserDTO{}, false
}
userInfo.Username = username
now := time.Now()
user := contracts.UserDTO{
ID: idgen.NextUint64ID(),
Username: userInfo.Username,
Nickname: userInfo.Name,
Email: userInfo.Email,
AvatarURL: userInfo.AvatarURL,
IsActive: userInfo.Active,
LastLoginAt: now,
CreatedAt: now,
UpdatedAt: now,
}
if err := InsertUser(ctx, &user); err != nil {
response.AbortInternal(c, err.Error())
return contracts.UserDTO{}, false
}
if err := BindExternalAccount(ctx, &ExternalAccount{
AuthSourceID: source.ID,
UserID: user.ID,
ExternalID: userInfo.Sub,
ExternalUsername: userInfo.Username,
Email: userInfo.Email,
}); err != nil {
response.AbortBadRequest(c, err.Error())
return contracts.UserDTO{}, false
}
logger.InfoF(ctx, "[LoginAudit] successful OAuth registration via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
return user, true
}
// UserInfo 获取当前登录用户信息
// @Summary 获取当前登录用户信息
// @Description 返回当前登录用户的基本信息,需要登录。
// @Tags oauth
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=auth.BasicUserInfo} "用户信息"
// @Failure 401 {object} response.Any "未登录"
// @Router /api/v1/oauth/user-info [get]
// @Router /api/v1/user-info [get]
func 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(BuildBasicUserInfo(user, needChange)),
)
}
// 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 Logout(c *gin.Context) {
session := sessions.Default(c)
userID := session.Get(UserIDKey)
username := session.Get(UserNameKey)
if userID != nil {
logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP())
if id := ParseUserID(userID); id > 0 {
InvalidateCachedUser(c.Request.Context(), id)
}
}
session.Options(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 ListExternalAccounts(c *gin.Context) {
userID := GetUserIDFromContext(c)
accounts, err := ListExternalAccountsByUserID(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 DeleteExternalAccount(c *gin.Context) {
userID := GetUserIDFromContext(c)
if userID == 0 {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
rawID := strings.TrimSpace(c.Param("id"))
id, err := strconv.ParseUint(rawID, 10, 64)
if err != nil || id == 0 {
response.AbortBadRequest(c, errInvalidExternalAccountBindingID)
return
}
if err := UnbindExternalAccount(c.Request.Context(), id, userID); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
-195
View File
@@ -1,195 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/response"
"Wavelet/pkg/trace"
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"github.com/gin-gonic/gin"
)
// whitelist holds the no-auth route patterns. They are registered during Apply and
// matched on every request, so PathWhitelist parses them once up front.
var whitelist = extpoints.NewPathWhitelist()
// RegisterWhitelist registers route patterns that bypass mandatory authentication.
func RegisterWhitelist(patterns ...string) {
whitelist.Add(patterns...)
}
// IsWhitelisted checks if the specified path matches the auth whitelist.
func IsWhitelisted(path string) bool {
return whitelist.Match(path)
}
func hashToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
// currentUserIDFromRequestContext 是接入层向 Service 层暴露的登录态桥接。
//
// Session 读取必须依赖 *gin.Context,而 Service 层禁止 import gin,
// 因此该类型断言收敛在本(接入层)文件中。ok 为 false 表示 ctx 不是 *gin.Context。
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, tokenStr string) (*contracts.UserDTO, *CachedToken, error) {
tokenHash := hashToken(tokenStr)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil || tokenRecord == nil {
tokenRecord, err = GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, nil, err
}
SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || user == nil || !user.IsActive {
user, err = GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, nil, err
}
SetCachedUser(ctx, tokenRecord.UserID, user)
}
return user, tokenRecord, nil
}
// GetUserFromRequest 从请求中获取当前用户(优先 Access Token,其次 Session)
func GetUserFromRequest(c *gin.Context) (*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, tokenStr); err == nil {
if user.Username == SystemUsername {
return nil, errors.New(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(errUnauthorizedInternal)
}
user, err := GetCachedUser(ctx, userID)
if err != nil || user == nil || !user.IsActive {
user, err = GetActiveUserByID(ctx, userID)
if err != nil {
return nil, err
}
SetCachedUser(ctx, userID, user)
}
ginutil.SetToContext(c, contracts.AuthTokenAuthKey, false)
ginutil.SetToContext(c, contracts.AuthTokenAdminKey, false)
if user.Username == "system" {
return nil, errors.New(errSystemUserLoginNotAllowed)
}
return user, nil
}
// LoginRequired 返回登录鉴权中间件,校验 Access Token 或 Session
func LoginRequired() gin.HandlerFunc {
return func(c *gin.Context) {
if IsWhitelisted(c.Request.URL.Path) {
c.Next()
return
}
_, span := trace.Start(c.Request.Context(), "LoginRequired")
defer span.End()
user, err := GetUserFromRequest(c)
if err != nil {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
LogForAudit(c.Request.Context(), user, c)
ginutil.SetToContext(c, contracts.AuthUserObjKey, user)
c.Next()
}
}
// AdminRequired 校验管理员权限(支持 Session 和 Token 鉴权)
func AdminRequired() gin.HandlerFunc {
return func(c *gin.Context) {
_, span := trace.Start(c.Request.Context(), "AdminRequired")
defer span.End()
user, err := GetUserFromRequest(c)
if err != nil {
response.AbortUnauthorized(c, 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, errInsufficientPermission)
return
}
if !isTokenAuth && !user.IsAdmin {
response.AbortForbidden(c, errInsufficientPermission)
return
}
LogForAudit(c.Request.Context(), user, c)
ginutil.SetToContext(c, contracts.AuthUserObjKey, user)
c.Next()
}
}
// LoginAdminRequired is an alias for AdminRequired.
func LoginAdminRequired() gin.HandlerFunc {
return AdminRequired()
}
// DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点
func DisallowTokenAuth() gin.HandlerFunc {
return func(c *gin.Context) {
if tokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
response.AbortForbidden(c, ErrTokenAuthNotAllowed)
return
}
c.Next()
}
}
@@ -0,0 +1,12 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package do provides domain data objects for the auth plugin.
package do
// CachedToken represents the minimal cached representation of an access token.
type CachedToken struct {
ID uint64 `json:"id"`
UserID uint64 `json:"user_id"`
IsAdmin bool `json:"is_admin"`
}
@@ -0,0 +1,75 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package do provides domain data objects for the auth plugin.
package do
import (
"Wavelet/plugins/domain/auth/consts"
"strconv"
"time"
)
// CapRuntimeSettings is the parsed CAPTCHA runtime configuration loaded from system_configs.
type CapRuntimeSettings struct {
LoginEnabled bool
ChallengeCount int
ChallengeSize int
ChallengeDifficulty int
ChallengeTTL time.Duration
TokenTTL time.Duration
}
// CapConfigRecord maps the columns selected from the system config table.
type CapConfigRecord struct {
Key string `gorm:"column:key"`
Value string `gorm:"column:value"`
}
// ParseCapRuntimeSettings parses system config key-value map into CapRuntimeSettings with fallback defaults.
func ParseCapRuntimeSettings(configs map[string]string) CapRuntimeSettings {
settings := CapRuntimeSettings{
ChallengeCount: consts.DefaultCapChallengeCount,
ChallengeSize: consts.DefaultCapChallengeSize,
ChallengeDifficulty: consts.DefaultCapChallengeDifficulty,
ChallengeTTL: consts.DefaultCapChallengeTTL,
TokenTTL: consts.DefaultCapTokenTTL,
}
if len(configs) == 0 {
return settings
}
if val, ok := configs[consts.ConfigKeyCapLoginEnabled]; ok {
if enabled, err := strconv.ParseBool(val); err == nil {
settings.LoginEnabled = enabled
}
}
if val, ok := configs[consts.ConfigKeyCapChallengeCount]; ok {
if count, err := strconv.Atoi(val); err == nil && count > 0 {
settings.ChallengeCount = count
}
}
if val, ok := configs[consts.ConfigKeyCapChallengeSize]; ok {
if size, err := strconv.Atoi(val); err == nil && size > 0 {
settings.ChallengeSize = size
}
}
if val, ok := configs[consts.ConfigKeyCapChallengeDifficulty]; ok {
if diff, err := strconv.Atoi(val); err == nil && diff > 0 {
settings.ChallengeDifficulty = diff
}
}
if val, ok := configs[consts.ConfigKeyCapChallengeTTL]; ok {
if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 {
settings.ChallengeTTL = time.Duration(ttlSeconds) * time.Second
}
}
if val, ok := configs[consts.ConfigKeyCapTokenTTL]; ok {
if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 {
settings.TokenTTL = time.Duration(ttlSeconds) * time.Second
}
}
return settings
}
@@ -0,0 +1,43 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package do_test
import (
"Wavelet/plugins/domain/auth/consts"
"Wavelet/plugins/domain/auth/model/do"
"testing"
"time"
"github.com/stretchr/testify/assert"
)
func TestParseCapRuntimeSettings(t *testing.T) {
t.Run("Default fallback on empty config", func(t *testing.T) {
settings := do.ParseCapRuntimeSettings(nil)
assert.False(t, settings.LoginEnabled)
assert.Equal(t, consts.DefaultCapChallengeCount, settings.ChallengeCount)
assert.Equal(t, consts.DefaultCapChallengeSize, settings.ChallengeSize)
assert.Equal(t, consts.DefaultCapChallengeDifficulty, settings.ChallengeDifficulty)
assert.Equal(t, consts.DefaultCapChallengeTTL, settings.ChallengeTTL)
assert.Equal(t, consts.DefaultCapTokenTTL, settings.TokenTTL)
})
t.Run("Parsed custom configs", func(t *testing.T) {
configs := map[string]string{
consts.ConfigKeyCapLoginEnabled: "true",
consts.ConfigKeyCapChallengeCount: "3",
consts.ConfigKeyCapChallengeSize: "64",
consts.ConfigKeyCapChallengeDifficulty: "5",
consts.ConfigKeyCapChallengeTTL: "300",
consts.ConfigKeyCapTokenTTL: "600",
}
settings := do.ParseCapRuntimeSettings(configs)
assert.True(t, settings.LoginEnabled)
assert.Equal(t, 3, settings.ChallengeCount)
assert.Equal(t, 64, settings.ChallengeSize)
assert.Equal(t, 5, settings.ChallengeDifficulty)
assert.Equal(t, 300*time.Second, settings.ChallengeTTL)
assert.Equal(t, 600*time.Second, settings.TokenTTL)
})
}
@@ -0,0 +1,33 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package do provides domain data objects for the auth plugin.
package do
import "encoding/json"
// OAuthStatePayload represents the cached state verification payload for OAuth flow.
type OAuthStatePayload struct {
SourceName string `json:"source_name"`
Purpose string `json:"purpose"`
UserID uint64 `json:"user_id,omitempty"`
SessionHash string `json:"session_hash"`
}
// Encode converts OAuthStatePayload to a JSON string.
func (p OAuthStatePayload) Encode() (string, error) {
data, err := json.Marshal(p)
if err != nil {
return "", err
}
return string(data), nil
}
// DecodeOAuthStatePayload parses a JSON string into OAuthStatePayload.
func DecodeOAuthStatePayload(value string) (OAuthStatePayload, error) {
var payload OAuthStatePayload
if err := json.Unmarshal([]byte(value), &payload); err != nil {
return OAuthStatePayload{}, err
}
return payload, nil
}
@@ -0,0 +1,32 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package do_test
import (
"Wavelet/plugins/domain/auth/model/do"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestOAuthStatePayload(t *testing.T) {
payload := do.OAuthStatePayload{
SourceName: "github",
Purpose: "login",
UserID: 12345,
SessionHash: "hash-abc-123",
}
encoded, err := payload.Encode()
require.NoError(t, err)
assert.NotEmpty(t, encoded)
decoded, err := do.DecodeOAuthStatePayload(encoded)
require.NoError(t, err)
assert.Equal(t, payload, decoded)
_, err = do.DecodeOAuthStatePayload("invalid-json")
assert.Error(t, err)
}
@@ -0,0 +1,45 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package dto provides data transfer objects and views for the auth plugin.
package dto
// AuthSourceView 登录源展示信息
type AuthSourceView struct {
ID uint64 `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
DisplayName string `json:"display_name"`
IsActive bool `json:"is_active"`
IconURL string `json:"icon_url"`
ClientSecretConfigured bool `json:"client_secret_configured"`
}
// OAuthAuthorizeResponse 授权 URL 响应
type OAuthAuthorizeResponse struct {
AuthorizeURL string `json:"authorize_url"`
}
// OAuthCallbackResult 回调处理结果
type OAuthCallbackResult struct {
Status string `json:"status"`
User *BasicUserInfo `json:"user,omitempty"`
}
// CallbackRequest OAuth 回调请求参数
type CallbackRequest struct {
State string `json:"state" binding:"required"`
Code string `json:"code" binding:"required"`
}
// ExternalAccountView 外部帐号绑定视图(脱敏展示用)
type ExternalAccountView struct {
ID uint64 `json:"id"`
AuthSourceID uint64 `json:"auth_source_id"`
AuthSourceName string `json:"auth_source_name"`
AuthSourceType string `json:"auth_source_type"`
AuthSourceLabel string `json:"auth_source_label"`
ExternalUsername string `json:"external_username"`
Email string `json:"email"`
CreatedAt string `json:"created_at"`
}
@@ -1,7 +1,8 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// Package dto provides data transfer objects and views for the auth plugin.
package dto
import (
"Wavelet/plugins/domain/auth/pow"
@@ -10,13 +11,13 @@ import (
// ChallengeResponse is a local type alias for the pow.ChallengeResponse struct
type ChallengeResponse = pow.ChallengeResponse
// challengeRequest is the CAPTCHA challenge request payload.
type challengeRequest struct {
// ChallengeRequest is the CAPTCHA challenge request payload.
type ChallengeRequest struct {
Scope string `json:"scope" form:"scope"`
}
// redeemRequest is the CAPTCHA redeem request payload.
type redeemRequest struct {
// RedeemRequest is the CAPTCHA redeem request payload.
type RedeemRequest struct {
Token string `json:"token" binding:"required"`
Solutions []int `json:"solutions" binding:"required"`
Scope string `json:"scope" form:"scope"`
@@ -29,9 +30,3 @@ type RedeemResponse struct {
Expires int64 `json:"expires,omitempty"`
Error string `json:"error,omitempty"`
}
// capConfigRecord maps the columns selected from the system config table.
type capConfigRecord struct {
Key string `gorm:"column:key"`
Value string `gorm:"column:value"`
}
@@ -0,0 +1,84 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package dto provides data transfer objects and views for the auth plugin.
package dto
import (
"Wavelet/core/contracts"
"strconv"
)
// BasicUserInfo 用户基本信息结构体
type BasicUserInfo struct {
ID uint64 `json:"id,string"`
Username string `json:"username"`
Nickname string `json:"nickname"`
Email string `json:"email"`
AvatarURL string `json:"avatar_url"`
IsAdmin bool `json:"is_admin"`
NeedChangePassword bool `json:"need_change_password"`
Bio string `json:"bio"`
Phone string `json:"phone"`
Gender string `json:"gender"`
Website string `json:"website"`
Location string `json:"location"`
}
// BuildBasicUserInfo 将 UserDTO 转换为 BasicUserInfo
func BuildBasicUserInfo(user *contracts.UserDTO, needChange bool) BasicUserInfo {
if user == nil {
return BasicUserInfo{}
}
return BasicUserInfo{
ID: user.ID,
Username: user.Username,
Nickname: user.Nickname,
Email: user.Email,
AvatarURL: user.AvatarURL,
IsAdmin: user.IsAdmin,
NeedChangePassword: needChange || user.NeedChangePassword,
Bio: user.Bio,
Phone: user.Phone,
Gender: user.Gender,
Website: user.Website,
Location: user.Location,
}
}
// LoginRequiredAuditLog 审计日志结构体
type LoginRequiredAuditLog struct {
UserID uint64 `json:"user_id"`
Username string `json:"username"`
ClientIP string `json:"client_ip"`
Method string `json:"method"`
Path string `json:"path"`
RequestURI string `json:"request_uri"`
UserAgent string `json:"user_agent"`
Referer string `json:"referer"`
}
// ParseUserID parses a string, int, or float64 user ID representation.
func ParseUserID(v any) uint64 {
switch val := v.(type) {
case uint64:
return val
case int64:
if val > 0 {
return uint64(val)
}
case int:
if val > 0 {
return uint64(val)
}
case float64:
if val > 0 {
return uint64(val)
}
case string:
if id, err := strconv.ParseUint(val, 10, 64); err == nil {
return id
}
}
return 0
}
@@ -0,0 +1,83 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package entity provides database model entities for the auth domain plugin.
package entity
import (
"Wavelet/plugins/domain/auth/consts"
"errors"
"regexp"
"strings"
"time"
)
var authSourceNamePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_-]{0,79}$`)
// AuthSource 认证源实体
type AuthSource struct {
ID uint64 `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"uniqueIndex;size:80;not null"`
Type string `json:"type" gorm:"size:20;not null"`
DisplayName string `json:"display_name" gorm:"size:100"`
IsActive bool `json:"is_active" gorm:"index;not null;default:false"`
ClientID string `json:"client_id" gorm:"size:255"`
ClientSecret string `json:"-" gorm:"size:1024"`
OpenIDDiscoveryURL string `json:"openid_discovery_url" gorm:"column:openid_discovery_url;size:1024"`
Scopes string `json:"scopes" gorm:"size:255"`
IconURL string `json:"icon_url" gorm:"size:1024"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
ClientSecretConfigured bool `json:"client_secret_configured" gorm:"-"`
}
// TableName 表名
func (AuthSource) TableName() string {
return "w_auth_sources"
}
// Normalize 对认证源字段进行标准化处理
func (source *AuthSource) Normalize() {
source.Type = strings.ToLower(strings.TrimSpace(source.Type))
source.Name = strings.TrimSpace(source.Name)
source.DisplayName = strings.TrimSpace(source.DisplayName)
source.ClientID = strings.TrimSpace(source.ClientID)
source.ClientSecret = strings.TrimSpace(source.ClientSecret)
source.OpenIDDiscoveryURL = strings.TrimSpace(source.OpenIDDiscoveryURL)
source.Scopes = strings.TrimSpace(source.Scopes)
source.IconURL = strings.TrimSpace(source.IconURL)
if source.DisplayName == "" {
source.DisplayName = source.Name
}
if source.Type == consts.AuthSourceTypeOIDC && source.Scopes == "" {
source.Scopes = "openid profile email"
}
}
// Validate 校验认证源字段合法性
func (source *AuthSource) Validate() error {
source.Normalize()
if source.Name == "" {
return errors.New(consts.ErrAuthSourceNameRequired)
}
if !authSourceNamePattern.MatchString(source.Name) {
return errors.New(consts.ErrAuthSourceNameInvalid)
}
if source.Type != consts.AuthSourceTypeOIDC {
return errors.New(consts.ErrAuthSourceTypeUnsupported)
}
if source.OpenIDDiscoveryURL == "" {
//nolint:staticcheck // descriptive error constant
return errors.New(consts.ErrAuthSourceDiscoveryURLRequired)
}
if source.IsActive && (source.ClientID == "" || source.ClientSecret == "") {
return errors.New(consts.ErrAuthSourceClientCredentialsRequired)
}
return nil
}
// Sanitize 脱敏处理,将 ClientSecret 清空并设置 ClientSecretConfigured 标志
func (source *AuthSource) Sanitize() {
source.ClientSecretConfigured = source.ClientSecret != ""
source.ClientSecret = ""
}
@@ -0,0 +1,75 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package entity_test
import (
"Wavelet/plugins/domain/auth/model/entity"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestAuthSourceValidation(t *testing.T) {
t.Run("Valid OIDC Source", func(t *testing.T) {
src := entity.AuthSource{
Name: "google",
Type: "oidc",
DisplayName: "Google Sign-In",
ClientID: "client-123",
ClientSecret: "secret-456",
OpenIDDiscoveryURL: "https://accounts.google.com",
IsActive: true,
}
require.NoError(t, src.Validate())
assert.Equal(t, "openid profile email", src.Scopes)
assert.Equal(t, "w_auth_sources", src.TableName())
src.Sanitize()
assert.True(t, src.ClientSecretConfigured)
assert.Empty(t, src.ClientSecret)
})
t.Run("Empty Name Fails", func(t *testing.T) {
src := entity.AuthSource{
Name: "",
Type: "oidc",
}
assert.Error(t, src.Validate())
})
t.Run("Invalid Name Format Fails", func(t *testing.T) {
src := entity.AuthSource{
Name: "invalid name with spaces!",
Type: "oidc",
}
assert.Error(t, src.Validate())
})
t.Run("Unsupported Type Fails", func(t *testing.T) {
src := entity.AuthSource{
Name: "ldap_source",
Type: "ldap",
}
assert.Error(t, src.Validate())
})
t.Run("Missing Discovery URL Fails", func(t *testing.T) {
src := entity.AuthSource{
Name: "google",
Type: "oidc",
}
assert.Error(t, src.Validate())
})
t.Run("Active Source Missing Credentials Fails", func(t *testing.T) {
src := entity.AuthSource{
Name: "google",
Type: "oidc",
OpenIDDiscoveryURL: "https://accounts.google.com",
IsActive: true,
}
assert.Error(t, src.Validate())
})
}
@@ -0,0 +1,26 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package entity provides database model entities for the auth domain plugin.
package entity
import (
"time"
)
// ExternalAccount 外部账号绑定实体
type ExternalAccount struct {
ID uint64 `json:"id" gorm:"primaryKey"`
AuthSourceID uint64 `json:"auth_source_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:1;index"`
UserID uint64 `json:"user_id" gorm:"index;not null"`
ExternalID string `json:"external_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:2;size:255;not null"`
ExternalUsername string `json:"external_username" gorm:"size:255"`
Email string `json:"email" gorm:"size:255"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// TableName 表名
func (ExternalAccount) TableName() string {
return "w_external_accounts"
}
-241
View File
@@ -1,241 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"encoding/json"
"errors"
"regexp"
"strconv"
"strings"
"time"
)
var authSourceNamePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_-]{0,79}$`)
// AuthSource 认证源实体
//
//nolint:revive // auth.AuthSource is standard domain entity name
type AuthSource struct {
ID uint64 `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"uniqueIndex;size:80;not null"`
Type string `json:"type" gorm:"size:20;not null"`
DisplayName string `json:"display_name" gorm:"size:100"`
IsActive bool `json:"is_active" gorm:"index;not null;default:false"`
ClientID string `json:"client_id" gorm:"size:255"`
ClientSecret string `json:"-" gorm:"size:1024"`
OpenIDDiscoveryURL string `json:"openid_discovery_url" gorm:"column:openid_discovery_url;size:1024"`
Scopes string `json:"scopes" gorm:"size:255"`
IconURL string `json:"icon_url" gorm:"size:1024"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
ClientSecretConfigured bool `json:"client_secret_configured" gorm:"-"`
}
// TableName 表名
func (AuthSource) TableName() string {
return "w_auth_sources"
}
// Normalize 对认证源字段进行标准化处理
func (source *AuthSource) Normalize() {
source.Type = strings.ToLower(strings.TrimSpace(source.Type))
source.Name = strings.TrimSpace(source.Name)
source.DisplayName = strings.TrimSpace(source.DisplayName)
source.ClientID = strings.TrimSpace(source.ClientID)
source.ClientSecret = strings.TrimSpace(source.ClientSecret)
source.OpenIDDiscoveryURL = strings.TrimSpace(source.OpenIDDiscoveryURL)
source.Scopes = strings.TrimSpace(source.Scopes)
source.IconURL = strings.TrimSpace(source.IconURL)
if source.DisplayName == "" {
source.DisplayName = source.Name
}
if source.Type == AuthSourceTypeOIDC && source.Scopes == "" {
source.Scopes = "openid profile email"
}
}
// Validate 校验认证源字段合法性
func (source *AuthSource) Validate() error {
source.Normalize()
if source.Name == "" {
return errors.New(errAuthSourceNameRequired)
}
if !authSourceNamePattern.MatchString(source.Name) {
return errors.New(errAuthSourceNameInvalid)
}
if source.Type != AuthSourceTypeOIDC {
return errors.New(errAuthSourceTypeUnsupported)
}
if source.OpenIDDiscoveryURL == "" {
//nolint:staticcheck // descriptive error constant
return errors.New(errAuthSourceDiscoveryURLRequired)
}
if source.IsActive && (source.ClientID == "" || source.ClientSecret == "") {
return errors.New(errAuthSourceClientCredentialsRequired)
}
return nil
}
// Sanitize 脱敏处理,将 ClientSecret 清空并设置 ClientSecretConfigured 标志
func (source *AuthSource) Sanitize() {
source.ClientSecretConfigured = source.ClientSecret != ""
source.ClientSecret = ""
}
// ExternalAccount 外部账号绑定实体
type ExternalAccount struct {
ID uint64 `json:"id" gorm:"primaryKey"`
AuthSourceID uint64 `json:"auth_source_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:1;index"`
UserID uint64 `json:"user_id" gorm:"index;not null"`
ExternalID string `json:"external_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:2;size:255;not null"`
ExternalUsername string `json:"external_username" gorm:"size:255"`
Email string `json:"email" gorm:"size:255"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// TableName 表名
func (ExternalAccount) TableName() string {
return "w_external_accounts"
}
// ExternalAccountView 外部帐号绑定视图(脱敏展示用)
type ExternalAccountView struct {
ID uint64 `json:"id"`
AuthSourceID uint64 `json:"auth_source_id"`
AuthSourceName string `json:"auth_source_name"`
AuthSourceType string `json:"auth_source_type"`
AuthSourceLabel string `json:"auth_source_label"`
ExternalUsername string `json:"external_username"`
Email string `json:"email"`
CreatedAt time.Time `json:"created_at"`
}
// AuthSourceView 登录源展示信息
//
//nolint:revive // auth.AuthSourceView is standard domain presentation struct
type AuthSourceView struct {
ID uint64 `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
DisplayName string `json:"display_name"`
IsActive bool `json:"is_active"`
IconURL string `json:"icon_url"`
ClientSecretConfigured bool `json:"client_secret_configured"`
}
// OAuthAuthorizeResponse 授权 URL 响应
type OAuthAuthorizeResponse struct {
AuthorizeURL string `json:"authorize_url"`
}
// OAuthCallbackResult 回调处理结果
type OAuthCallbackResult struct {
Status string `json:"status"`
User *BasicUserInfo `json:"user,omitempty"`
}
// CallbackRequest OAuth 回调请求参数
type CallbackRequest struct {
State string `json:"state" binding:"required"`
Code string `json:"code" binding:"required"`
}
// BasicUserInfo 用户基本信息结构体
type BasicUserInfo struct {
ID uint64 `json:"id,string"`
Username string `json:"username"`
Nickname string `json:"nickname"`
Email string `json:"email"`
AvatarURL string `json:"avatar_url"`
IsAdmin bool `json:"is_admin"`
NeedChangePassword bool `json:"need_change_password"`
Bio string `json:"bio"`
Phone string `json:"phone"`
Gender string `json:"gender"`
Website string `json:"website"`
Location string `json:"location"`
}
// BuildBasicUserInfo 将 UserDTO 转换为 BasicUserInfo
func BuildBasicUserInfo(user *contracts.UserDTO, needChange bool) BasicUserInfo {
if user == nil {
return BasicUserInfo{}
}
return BasicUserInfo{
ID: user.ID,
Username: user.Username,
Nickname: user.Nickname,
Email: user.Email,
AvatarURL: user.AvatarURL,
IsAdmin: user.IsAdmin,
NeedChangePassword: needChange || user.NeedChangePassword,
Bio: user.Bio,
Phone: user.Phone,
Gender: user.Gender,
Website: user.Website,
Location: user.Location,
}
}
type oauthStatePayload struct {
SourceName string `json:"source_name"`
Purpose string `json:"purpose"`
UserID uint64 `json:"user_id,omitempty"`
SessionHash string `json:"session_hash"`
}
func encodeOAuthStatePayload(payload oauthStatePayload) (string, error) {
data, err := json.Marshal(payload)
if err != nil {
return "", err
}
return string(data), nil
}
func decodeOAuthStatePayload(value string) (oauthStatePayload, error) {
var payload oauthStatePayload
if err := json.Unmarshal([]byte(value), &payload); err != nil {
return oauthStatePayload{}, err
}
return payload, nil
}
type loginRequiredAuditLog struct {
UserID uint64 `json:"user_id"`
Username string `json:"username"`
ClientIP string `json:"client_ip"`
Method string `json:"method"`
Path string `json:"path"`
RequestURI string `json:"request_uri"`
UserAgent string `json:"user_agent"`
Referer string `json:"referer"`
}
// ParseUserID parses a string or float64 user ID representation.
func ParseUserID(v any) uint64 {
switch val := v.(type) {
case uint64:
return val
case int64:
if val > 0 {
return uint64(val)
}
case int:
if val > 0 {
return uint64(val)
}
case float64:
if val > 0 {
return uint64(val)
}
case string:
if id, err := strconv.ParseUint(val, 10, 64); err == nil {
return id
}
}
return 0
}
+52 -47
View File
@@ -8,6 +8,9 @@ import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"Wavelet/plugins/domain/auth/controller"
"Wavelet/plugins/domain/auth/dao"
"Wavelet/plugins/domain/auth/service"
"context"
"embed"
"reflect"
@@ -83,39 +86,60 @@ func (p *Plugin) DeclareConfig() []core.ConfigBinding {
// Apply registers the auth migrations, services, routes, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
var cfg SessionConfig
if err := ctx.Config().Bind("app", &cfg); err == nil {
SetSessionConfig(cfg)
if cfg.SessionSecret != "" {
SetCapSecret([]byte(cfg.SessionSecret))
if err := ctx.Config().Bind("app", &cfg); err != nil {
cfg = SessionConfig{
SessionCookieName: "wavelet_session",
SessionAge: 86400,
SessionHTTPOnly: true,
}
}
core.Bind[contracts.DBService](ctx, setDBService)
core.Bind[contracts.CacheService](ctx, setCacheService)
core.Bind[contracts.LimiterService](ctx, setLimiterService)
d := dao.New(nil, nil, nil)
core.Bind[contracts.DBService](ctx, d.SetDBService)
core.Bind[contracts.CacheService](ctx, d.SetCacheService)
core.Bind[contracts.LimiterService](ctx, d.SetLimiterService)
ctx.OnDispose(func() error {
setDBService(nil)
setCacheService(nil)
setLimiterService(nil)
d.SetDBService(nil)
d.SetCacheService(nil)
d.SetLimiterService(nil)
return nil
})
var capSecret []byte
if cfg.SessionSecret != "" {
capSecret = []byte(cfg.SessionSecret)
}
svc := service.New(d, cfg, capSecret)
if p.authSvc != nil {
// Custom injected auth service override
core.Provide[contracts.AuthService](ctx, p.authSvc)
} else {
core.Provide[contracts.AuthService](ctx, svc.AuthSvc)
}
if p.authRegistry != nil {
core.Provide[contracts.AuthRegistry](ctx, p.authRegistry)
} else {
core.Provide[contracts.AuthRegistry](ctx, svc.AuthRegistry)
}
ctrl := controller.New(svc)
setDefaultRuntime(d, svc, ctrl)
// Register CaptchaService
captchaSvc := service.NewCaptchaService(
svc.CapManager,
func(scope string) any { return ctrl.VerifyCaptcha(scope) },
ctrl.Captcha.Challenge,
ctrl.Captcha.Redeem,
)
core.Provide[contracts.CaptchaService](ctx, captchaSvc)
// 1. Register migrations
ctx.Migrations().Register("auth", authMigrations)
// 2. Initialize and provide AuthService, AuthRegistry & CaptchaService
if p.authSvc == nil {
p.authSvc = newAuthService()
}
if p.authRegistry == nil {
p.authRegistry = newAuthRegistry()
}
core.Provide[contracts.AuthService](ctx, p.authSvc)
core.Provide[contracts.AuthRegistry](ctx, p.authRegistry)
core.Provide[contracts.CaptchaService](ctx, captchaService{})
// 2.1 Register Public / Auth Whitelist Endpoints
// 2. Register Public / Auth Whitelist Endpoints
publicEndpoints := []string{
"/api/v1/oauth/sources",
"/api/v1/oauth/login",
@@ -131,30 +155,11 @@ func (p *Plugin) Apply(ctx *core.Context) error {
"/api/healthz",
"/metrics",
}
RegisterWhitelist(publicEndpoints...)
ctrl.RegisterWhitelist(publicEndpoints...)
ctx.Router().RegisterWhitelist(publicEndpoints...)
// 3. Register HTTP Routes
oauthGroup := ctx.Router().Group("/api/v1/oauth")
{
oauthGroup.GET("/sources", GetLoginSources)
oauthGroup.GET("/login", GetLoginURL)
oauthGroup.GET("/:source/authorize", Authorize)
oauthGroup.GET("/logout", Logout)
oauthGroup.POST("/callback", Callback)
oauthGroup.GET("/user-info", LoginRequired(), UserInfo)
oauthGroup.GET("/external-accounts", LoginRequired(), ListExternalAccounts)
oauthGroup.POST("/external-accounts/:id/delete", LoginRequired(), DeleteExternalAccount)
}
ctx.Router().GET("/api/v1/user-info", LoginRequired(), UserInfo)
// 3.1 Register CAPTCHA HTTP Routes
capGroup := ctx.Router().Group("/api/v1/cap")
{
capGroup.GET("/challenge", Challenge)
capGroup.POST("/challenge", Challenge)
capGroup.POST("/redeem", Redeem)
}
ctrl.RegisterRoutes(ctx.Router())
// 4. Register Settings Schemas
const (
@@ -193,17 +198,17 @@ func (p *Plugin) Apply(ctx *core.Context) error {
// 5. Register Event Listeners for domain events
ctx.Events().On(contracts.EventTopicUserStatusChanged, func(c context.Context, e contracts.UserStatusChangedEvent) error {
InvalidateCachedUser(c, e.UserID)
svc.DAO.InvalidateCachedUser(c, e.UserID)
return nil
})
ctx.Events().On(contracts.EventTopicUserDeleted, func(c context.Context, e contracts.UserDeletedEvent) error {
InvalidateCachedUser(c, e.TargetUserID)
svc.DAO.InvalidateCachedUser(c, e.TargetUserID)
return nil
})
ctx.Events().On(contracts.EventTopicConfigChanged, func(_ any) {
InvalidateCapRuntimeSettings()
svc.CapSettings.Invalidate()
})
return nil
-229
View File
@@ -1,229 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/pkg/util"
"context"
"sync"
"time"
"gorm.io/gorm"
)
var (
dbMu sync.RWMutex
dbSvc contracts.DBService
cacheMu sync.RWMutex
cacheSvc contracts.CacheService
limiterMu sync.RWMutex
limiterSvc contracts.LimiterService
)
func setDBService(s contracts.DBService) {
dbMu.Lock()
defer dbMu.Unlock()
dbSvc = s
}
func setCacheService(s contracts.CacheService) {
cacheMu.Lock()
defer cacheMu.Unlock()
cacheSvc = s
}
func setLimiterService(s contracts.LimiterService) {
limiterMu.Lock()
defer limiterMu.Unlock()
limiterSvc = s
}
func getDB(ctx context.Context) *gorm.DB {
if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil {
return s.DB(ctx)
}
dbMu.RLock()
s := dbSvc
dbMu.RUnlock()
if s != nil {
return s.DB(ctx)
}
return nil
}
func getCache(ctx context.Context) contracts.CacheService {
if s, err := core.InjectFrom[contracts.CacheService](ctx); err == nil && s != nil {
return s
}
cacheMu.RLock()
s := cacheSvc
cacheMu.RUnlock()
return s
}
func getLimiter(ctx context.Context) contracts.LimiterService {
if s, err := core.InjectFrom[contracts.LimiterService](ctx); err == nil && s != nil {
return s
}
limiterMu.RLock()
s := limiterSvc
limiterMu.RUnlock()
return s
}
// GetAccessTokenByHash 按令牌哈希读取访问令牌记录(仅取鉴权所需字段)
func GetAccessTokenByHash(ctx context.Context, tokenHash string) (*CachedToken, error) {
var row struct {
ID uint64
UserID uint64
IsAdmin bool
}
if err := getDB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&row).Error; err != nil {
return nil, err
}
return &CachedToken{
ID: row.ID,
UserID: row.UserID,
IsAdmin: row.IsAdmin,
}, nil
}
// GetActiveUserByID 读取仍处于启用状态的用户
func GetActiveUserByID(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
var user contracts.UserDTO
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&user).Error; err != nil {
return nil, err
}
return &user, nil
}
// GetUserByID 按 ID 读取用户(不限制启用状态)
func GetUserByID(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
var user contracts.UserDTO
if err := getDB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil {
return nil, err
}
return &user, nil
}
// InsertUser 新建用户记录
func InsertUser(ctx context.Context, user *contracts.UserDTO) error {
return getDB(ctx).Table("w_users").Create(user).Error
}
// TouchUserLastLogin 刷新用户最后登录时间
func TouchUserLastLogin(ctx context.Context, userID uint64, at time.Time) error {
return getDB(ctx).Table("w_users").Where("id = ?", userID).Update("last_login_at", at).Error
}
// ListSimilarUsernames 查询与基础用户名相同或带 `-序号` 后缀的用户名(用于用户名去重)
func ListSimilarUsernames(ctx context.Context, base string) ([]string, error) {
var existingUsernames []string
if err := getDB(ctx).Table("w_users").
Where("username = ? OR username LIKE ? ESCAPE '\\'", base, util.EscapeLike(base)+"-%").
Pluck("username", &existingUsernames).Error; err != nil {
return nil, err
}
return existingUsernames, nil
}
// GetSystemConfigValue 读取系统配置项原始值
func GetSystemConfigValue(ctx context.Context, key string) (string, error) {
var val string
if err := getDB(ctx).Table("w_system_configs").Where("key = ?", key).Pluck("value", &val).Error; err != nil {
return "", err
}
return val, nil
}
// ListAllAuthSources 获取全部认证源(含未启用),按 ID 升序
func ListAllAuthSources(ctx context.Context) ([]AuthSource, error) {
var sources []AuthSource
if err := getDB(ctx).Order("id ASC").Find(&sources).Error; err != nil {
return nil, err
}
return sources, nil
}
// GetAuthSourceByID 根据 ID 获取认证源
func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) {
var src AuthSource
if err := getDB(ctx).First(&src, id).Error; err != nil {
return nil, err
}
return &src, nil
}
// GetAuthSourceByName 根据名称获取认证源
func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error) {
var src AuthSource
if err := getDB(ctx).Where("name = ?", name).First(&src).Error; err != nil {
return nil, err
}
return &src, nil
}
// ListActiveAuthSources 获取所有启用的认证源
func ListActiveAuthSources(ctx context.Context) ([]AuthSource, error) {
var sources []AuthSource
if err := getDB(ctx).Where("is_active = ?", true).Order("id ASC").Find(&sources).Error; err != nil {
return nil, err
}
return sources, nil
}
// CreateAuthSourceRecord 新建认证源记录
func CreateAuthSourceRecord(ctx context.Context, source *AuthSource) error {
return getDB(ctx).Create(source).Error
}
// SaveAuthSourceRecord 全量保存认证源记录
func SaveAuthSourceRecord(ctx context.Context, source *AuthSource) error {
return getDB(ctx).Save(source).Error
}
// DeleteAuthSourceRecord 删除认证源记录
func DeleteAuthSourceRecord(ctx context.Context, source *AuthSource) error {
return getDB(ctx).Delete(source).Error
}
// GetActiveAuthSourcesCached 获取所有启用的认证源(带缓存或直接查询)
func GetActiveAuthSourcesCached(ctx context.Context) ([]AuthSource, error) {
return ListActiveAuthSources(ctx)
}
// GetAuthSourceByNameCached 根据名称获取认证源(带缓存或直接查询)
func GetAuthSourceByNameCached(ctx context.Context, name string) (*AuthSource, error) {
return GetAuthSourceByName(ctx, name)
}
// FindExternalAccount 查询指定认证源的外部账号绑定
func FindExternalAccount(ctx context.Context, authSourceID uint64, externalID string) (*ExternalAccount, error) {
var account ExternalAccount
if err := getDB(ctx).Where("auth_source_id = ? AND external_id = ?", authSourceID, externalID).First(&account).Error; err != nil {
return nil, err
}
return &account, nil
}
// BindExternalAccount 绑定外部账号
func BindExternalAccount(ctx context.Context, account *ExternalAccount) error {
return getDB(ctx).Create(account).Error
}
// ListExternalAccountsByUserID 获取用户绑定的所有外部账号
func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]ExternalAccount, error) {
var accounts []ExternalAccount
if err := getDB(ctx).Where("user_id = ?", userID).Find(&accounts).Error; err != nil {
return nil, err
}
return accounts, nil
}
// UnbindExternalAccount 解绑外部账号
func UnbindExternalAccount(ctx context.Context, id, userID uint64) error {
return getDB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error
}
-262
View File
@@ -1,262 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"context"
"errors"
"sync"
)
type authServiceImpl struct{}
func newAuthService() contracts.AuthService {
return &authServiceImpl{}
}
func (s *authServiceImpl) RequireAuthMiddleware() any {
return LoginRequired()
}
func (s *authServiceImpl) RequireAdminMiddleware() any {
return AdminRequired()
}
// GetCurrentUser 从 context 中读取登录用户。
//
// 中间件通过 gin 的 c.Set(contracts.AuthUserObjKey, user) 写入登录态;
// *gin.Context 自身实现了 context.Context,且其 Value(key) 对 string 类型 key
// 等价于 c.Get(key)(未命中时再回落到 Request.Context().Value),
// 因此这里无需感知 gin 即可读取同一份登录态。
func (s *authServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
if v := ctx.Value(contracts.AuthUserObjKey); v != nil {
if u, ok := v.(*contracts.UserDTO); ok && u != nil {
return u, nil
}
}
return nil, errors.New(errUserNotInContext)
}
func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contracts.UserDTO, error) {
if token == "" {
return nil, errors.New(errEmptyToken)
}
tokenHash := hashToken(token)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil {
tokenRecord, err = GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, err
}
SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || user == nil || !user.IsActive {
user, err = GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, err
}
SetCachedUser(ctx, tokenRecord.UserID, user)
}
if user.Username == SystemUsername {
return nil, errors.New(errSystemUserTokenNotAllowed)
}
return user, nil
}
func (s *authServiceImpl) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) {
return "", nil
}
func (s *authServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64) error {
InvalidateCachedUser(ctx, userID)
return nil
}
// GetCurrentUserID 从请求登录态中读取用户 ID。
//
// Session 读取依赖 gin,属于接入层职责,因此这里通过接入层桥接函数
// currentUserIDFromRequestContext(见 middleware.go)取值,Service 层本身不感知 gin。
func (s *authServiceImpl) GetCurrentUserID(ctx context.Context) (uint64, error) {
userID, ok := currentUserIDFromRequestContext(ctx)
if !ok {
return 0, errors.New(errUserNotInContext)
}
return userID, nil
}
func (s *authServiceImpl) RevokeToken(ctx context.Context, tokenHash string) error {
InvalidateCachedToken(ctx, tokenHash)
return nil
}
func (s *authServiceImpl) DisallowTokenAuthMiddleware() any {
return DisallowTokenAuth()
}
func (s *authServiceImpl) InvalidateCachedUser(ctx context.Context, userID uint64) {
InvalidateCachedUser(ctx, userID)
}
func (s *authServiceImpl) InvalidateCachedToken(ctx context.Context, tokenHash string) {
InvalidateCachedToken(ctx, tokenHash)
}
func (s *authServiceImpl) ListAuthSources(ctx context.Context) ([]contracts.AuthSourceViewDTO, error) {
sources, err := ListAllAuthSources(ctx)
if err != nil {
return nil, err
}
views := make([]contracts.AuthSourceViewDTO, len(sources))
for i := range sources {
views[i] = contracts.AuthSourceViewDTO{
ID: sources[i].ID,
Name: sources[i].Name,
Type: sources[i].Type,
DisplayName: sources[i].DisplayName,
IsActive: sources[i].IsActive,
IconURL: sources[i].IconURL,
ClientSecretConfigured: sources[i].ClientSecret != "",
}
}
return views, nil
}
func (s *authServiceImpl) CreateAuthSource(ctx context.Context, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
model := AuthSource{
ID: source.ID,
Name: source.Name,
Type: source.Type,
DisplayName: source.DisplayName,
ClientID: source.ClientID,
ClientSecret: source.ClientSecret,
OpenIDDiscoveryURL: source.OpenIDDiscoveryURL,
Scopes: source.Scopes,
IconURL: source.IconURL,
IsActive: source.IsActive,
}
if err := model.Validate(); err != nil {
return nil, err
}
if err := CreateAuthSourceRecord(ctx, &model); err != nil {
return nil, err
}
model.Sanitize()
return toAuthSourceDTO(&model), nil
}
func (s *authServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
existing, err := GetAuthSourceByID(ctx, id)
if err != nil {
return nil, err
}
existing.DisplayName = source.DisplayName
existing.ClientID = source.ClientID
if source.ClientSecret != "" {
existing.ClientSecret = source.ClientSecret
}
existing.OpenIDDiscoveryURL = source.OpenIDDiscoveryURL
existing.Scopes = source.Scopes
existing.IconURL = source.IconURL
if err := existing.Validate(); err != nil {
return nil, err
}
if err := SaveAuthSourceRecord(ctx, existing); err != nil {
return nil, err
}
existing.Sanitize()
return toAuthSourceDTO(existing), nil
}
func (s *authServiceImpl) DeleteAuthSource(ctx context.Context, id uint64) error {
existing, err := GetAuthSourceByID(ctx, id)
if err != nil {
return err
}
return DeleteAuthSourceRecord(ctx, existing)
}
func (s *authServiceImpl) ToggleAuthSource(ctx context.Context, id uint64) (*contracts.AuthSourceDTO, error) {
existing, err := GetAuthSourceByID(ctx, id)
if err != nil {
return nil, err
}
existing.IsActive = !existing.IsActive
if err := SaveAuthSourceRecord(ctx, existing); err != nil {
return nil, err
}
existing.Sanitize()
return toAuthSourceDTO(existing), nil
}
func toAuthSourceDTO(s *AuthSource) *contracts.AuthSourceDTO {
if s == nil {
return nil
}
return &contracts.AuthSourceDTO{
ID: s.ID,
Name: s.Name,
Type: s.Type,
DisplayName: s.DisplayName,
ClientID: s.ClientID,
ClientSecret: s.ClientSecret,
OpenIDDiscoveryURL: s.OpenIDDiscoveryURL,
Scopes: s.Scopes,
IconURL: s.IconURL,
IsActive: s.IsActive,
CreatedAt: s.CreatedAt,
UpdatedAt: s.UpdatedAt,
}
}
type authRegistryImpl struct {
mu sync.RWMutex
providers map[string]contracts.OAuthProvider
}
func newAuthRegistry() contracts.AuthRegistry {
return &authRegistryImpl{
providers: make(map[string]contracts.OAuthProvider),
}
}
func (r *authRegistryImpl) RegisterOAuthProvider(name string, provider contracts.OAuthProvider) {
r.mu.Lock()
defer r.mu.Unlock()
r.providers[name] = provider
}
func (r *authRegistryImpl) GetOAuthProvider(name string) (contracts.OAuthProvider, bool) {
r.mu.RLock()
defer r.mu.RUnlock()
p, ok := r.providers[name]
return p, ok
}
func (r *authRegistryImpl) ListOAuthProviders() []string {
r.mu.RLock()
defer r.mu.RUnlock()
res := make([]string, 0, len(r.providers))
for name := range r.providers {
res = append(res, name)
}
return res
}
@@ -0,0 +1,49 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package service implements domain business services and orchestration for the auth plugin.
package service
import (
"Wavelet/core/contracts"
"sync"
)
// AuthRegistryImpl implements contracts.AuthRegistry.
type AuthRegistryImpl struct {
mu sync.RWMutex
providers map[string]contracts.OAuthProvider
}
// NewAuthRegistry creates a new AuthRegistryImpl.
func NewAuthRegistry() *AuthRegistryImpl {
return &AuthRegistryImpl{
providers: make(map[string]contracts.OAuthProvider),
}
}
// RegisterOAuthProvider registers an OAuthProvider by name.
func (r *AuthRegistryImpl) RegisterOAuthProvider(name string, provider contracts.OAuthProvider) {
r.mu.Lock()
defer r.mu.Unlock()
r.providers[name] = provider
}
// GetOAuthProvider retrieves an OAuthProvider by name.
func (r *AuthRegistryImpl) GetOAuthProvider(name string) (contracts.OAuthProvider, bool) {
r.mu.RLock()
defer r.mu.RUnlock()
p, ok := r.providers[name]
return p, ok
}
// ListOAuthProviders lists all registered provider names.
func (r *AuthRegistryImpl) ListOAuthProviders() []string {
r.mu.RLock()
defer r.mu.RUnlock()
res := make([]string, 0, len(r.providers))
for name := range r.providers {
res = append(res, name)
}
return res
}
@@ -0,0 +1,278 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package service implements domain business services and orchestration for the auth plugin.
package service
import (
"Wavelet/core/contracts"
"Wavelet/plugins/domain/auth/consts"
"Wavelet/plugins/domain/auth/dao"
"Wavelet/plugins/domain/auth/model/entity"
"context"
"crypto/sha256"
"encoding/hex"
"errors"
)
// HashToken computes SHA-256 hex digest of access token.
func HashToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
// UserIDExtractor extracts user ID from a request context.
type UserIDExtractor func(ctx context.Context) (uint64, bool)
// AuthServiceImpl implements contracts.AuthService.
type AuthServiceImpl struct {
dao *dao.DAO
requireAuthMiddleware any
requireAdminMiddleware any
disallowTokenMiddleware any
userIDExtractor UserIDExtractor
}
// NewAuthService creates a new AuthServiceImpl.
func NewAuthService(
d *dao.DAO,
requireAuth any,
requireAdmin any,
disallowToken any,
extractor UserIDExtractor,
) *AuthServiceImpl {
return &AuthServiceImpl{
dao: d,
requireAuthMiddleware: requireAuth,
requireAdminMiddleware: requireAdmin,
disallowTokenMiddleware: disallowToken,
userIDExtractor: extractor,
}
}
// SetMiddlewareHandlers wires middleware handlers into AuthService after controller initialization.
func (s *AuthServiceImpl) SetMiddlewareHandlers(requireAuth, requireAdmin, disallowToken any, extractor UserIDExtractor) {
s.requireAuthMiddleware = requireAuth
s.requireAdminMiddleware = requireAdmin
s.disallowTokenMiddleware = disallowToken
s.userIDExtractor = extractor
}
// RequireAuthMiddleware returns the authentication check middleware.
func (s *AuthServiceImpl) RequireAuthMiddleware() any {
return s.requireAuthMiddleware
}
// RequireAdminMiddleware returns the admin authorization middleware.
func (s *AuthServiceImpl) RequireAdminMiddleware() any {
return s.requireAdminMiddleware
}
// DisallowTokenAuthMiddleware returns the token rejection middleware.
func (s *AuthServiceImpl) DisallowTokenAuthMiddleware() any {
return s.disallowTokenMiddleware
}
// GetCurrentUser 从 context 中读取登录用户。
func (s *AuthServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
if v := ctx.Value(contracts.AuthUserObjKey); v != nil {
if u, ok := v.(*contracts.UserDTO); ok && u != nil {
return u, nil
}
}
return nil, errors.New(consts.ErrUserNotInContext)
}
// GetCurrentUserID 从请求登录态中读取用户 ID。
func (s *AuthServiceImpl) GetCurrentUserID(ctx context.Context) (uint64, error) {
if s.userIDExtractor != nil {
if userID, ok := s.userIDExtractor(ctx); ok {
return userID, nil
}
}
return 0, errors.New(consts.ErrUserNotInContext)
}
// VerifyToken 验证访问令牌并返回对应的用户。
func (s *AuthServiceImpl) VerifyToken(ctx context.Context, token string) (*contracts.UserDTO, error) {
if token == "" {
return nil, errors.New(consts.ErrEmptyToken)
}
tokenHash := HashToken(token)
tokenRecord, err := s.dao.GetCachedToken(ctx, tokenHash)
if err != nil {
tokenRecord, err = s.dao.GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, err
}
s.dao.SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := s.dao.GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || user == nil || !user.IsActive {
user, err = s.dao.GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, err
}
s.dao.SetCachedUser(ctx, tokenRecord.UserID, user)
}
if user.Username == consts.SystemUsername {
return nil, errors.New(consts.ErrSystemUserTokenNotAllowed)
}
return user, nil
}
// CreateSession establishes an authenticated session.
func (s *AuthServiceImpl) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) {
return "", nil
}
// RevokeUserSessions revokes active sessions and cached tokens for a user.
func (s *AuthServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64) error {
s.dao.InvalidateCachedUser(ctx, userID)
return nil
}
// RevokeToken invalidates a cached token by its hash.
func (s *AuthServiceImpl) RevokeToken(ctx context.Context, tokenHash string) error {
s.dao.InvalidateCachedToken(ctx, tokenHash)
return nil
}
// InvalidateCachedUser invalidates cached user profile.
func (s *AuthServiceImpl) InvalidateCachedUser(ctx context.Context, userID uint64) {
s.dao.InvalidateCachedUser(ctx, userID)
}
// InvalidateCachedToken invalidates cached access token.
func (s *AuthServiceImpl) InvalidateCachedToken(ctx context.Context, tokenHash string) {
s.dao.InvalidateCachedToken(ctx, tokenHash)
}
// ListAuthSources lists all configured authentication sources.
func (s *AuthServiceImpl) ListAuthSources(ctx context.Context) ([]contracts.AuthSourceViewDTO, error) {
sources, err := s.dao.ListAllAuthSources(ctx)
if err != nil {
return nil, err
}
views := make([]contracts.AuthSourceViewDTO, len(sources))
for i := range sources {
views[i] = contracts.AuthSourceViewDTO{
ID: sources[i].ID,
Name: sources[i].Name,
Type: sources[i].Type,
DisplayName: sources[i].DisplayName,
IsActive: sources[i].IsActive,
IconURL: sources[i].IconURL,
ClientSecretConfigured: sources[i].ClientSecret != "",
}
}
return views, nil
}
// CreateAuthSource creates a new authentication source.
func (s *AuthServiceImpl) CreateAuthSource(ctx context.Context, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
model := entity.AuthSource{
ID: source.ID,
Name: source.Name,
Type: source.Type,
DisplayName: source.DisplayName,
ClientID: source.ClientID,
ClientSecret: source.ClientSecret,
OpenIDDiscoveryURL: source.OpenIDDiscoveryURL,
Scopes: source.Scopes,
IconURL: source.IconURL,
IsActive: source.IsActive,
}
if err := model.Validate(); err != nil {
return nil, err
}
if err := s.dao.CreateAuthSource(ctx, &model); err != nil {
return nil, err
}
model.Sanitize()
return toAuthSourceDTO(&model), nil
}
// UpdateAuthSource updates an existing authentication source.
func (s *AuthServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
existing, err := s.dao.GetAuthSourceByID(ctx, id)
if err != nil {
return nil, err
}
existing.DisplayName = source.DisplayName
existing.ClientID = source.ClientID
if source.ClientSecret != "" {
existing.ClientSecret = source.ClientSecret
}
existing.OpenIDDiscoveryURL = source.OpenIDDiscoveryURL
existing.Scopes = source.Scopes
existing.IconURL = source.IconURL
if err := existing.Validate(); err != nil {
return nil, err
}
if err := s.dao.SaveAuthSource(ctx, existing); err != nil {
return nil, err
}
existing.Sanitize()
return toAuthSourceDTO(existing), nil
}
// DeleteAuthSource deletes an authentication source.
func (s *AuthServiceImpl) DeleteAuthSource(ctx context.Context, id uint64) error {
existing, err := s.dao.GetAuthSourceByID(ctx, id)
if err != nil {
return err
}
return s.dao.DeleteAuthSource(ctx, existing)
}
// ToggleAuthSource toggles active status of an authentication source.
func (s *AuthServiceImpl) ToggleAuthSource(ctx context.Context, id uint64) (*contracts.AuthSourceDTO, error) {
existing, err := s.dao.GetAuthSourceByID(ctx, id)
if err != nil {
return nil, err
}
existing.IsActive = !existing.IsActive
if err := s.dao.SaveAuthSource(ctx, existing); err != nil {
return nil, err
}
existing.Sanitize()
return toAuthSourceDTO(existing), nil
}
func toAuthSourceDTO(s *entity.AuthSource) *contracts.AuthSourceDTO {
if s == nil {
return nil
}
return &contracts.AuthSourceDTO{
ID: s.ID,
Name: s.Name,
Type: s.Type,
DisplayName: s.DisplayName,
ClientID: s.ClientID,
ClientSecret: s.ClientSecret,
OpenIDDiscoveryURL: s.OpenIDDiscoveryURL,
Scopes: s.Scopes,
IconURL: s.IconURL,
IsActive: s.IsActive,
CreatedAt: s.CreatedAt,
UpdatedAt: s.UpdatedAt,
}
}
@@ -1,43 +1,46 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// Package service implements domain business services and orchestration for the auth plugin.
package service
import (
"Wavelet/core/contracts"
"Wavelet/plugins/domain/auth/consts"
"Wavelet/plugins/domain/auth/model/dto"
"Wavelet/plugins/domain/auth/pow"
"context"
"crypto/sha256"
"encoding/hex"
"strconv"
"strings"
"sync"
"time"
)
const (
redeemTokenIDLength = 8 // 兑换 Token ID 字节长度
redeemVerTokenLength = 15 // 兑换验证 Token 字节长度
tokenPartsCount = 2 // 兑换 Token 由两部分组成
valuePartsCount = 2 // 存储值由 scope 和过期时间组成
)
// CaptchaManager orchestrates challenge generation and solution validation.
type CaptchaManager struct {
secret []byte
store pow.Store
secret []byte
store pow.Store
settingsMgr *CapSettingsManager
}
// NewCaptchaManager creates a new CAPTCHA Manager.
func NewCaptchaManager(secret []byte, store pow.Store) *CaptchaManager {
func NewCaptchaManager(secret []byte, store pow.Store, settingsMgr *CapSettingsManager) *CaptchaManager {
return &CaptchaManager{
secret: secret,
store: store,
secret: secret,
store: store,
settingsMgr: settingsMgr,
}
}
// SetSecret updates the shared secret used for PoW generation and validation.
func (m *CaptchaManager) SetSecret(secret []byte) {
m.secret = secret
}
// Generate creates a challenge response.
func (m *CaptchaManager) Generate(ctx context.Context, scope string) (*pow.ChallengeResponse, error) {
settings, err := CurrentCapSettings(ctx)
settings, err := m.settingsMgr.Current(ctx)
if err != nil {
return nil, err
}
@@ -52,17 +55,17 @@ func (m *CaptchaManager) Generate(ctx context.Context, scope string) (*pow.Chall
}
// Redeem verifies PoW solutions and returns a one-time redeem token.
func (m *CaptchaManager) Redeem(ctx context.Context, token string, solutions []int, scope string) (*RedeemResponse, error) {
func (m *CaptchaManager) Redeem(ctx context.Context, token string, solutions []int, scope string) (*dto.RedeemResponse, error) {
sigHex := pow.JwtSigHex(token)
if sigHex == "" {
return &RedeemResponse{Success: false, Error: redeemErrInvalidToken}, nil
return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrInvalidToken}, nil
}
nonceKey := "cap:nonce:" + sigHex
payload, err := pow.VerifyChallengeSolutions(token, solutions, m.secret, scope)
if err != nil {
return &RedeemResponse{Success: false, Error: err.Error()}, nil //nolint:nilerr // validation errors are returned as response, not system errors
return &dto.RedeemResponse{Success: false, Error: err.Error()}, nil //nolint:nilerr // validation errors returned as response
}
now := time.Now().UnixNano() / int64(time.Millisecond)
@@ -73,19 +76,19 @@ func (m *CaptchaManager) Redeem(ctx context.Context, token string, solutions []i
set, err := m.store.SetNX(ctx, nonceKey, "1", nonceTTL)
if err != nil {
return &RedeemResponse{Success: false, Error: redeemErrNonceStoreFailed}, err
return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrNonceStoreFailed}, err
}
if !set {
return &RedeemResponse{Success: false, Error: redeemErrAlreadyRedeemed}, nil
return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrAlreadyRedeemed}, nil
}
settings, err := CurrentCapSettings(ctx)
settings, err := m.settingsMgr.Current(ctx)
if err != nil {
return &RedeemResponse{Success: false, Error: redeemErrSettingsLoad}, err
return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrSettingsLoad}, err
}
id := pow.RandomHex(redeemTokenIDLength)
verToken := pow.RandomHex(redeemVerTokenLength)
id := pow.RandomHex(consts.RedeemTokenIDLength)
verToken := pow.RandomHex(consts.RedeemVerTokenLength)
verHashBytes := sha256.Sum256([]byte(verToken))
verHashHex := hex.EncodeToString(verHashBytes[:])
@@ -94,10 +97,10 @@ func (m *CaptchaManager) Redeem(ctx context.Context, token string, solutions []i
storeVal := strconv.FormatInt(tokenExpires.UnixNano(), 10) + "|" + scope
if err := m.store.Set(ctx, tokenKey, storeVal, settings.TokenTTL); err != nil {
return &RedeemResponse{Success: false, Error: redeemErrTokenStoreFailed}, err
return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrTokenStoreFailed}, err
}
return &RedeemResponse{
return &dto.RedeemResponse{
Success: true,
Token: id + ":" + verToken,
Expires: tokenExpires.UnixNano() / int64(time.Millisecond),
@@ -110,7 +113,7 @@ func (m *CaptchaManager) VerifyToken(ctx context.Context, token, expectedScope s
return false, nil
}
parts := strings.Split(token, ":")
if len(parts) != tokenPartsCount {
if len(parts) != consts.TokenPartsCount {
return false, nil
}
id := parts[0]
@@ -121,7 +124,10 @@ func (m *CaptchaManager) VerifyToken(ctx context.Context, token, expectedScope s
tokenKey := "cap:token:" + id + ":" + verHashHex
val, exists, err := sGetAndDelete(ctx, m.store, tokenKey)
if m.store == nil {
return false, nil
}
val, exists, err := m.store.GetAndDelete(ctx, tokenKey)
if err != nil {
return false, err
}
@@ -130,13 +136,13 @@ func (m *CaptchaManager) VerifyToken(ctx context.Context, token, expectedScope s
}
valParts := strings.Split(val, "|")
if len(valParts) != valuePartsCount {
if len(valParts) != consts.ValuePartsCount {
return false, nil
}
expNano, err := strconv.ParseInt(valParts[0], 10, 64)
if err != nil {
return false, nil //nolint:nilerr // invalid format is treated as validation failure
return false, nil //nolint:nilerr // invalid format is failure
}
tokenScope := valParts[1]
@@ -151,41 +157,38 @@ func (m *CaptchaManager) VerifyToken(ctx context.Context, token, expectedScope s
return true, nil
}
func sGetAndDelete(ctx context.Context, store pow.Store, key string) (string, bool, error) {
if store == nil {
return "", false, nil
}
return store.GetAndDelete(ctx, key)
// CaptchaServiceImpl implements contracts.CaptchaService.
type CaptchaServiceImpl struct {
manager *CaptchaManager
verifyMiddleware func(scope string) any
challengeHandler any
redeemHandler any
}
var (
defaultCapManagerMu sync.RWMutex
defaultCapManager *CaptchaManager
)
// SetCapSecret sets the shared secret used by the default CAPTCHA manager.
func SetCapSecret(secret []byte) {
defaultCapManagerMu.Lock()
defer defaultCapManagerMu.Unlock()
if len(secret) > 0 {
store := pow.NewMemoryStore(1 * time.Minute)
defaultCapManager = NewCaptchaManager(secret, store)
// NewCaptchaService creates a new CaptchaServiceImpl.
func NewCaptchaService(mgr *CaptchaManager, verifyMiddleware func(scope string) any, challengeHandler any, redeemHandler any) contracts.CaptchaService {
return &CaptchaServiceImpl{
manager: mgr,
verifyMiddleware: verifyMiddleware,
challengeHandler: challengeHandler,
redeemHandler: redeemHandler,
}
}
// GetDefaultCapManager yields the global singleton CAPTCHA manager.
func GetDefaultCapManager() *CaptchaManager {
defaultCapManagerMu.RLock()
defer defaultCapManagerMu.RUnlock()
return defaultCapManager
// VerifyMiddleware returns the captcha verification middleware.
func (s *CaptchaServiceImpl) VerifyMiddleware(scope string) any {
if s.verifyMiddleware != nil {
return s.verifyMiddleware(scope)
}
return nil
}
type captchaService struct{}
func (captchaService) VerifyMiddleware(scope string) any {
return VerifyCaptchaMiddleware(GetDefaultCapManager(), scope)
// ChallengeHandler returns the challenge HTTP handler.
func (s *CaptchaServiceImpl) ChallengeHandler() any {
return s.challengeHandler
}
func (captchaService) ChallengeHandler() any { return Challenge }
func (captchaService) RedeemHandler() any { return Redeem }
// RedeemHandler returns the redeem HTTP handler.
func (s *CaptchaServiceImpl) RedeemHandler() any {
return s.redeemHandler
}
@@ -0,0 +1,119 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package service implements domain business services and orchestration for the auth plugin.
package service
import (
"Wavelet/plugins/domain/auth/consts"
"Wavelet/plugins/domain/auth/dao"
"Wavelet/plugins/domain/auth/model/do"
"context"
"errors"
"sync/atomic"
"golang.org/x/sync/singleflight"
)
var capRuntimeConfigKeys = []string{
consts.ConfigKeyCapLoginEnabled,
consts.ConfigKeyCapChallengeCount,
consts.ConfigKeyCapChallengeSize,
consts.ConfigKeyCapChallengeDifficulty,
consts.ConfigKeyCapChallengeTTL,
consts.ConfigKeyCapTokenTTL,
}
var capRuntimeConfigKeySet = func() map[string]struct{} {
set := make(map[string]struct{}, len(capRuntimeConfigKeys))
for _, key := range capRuntimeConfigKeys {
set[key] = struct{}{}
}
return set
}()
// IsCapRuntimeConfigKey reports whether a system config key affects CAPTCHA runtime settings.
func IsCapRuntimeConfigKey(key string) bool {
_, ok := capRuntimeConfigKeySet[key]
return ok
}
// CapSettingsManager manages dynamic CAPTCHA configuration cache.
type CapSettingsManager struct {
dao *dao.DAO
snapshot atomic.Pointer[do.CapRuntimeSettings]
loadGroup singleflight.Group
}
// NewCapSettingsManager creates a new CapSettingsManager.
func NewCapSettingsManager(d *dao.DAO) *CapSettingsManager {
return &CapSettingsManager{
dao: d,
}
}
// Invalidate drops the in-process CAPTCHA settings snapshot.
func (m *CapSettingsManager) Invalidate() {
m.snapshot.Store(nil)
}
// Current returns the cached CAPTCHA runtime settings snapshot.
func (m *CapSettingsManager) Current(ctx context.Context) (do.CapRuntimeSettings, error) {
if snapshot := m.snapshot.Load(); snapshot != nil {
return *snapshot, nil
}
loaded, err, _ := m.loadGroup.Do("cap-runtime-settings", func() (any, error) {
if snapshot := m.snapshot.Load(); snapshot != nil {
return *snapshot, nil
}
settings, loadErr := m.loadSettings(ctx)
if loadErr != nil {
return do.CapRuntimeSettings{}, loadErr
}
m.snapshot.Store(&settings)
return settings, nil
})
if err != nil {
return do.CapRuntimeSettings{}, err
}
settings, ok := loaded.(do.CapRuntimeSettings)
if !ok {
return do.CapRuntimeSettings{}, errors.New("cap runtime settings loader returned unexpected type")
}
return settings, nil
}
// CapProtectionEnabled reports whether CAPTCHA verification is required for protected routes.
func (m *CapSettingsManager) CapProtectionEnabled(ctx context.Context) bool {
settings, err := m.Current(ctx)
if err != nil {
return false
}
return settings.LoginEnabled
}
// InstallTestSnapshot installs a fixed snapshot for unit tests.
func (m *CapSettingsManager) InstallTestSnapshot(settings do.CapRuntimeSettings) func() {
snapshot := settings
m.snapshot.Store(&snapshot)
return m.Invalidate
}
func (m *CapSettingsManager) loadSettings(ctx context.Context) (do.CapRuntimeSettings, error) {
if m.dao == nil {
return do.ParseCapRuntimeSettings(nil), nil
}
records, err := m.dao.ListSystemConfigsByKeys(ctx, capRuntimeConfigKeys)
if err != nil {
return do.CapRuntimeSettings{}, err
}
configs := make(map[string]string, len(records))
for _, r := range records {
configs[r.Key] = r.Value
}
return do.ParseCapRuntimeSettings(configs), nil
}
@@ -0,0 +1,440 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package service implements domain business services and orchestration for the auth plugin.
package service
import (
"Wavelet/core/contracts"
"Wavelet/pkg/idgen"
"Wavelet/plugins/domain/auth/consts"
"Wavelet/plugins/domain/auth/dao"
"Wavelet/plugins/domain/auth/model/dto"
"Wavelet/plugins/domain/auth/model/entity"
"context"
"errors"
"fmt"
"strconv"
"strings"
"time"
"github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2"
"gorm.io/gorm"
)
// OAuthService orchestrates OAuth/OIDC operations.
type OAuthService struct {
dao *dao.DAO
providerCache *OIDCProviderCache
sessionSvc *SessionService
}
// NewOAuthService creates a new OAuthService.
func NewOAuthService(d *dao.DAO, cache *OIDCProviderCache, sessSvc *SessionService) *OAuthService {
return &OAuthService{
dao: d,
providerCache: cache,
sessionSvc: sessSvc,
}
}
// IsOIDCLoginEnabled checks if OIDC login is globally enabled.
func (s *OAuthService) IsOIDCLoginEnabled(ctx context.Context) bool {
val, err := s.dao.GetSystemConfigValue(ctx, "oidc_login_enabled")
if err != nil || val == "" {
return true
}
b, err := strconv.ParseBool(val)
if err != nil {
return true
}
return b
}
// ResolveAuthSource retrieves the specified or default active auth source.
func (s *OAuthService) ResolveAuthSource(ctx context.Context, sourceName string) (*entity.AuthSource, error) {
name := strings.TrimSpace(strings.ToLower(sourceName))
if name == "" {
sources, err := s.dao.ListActiveAuthSources(ctx)
if err != nil {
return nil, err
}
if len(sources) == 0 {
return nil, errors.New(consts.ErrNoActiveAuthSource)
}
src, err := s.dao.GetAuthSourceByName(ctx, sources[0].Name)
if err != nil {
return nil, err
}
return src, nil
}
src, err := s.dao.GetAuthSourceByName(ctx, name)
if err != nil {
return nil, err
}
return src, nil
}
// ActiveLoginSources returns all active login sources formatted for display.
func (s *OAuthService) ActiveLoginSources(ctx context.Context) ([]dto.AuthSourceView, error) {
if !s.IsOIDCLoginEnabled(ctx) {
return nil, nil
}
dbSources, err := s.dao.ListActiveAuthSources(ctx)
if err != nil {
return nil, err
}
sources := make([]dto.AuthSourceView, 0, len(dbSources))
for _, source := range dbSources {
sources = append(sources, dto.AuthSourceView{
ID: source.ID,
Name: source.Name,
Type: source.Type,
DisplayName: source.DisplayName,
IsActive: source.IsActive,
IconURL: source.IconURL,
ClientSecretConfigured: source.ClientSecretConfigured,
})
}
return sources, nil
}
// GetFrontendLoginRedirectURL constructs the OAuth frontend redirect URL.
func (s *OAuthService) GetFrontendLoginRedirectURL(ctx context.Context) (string, error) {
val, err := s.dao.GetSystemConfigValue(ctx, "server_address")
if err != nil || strings.TrimSpace(val) == "" {
return "", errors.New(consts.ErrServerAddressMissing)
}
return strings.TrimRight(val, "/") + "/login", nil
}
// ReserveOAuthStateSlot ensures that a session does not abuse OAuth state generation.
func (s *OAuthService) ReserveOAuthStateSlot(ctx context.Context, sessionHash string) error {
if sessionHash == "" {
return nil
}
if limiter := s.dao.Limiter(); limiter != nil {
key := fmt.Sprintf(consts.OAuthStateLimitKeyFormat, sessionHash)
res, err := limiter.Allow(ctx, key, contracts.Rate{
Limit: consts.OAuthStateLimitMax,
Period: consts.OAuthStateCacheKeyExpiration,
})
if err != nil {
return err
}
if !res.Allowed {
return errors.New(consts.ErrOAuthStateRateLimited)
}
return nil
}
cache := s.dao.Cache()
if cache == nil {
return nil
}
key := fmt.Sprintf(consts.OAuthStateLimitKeyFormat, sessionHash)
var count int
_ = cache.Get(ctx, key, &count)
count++
_ = cache.Set(ctx, key, count, consts.OAuthStateCacheKeyExpiration)
if count > consts.OAuthStateLimitMax {
return errors.New(consts.ErrOAuthStateRateLimited)
}
return nil
}
// BuildOAuthConfig builds oauth2.Config and oidc.IDTokenVerifier.
func (s *OAuthService) BuildOAuthConfig(ctx context.Context, source *entity.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
if source == nil {
return nil, nil, errors.New(consts.ErrAuthSourceRequired)
}
if source.OpenIDDiscoveryURL == "" {
return nil, nil, errors.New(consts.ErrDiscoveryURLRequired)
}
issuer := strings.TrimSuffix(strings.TrimSpace(source.OpenIDDiscoveryURL), "/")
issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration")
issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server")
provider, err := s.providerCache.Get(ctx, issuer)
if err != nil {
return nil, nil, err
}
verifier := provider.Verifier(&oidc.Config{ClientID: source.ClientID})
scopes := strings.Fields(source.Scopes)
if len(scopes) == 0 {
scopes = []string{oidc.ScopeOpenID, "profile", "email"}
}
if !containsScope(scopes, oidc.ScopeOpenID) {
scopes = append([]string{oidc.ScopeOpenID}, scopes...)
}
return &oauth2.Config{
ClientID: source.ClientID,
ClientSecret: source.ClientSecret,
RedirectURL: redirectURL,
Scopes: scopes,
Endpoint: provider.Endpoint(),
}, verifier, nil
}
func containsScope(scopes []string, scope string) bool {
for _, item := range scopes {
if item == scope {
return true
}
}
return false
}
// BuildAuthorizeURL generates the redirect authorize URL for the source and state.
func (s *OAuthService) BuildAuthorizeURL(ctx context.Context, source *entity.AuthSource, state string) (string, error) {
redirectURL, err := s.GetFrontendLoginRedirectURL(ctx)
if err != nil {
return "", err
}
authConfig, verifier, err := s.BuildOAuthConfig(ctx, source, redirectURL)
if err != nil {
return "", err
}
if verifier != nil {
return authConfig.AuthCodeURL(state, oidc.Nonce(state)), nil
}
return authConfig.AuthCodeURL(state), nil
}
// BuildOAuthUserInfo exchanges the auth code and retrieves user identity claims.
func (s *OAuthService) BuildOAuthUserInfo(ctx context.Context, source *entity.AuthSource, code, nonce, redirectURL string) (*contracts.OAuthUserInfoDTO, error) {
authConfig, verifier, err := s.BuildOAuthConfig(ctx, source, redirectURL)
if err != nil {
return nil, err
}
token, err := authConfig.Exchange(ctx, code)
if err != nil {
return nil, err
}
userInfo := &contracts.OAuthUserInfoDTO{Active: true}
if verifier != nil {
if verifyErr := s.verifyIDToken(ctx, verifier, token, nonce, userInfo); verifyErr != nil {
return nil, verifyErr
}
}
if userInfo.Username == "" && userInfo.PreferredUsername != "" {
userInfo.Username = userInfo.PreferredUsername
}
if userInfo.Username == "" && userInfo.Email != "" {
userInfo.Username = strings.Split(userInfo.Email, "@")[0]
}
if userInfo.Username == "" && userInfo.Sub != "" {
userInfo.Username = userInfo.Sub
}
if userInfo.Name == "" {
userInfo.Name = userInfo.Username
}
return userInfo, nil
}
func (s *OAuthService) verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *contracts.OAuthUserInfoDTO) error {
rawIDToken, ok := token.Extra("id_token").(string)
if !ok {
return nil
}
idToken, verifyErr := verifier.Verify(ctx, rawIDToken)
if verifyErr != nil {
return fmt.Errorf(consts.ErrIDTokenVerifyFailedFormat, consts.ErrIDTokenVerifyFailed, verifyErr)
}
if nonce != "" && idToken.Nonce != nonce {
return errors.New(consts.ErrNonceMismatch)
}
if claimsErr := idToken.Claims(userInfo); claimsErr != nil {
return claimsErr
}
return nil
}
// NormalizeOAuthUserInfo sanitizes user claims.
func (s *OAuthService) NormalizeOAuthUserInfo(userInfo *contracts.OAuthUserInfoDTO) error {
userInfo.Username = strings.TrimSpace(userInfo.Username)
userInfo.PreferredUsername = strings.TrimSpace(userInfo.PreferredUsername)
userInfo.Email = strings.TrimSpace(userInfo.Email)
userInfo.Name = strings.TrimSpace(userInfo.Name)
userInfo.AvatarURL = strings.TrimSpace(userInfo.AvatarURL)
if userInfo.Username == "" && userInfo.PreferredUsername != "" {
userInfo.Username = userInfo.PreferredUsername
}
if userInfo.Username == "" && userInfo.Email != "" {
userInfo.Username = strings.Split(userInfo.Email, "@")[0]
}
if userInfo.Username == "" && userInfo.Sub != "" {
userInfo.Username = userInfo.Sub
}
if userInfo.Username == "" {
return errors.New(consts.ErrUsernameFromSourceFailed)
}
if userInfo.Name == "" {
userInfo.Name = userInfo.Username
}
if !userInfo.Active {
userInfo.Active = true
}
return nil
}
// UniqueUsername generates a unique username given a base candidate.
func (s *OAuthService) UniqueUsername(ctx context.Context, base string) (string, error) {
base = strings.TrimSpace(base)
if base == "" {
base = "user"
}
existingUsernames, err := s.dao.ListSimilarUsernames(ctx, base)
if err != nil {
return "", err
}
exists := make(map[string]bool, len(existingUsernames))
for _, u := range existingUsernames {
exists[strings.ToLower(u)] = true
}
if !exists[strings.ToLower(base)] {
return base, nil
}
for i := 1; i <= 1000; i++ {
candidate := fmt.Sprintf("%s-%d", base, i)
if !exists[strings.ToLower(candidate)] {
return candidate, nil
}
}
return "", errors.New(consts.ErrUsernameGenerateFailed)
}
// BindExternalAccount binds an external identity to an existing user.
func (s *OAuthService) BindExternalAccount(ctx context.Context, sourceID, userID uint64, userInfo *contracts.OAuthUserInfoDTO) error {
user, err := s.dao.GetUserByID(ctx, userID)
if err != nil {
return err
}
if err := s.dao.BindExternalAccount(ctx, &entity.ExternalAccount{
AuthSourceID: sourceID,
UserID: user.ID,
ExternalID: userInfo.Sub,
ExternalUsername: userInfo.Username,
Email: userInfo.Email,
}); err != nil {
return err
}
user.LastLoginAt = time.Now()
_ = s.dao.TouchUserLastLogin(ctx, user.ID, user.LastLoginAt)
return nil
}
// AuthenticateOrRegisterUser finds existing binding or creates a new user.
func (s *OAuthService) AuthenticateOrRegisterUser(ctx context.Context, source *entity.AuthSource, userInfo *contracts.OAuthUserInfoDTO) (*contracts.UserDTO, bool, error) {
account, err := s.dao.FindExternalAccount(ctx, source.ID, userInfo.Sub)
if err == nil {
user, loadErr := s.dao.GetUserByID(ctx, account.UserID)
if loadErr != nil {
return nil, false, loadErr
}
user.LastLoginAt = time.Now()
_ = s.dao.TouchUserLastLogin(ctx, user.ID, user.LastLoginAt)
return user, true, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, false, err
}
// Not found -> check registration
registrationEnabled := true
val, cfgErr := s.dao.GetSystemConfigValue(ctx, "registration_enabled")
if cfgErr == nil && val != "" {
if b, err := strconv.ParseBool(val); err == nil {
registrationEnabled = b
}
}
if !registrationEnabled {
return nil, false, nil // registration disabled -> need bind
}
username, uniqueErr := s.UniqueUsername(ctx, userInfo.Username)
if uniqueErr != nil {
return nil, false, uniqueErr
}
userInfo.Username = username
now := time.Now()
user := contracts.UserDTO{
ID: idgen.NextUint64ID(),
Username: userInfo.Username,
Nickname: userInfo.Name,
Email: userInfo.Email,
AvatarURL: userInfo.AvatarURL,
IsActive: userInfo.Active,
LastLoginAt: now,
CreatedAt: now,
UpdatedAt: now,
}
if err := s.dao.InsertUser(ctx, &user); err != nil {
return nil, false, err
}
if err := s.dao.BindExternalAccount(ctx, &entity.ExternalAccount{
AuthSourceID: source.ID,
UserID: user.ID,
ExternalID: userInfo.Sub,
ExternalUsername: userInfo.Username,
Email: userInfo.Email,
}); err != nil {
return nil, false, err
}
return &user, true, nil
}
// ListExternalAccounts returns sanitized external account bindings.
func (s *OAuthService) ListExternalAccounts(ctx context.Context, userID uint64) ([]dto.ExternalAccountView, error) {
accounts, err := s.dao.ListExternalAccountsByUserID(ctx, userID)
if err != nil {
return nil, err
}
views := make([]dto.ExternalAccountView, len(accounts))
for i, acc := range accounts {
source, _ := s.dao.GetAuthSourceByID(ctx, acc.AuthSourceID)
sourceName, sourceType, sourceLabel := "", "", ""
if source != nil {
sourceName = source.Name
sourceType = source.Type
sourceLabel = source.DisplayName
}
views[i] = dto.ExternalAccountView{
ID: acc.ID,
AuthSourceID: acc.AuthSourceID,
AuthSourceName: sourceName,
AuthSourceType: sourceType,
AuthSourceLabel: sourceLabel,
ExternalUsername: acc.ExternalUsername,
Email: acc.Email,
CreatedAt: acc.CreatedAt.Format(time.RFC3339),
}
}
return views, nil
}
// DeleteExternalAccount unbinds an external account.
func (s *OAuthService) DeleteExternalAccount(ctx context.Context, id, userID uint64) error {
return s.dao.UnbindExternalAccount(ctx, id, userID)
}
@@ -1,7 +1,8 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// Package service implements domain business services and orchestration for the auth plugin.
package service
import (
"context"
@@ -13,16 +14,18 @@ import (
"golang.org/x/sync/singleflight"
)
// oidcProviderCache 进程级 OIDC provider 缓存。
type oidcProviderCache struct {
// OIDCProviderCache 进程级 OIDC provider 缓存。
type OIDCProviderCache struct {
mu sync.RWMutex
entries map[string]*oidc.Provider // key: normalized issuer URL
sfGroup singleflight.Group
}
// globalOIDCProviderCache 是包级单例缓存,与进程同生命周期。
var globalOIDCProviderCache = &oidcProviderCache{
entries: make(map[string]*oidc.Provider),
// NewOIDCProviderCache creates a new OIDCProviderCache.
func NewOIDCProviderCache() *OIDCProviderCache {
return &OIDCProviderCache{
entries: make(map[string]*oidc.Provider),
}
}
// discoveryContext 从请求 ctx 提取 HTTP 客户端,并绑定到不可取消的 Background ctx。
@@ -34,8 +37,8 @@ func discoveryContext(ctx context.Context) context.Context {
return bg
}
// get 返回缓存的 provider;若无则通过 oidc.NewProvider 获取并写入缓存。
func (c *oidcProviderCache) get(ctx context.Context, issuer string) (*oidc.Provider, error) {
// Get 返回缓存的 provider;若无则通过 oidc.NewProvider 获取并写入缓存。
func (c *OIDCProviderCache) Get(ctx context.Context, issuer string) (*oidc.Provider, error) {
c.mu.RLock()
if p, ok := c.entries[issuer]; ok {
c.mu.RUnlock()
@@ -68,14 +71,9 @@ func (c *oidcProviderCache) get(ctx context.Context, issuer string) (*oidc.Provi
return v.(*oidc.Provider), nil //nolint:forcetypeassert
}
// invalidate 从缓存中移除指定 issuer 对应的 provider。
func (c *oidcProviderCache) invalidate(issuer string) {
// Invalidate 从缓存中移除指定 issuer 对应的 provider。
func (c *OIDCProviderCache) Invalidate(issuer string) {
c.mu.Lock()
delete(c.entries, issuer)
c.mu.Unlock()
}
// InvalidateOIDCProviderCache 从进程级缓存中清除指定 issuer 的 provider 条目。
func InvalidateOIDCProviderCache(issuer string) {
globalOIDCProviderCache.invalidate(issuer)
}
@@ -0,0 +1,51 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package service implements domain business services and orchestration for the auth plugin.
package service
import (
"Wavelet/plugins/domain/auth/dao"
"Wavelet/plugins/domain/auth/pow"
"time"
)
// Service aggregates all domain services for the auth plugin.
type Service struct {
DAO *dao.DAO
Session *SessionService
OAuth *OAuthService
OIDCProviderCache *OIDCProviderCache
CapSettings *CapSettingsManager
CapManager *CaptchaManager
AuthSvc *AuthServiceImpl
AuthRegistry *AuthRegistryImpl
}
// New creates a new Service container with all domain services wired up.
func New(d *dao.DAO, sessionCfg SessionConfig, capSecret []byte) *Service {
sessionSvc := NewSessionService(sessionCfg, d)
oidcCache := NewOIDCProviderCache()
oauthSvc := NewOAuthService(d, oidcCache, sessionSvc)
capSettings := NewCapSettingsManager(d)
var capStore pow.Store
if len(capSecret) > 0 {
capStore = pow.NewMemoryStore(1 * time.Minute)
}
capMgr := NewCaptchaManager(capSecret, capStore, capSettings)
authSvc := NewAuthService(d, nil, nil, nil, nil)
authRegistry := NewAuthRegistry()
return &Service{
DAO: d,
Session: sessionSvc,
OAuth: oauthSvc,
OIDCProviderCache: oidcCache,
CapSettings: capSettings,
CapManager: capMgr,
AuthSvc: authSvc,
AuthRegistry: authRegistry,
}
}
@@ -0,0 +1,177 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package service implements domain business services and orchestration for the auth plugin.
package service
import (
"Wavelet/core/contracts"
"Wavelet/plugins/domain/auth/consts"
"Wavelet/plugins/domain/auth/dao"
"context"
"crypto/sha256"
"encoding/hex"
"net/http"
"strconv"
"strings"
"sync"
"github.com/gin-contrib/sessions"
"github.com/google/uuid"
gsessions "github.com/gorilla/sessions"
)
// SessionConfig defines session settings.
type SessionConfig struct {
SessionCookieName string `config:"session_cookie_name" env:"APP_SESSION_COOKIE_NAME" default:"wavelet_session"`
SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"`
SessionDomain string `config:"session_domain" env:"APP_SESSION_DOMAIN"`
SessionAge int `config:"session_age" env:"APP_SESSION_AGE" default:"86400"`
SessionHTTPOnly bool `config:"session_http_only" env:"APP_SESSION_HTTP_ONLY" default:"true"`
SessionSecure bool `config:"session_secure" env:"APP_SESSION_SECURE"`
}
// SessionService manages HTTP session operations, cookies, and tokens.
type SessionService struct {
mu sync.RWMutex
config SessionConfig
dao *dao.DAO
}
// NewSessionService creates a new SessionService.
func NewSessionService(cfg SessionConfig, d *dao.DAO) *SessionService {
return &SessionService{
config: cfg,
dao: d,
}
}
// SetConfig updates the active session configuration.
func (s *SessionService) SetConfig(cfg SessionConfig) {
s.mu.Lock()
defer s.mu.Unlock()
s.config = cfg
}
// Config returns the current session configuration.
func (s *SessionService) Config() SessionConfig {
s.mu.RLock()
defer s.mu.RUnlock()
return s.config
}
// GetSessionOptions 根据配置构建 Session 选项
func (s *SessionService) GetSessionOptions(maxAge int) sessions.Options {
cfg := s.Config()
return sessions.Options{
Path: "/",
Domain: cfg.SessionDomain,
MaxAge: maxAge,
HttpOnly: cfg.SessionHTTPOnly,
Secure: cfg.SessionSecure,
SameSite: http.SameSiteLaxMode,
}
}
// StripCookieMaxAgeAndExpires 从 Set-Cookie 响应头中移除 Max-Age 和 Expires,从而使其成为浏览器会话 Cookie
func (s *SessionService) StripCookieMaxAgeAndExpires(header http.Header, cookieName string) {
headers := header["Set-Cookie"]
if len(headers) == 0 {
return
}
newHeaders := make([]string, 0, len(headers))
for _, h := range headers {
if strings.HasPrefix(h, cookieName+"=") {
parts := strings.Split(h, ";")
newParts := make([]string, 0, len(parts))
for _, p := range parts {
trimmed := strings.TrimSpace(p)
lower := strings.ToLower(trimmed)
if strings.HasPrefix(lower, "max-age=") || strings.HasPrefix(lower, "expires=") {
continue
}
newParts = append(newParts, p)
}
newHeaders = append(newHeaders, strings.Join(newParts, ";"))
} else {
newHeaders = append(newHeaders, h)
}
}
header["Set-Cookie"] = newHeaders
}
// EnsureSessionToken returns or generates the session unique token.
func (s *SessionService) EnsureSessionToken(session sessions.Session) (string, bool) {
token, ok := session.Get(consts.SessionTokenKey).(string)
if !ok || token == "" {
token = uuid.NewString()
session.Set(consts.SessionTokenKey, token)
return token, true
}
return token, false
}
// HashSessionToken hashes the session token using SHA-256.
func (s *SessionService) HashSessionToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
// RotateSessionID forces session ID rotation to prevent session fixation attacks.
func (s *SessionService) RotateSessionID(session sessions.Session) {
if inner, ok := session.(interface{ Session() *gsessions.Session }); ok {
if sess := inner.Session(); sess != nil {
sess.ID = ""
}
}
}
// CalculateSessionMaxAge dynamically calculates max age and whether it's a browser-session cookie.
func (s *SessionService) CalculateSessionMaxAge(ctx context.Context) (int, bool) {
cfg := s.Config()
maxAge := cfg.SessionAge
isSessionCookie := false
if s.dao != nil {
val, err := s.dao.GetSystemConfigValue(ctx, "login_session_ttl_hours")
if err == nil && val != "" {
if ttlHours, err := strconv.Atoi(val); err == nil {
switch {
case ttlHours == -1:
// 永不过期,设置为 10 年
maxAge = 10 * 365 * 24 * 3600
case ttlHours > 0:
maxAge = ttlHours * 3600
case ttlHours == 0:
isSessionCookie = true
}
}
}
}
return maxAge, isSessionCookie
}
// ApplyLoginSession writes the authenticated user into a freshly rotated session.
func (s *SessionService) ApplyLoginSession(ctx context.Context, session sessions.Session, user *contracts.UserDTO, extras ...map[string]any) (bool, error) {
session.Clear()
s.RotateSessionID(session)
session.Set(consts.UserIDKey, strconv.FormatUint(user.ID, 10))
session.Set(consts.UserNameKey, user.Username)
if len(extras) > 0 {
for key, value := range extras[0] {
session.Set(key, value)
}
}
maxAge, isSessionCookie := s.CalculateSessionMaxAge(ctx)
session.Options(s.GetSessionOptions(maxAge))
if err := session.Save(); err != nil {
return false, err
}
return isSessionCookie, nil
}
-169
View File
@@ -1,169 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"context"
"crypto/sha256"
"encoding/hex"
"net/http"
"strconv"
"strings"
"sync"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
gsessions "github.com/gorilla/sessions"
)
var (
sessConfigMu sync.RWMutex
sessConfig = SessionConfig{
SessionCookieName: "wavelet_session",
SessionAge: 86400,
SessionHTTPOnly: true,
}
)
// SetSessionConfig updates the active session configuration.
func SetSessionConfig(cfg SessionConfig) {
sessConfigMu.Lock()
defer sessConfigMu.Unlock()
sessConfig = cfg
}
// GetSessionConfig returns the active session configuration.
func GetSessionConfig() SessionConfig {
sessConfigMu.RLock()
defer sessConfigMu.RUnlock()
return sessConfig
}
// GetSessionOptions 根据配置构建 Session 选项
func GetSessionOptions(maxAge int) sessions.Options {
cfg := GetSessionConfig()
return sessions.Options{
Path: "/",
Domain: cfg.SessionDomain,
MaxAge: maxAge,
HttpOnly: cfg.SessionHTTPOnly,
Secure: cfg.SessionSecure,
SameSite: http.SameSiteLaxMode,
}
}
// StripCookieMaxAgeAndExpires 从 Set-Cookie 响应头中移除 Max-Age 和 Expires,从而使其成为浏览器会话 Cookie
func StripCookieMaxAgeAndExpires(header http.Header, cookieName string) {
headers := header["Set-Cookie"]
if len(headers) == 0 {
return
}
newHeaders := make([]string, 0, len(headers))
for _, h := range headers {
if strings.HasPrefix(h, cookieName+"=") {
parts := strings.Split(h, ";")
newParts := make([]string, 0, len(parts))
for _, p := range parts {
trimmed := strings.TrimSpace(p)
lower := strings.ToLower(trimmed)
if strings.HasPrefix(lower, "max-age=") || strings.HasPrefix(lower, "expires=") {
continue
}
newParts = append(newParts, p)
}
newHeaders = append(newHeaders, strings.Join(newParts, ";"))
} else {
newHeaders = append(newHeaders, h)
}
}
header["Set-Cookie"] = newHeaders
}
// GetUserIDFromSession 从 Session 中提取用户 ID
func GetUserIDFromSession(s sessions.Session) uint64 {
val := s.Get(UserIDKey)
return ParseUserID(val)
}
// GetUserIDFromContext 从 Gin Context 的 Session 中提取用户 ID
func GetUserIDFromContext(c *gin.Context) (uid uint64) {
defer func() {
_ = recover()
}()
session := sessions.Default(c)
return GetUserIDFromSession(session)
}
func ensureSessionToken(s sessions.Session) (string, bool) {
token, ok := s.Get(SessionTokenKey).(string)
if !ok || token == "" {
token = uuid.NewString()
s.Set(SessionTokenKey, token)
return token, true
}
return token, false
}
func hashSessionToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
func rotateSessionID(s sessions.Session) {
if inner, ok := s.(interface{ Session() *gsessions.Session }); ok {
if sess := inner.Session(); sess != nil {
sess.ID = ""
}
}
}
// SetLoginSession writes the authenticated user into a freshly rotated session.
func SetLoginSession(ctx context.Context, c *gin.Context, user *contracts.UserDTO, extras ...map[string]any) error {
session := sessions.Default(c)
session.Clear()
rotateSessionID(session)
session.Set(UserIDKey, strconv.FormatUint(user.ID, 10))
session.Set(UserNameKey, user.Username)
if len(extras) > 0 {
for key, value := range extras[0] {
session.Set(key, value)
}
}
// 根据系统配置动态设置 Session 过期时间
cfg := GetSessionConfig()
maxAge := cfg.SessionAge
isSessionCookie := false
val, err := GetSystemConfigValue(ctx, "login_session_ttl_hours")
if err == nil && val != "" {
if ttlHours, err := strconv.Atoi(val); err == nil {
switch {
case ttlHours == -1:
// 永不过期,设置为 10 年
maxAge = 10 * 365 * 24 * 3600
case ttlHours > 0:
maxAge = ttlHours * 3600
case ttlHours == 0:
isSessionCookie = true
}
}
}
session.Options(GetSessionOptions(maxAge))
if err := session.Save(); err != nil {
return err
}
if isSessionCookie {
StripCookieMaxAgeAndExpires(c.Writer.Header(), cfg.SessionCookieName)
}
return nil
}