diff --git a/backend/docs/docs.go b/backend/docs/docs.go index 38707569..b037db64 100644 --- a/backend/docs/docs.go +++ b/backend/docs/docs.go @@ -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": { diff --git a/backend/docs/swagger.json b/backend/docs/swagger.json index b5b17238..e073d072 100644 --- a/backend/docs/swagger.json +++ b/backend/docs/swagger.json @@ -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": { diff --git a/backend/docs/swagger.yaml b/backend/docs/swagger.yaml index 14d39049..db617aee 100644 --- a/backend/docs/swagger.yaml +++ b/backend/docs/swagger.yaml @@ -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: 未登录 diff --git a/backend/plugins/domain/auth/auth_source_resolver.go b/backend/plugins/domain/auth/auth_source_resolver.go deleted file mode 100644 index 4d1c3992..00000000 --- a/backend/plugins/domain/auth/auth_source_resolver.go +++ /dev/null @@ -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 -} diff --git a/backend/plugins/domain/auth/cap_errs.go b/backend/plugins/domain/auth/cap_errs.go deleted file mode 100644 index 1f47bfb9..00000000 --- a/backend/plugins/domain/auth/cap_errs.go +++ /dev/null @@ -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 -) diff --git a/backend/plugins/domain/auth/cap_middleware.go b/backend/plugins/domain/auth/cap_middleware.go deleted file mode 100644 index 1c6c7c7a..00000000 --- a/backend/plugins/domain/auth/cap_middleware.go +++ /dev/null @@ -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() - } -} diff --git a/backend/plugins/domain/auth/cap_runtime_settings.go b/backend/plugins/domain/auth/cap_runtime_settings.go deleted file mode 100644 index 411ce58c..00000000 --- a/backend/plugins/domain/auth/cap_runtime_settings.go +++ /dev/null @@ -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 -} diff --git a/backend/plugins/domain/auth/config.go b/backend/plugins/domain/auth/config.go index c01b416e..dd77ffd4 100644 --- a/backend/plugins/domain/auth/config.go +++ b/backend/plugins/domain/auth/config.go @@ -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 diff --git a/backend/plugins/domain/auth/consts/cap.go b/backend/plugins/domain/auth/consts/cap.go new file mode 100644 index 00000000..8981941d --- /dev/null +++ b/backend/plugins/domain/auth/consts/cap.go @@ -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 +) diff --git a/backend/plugins/domain/auth/constants.go b/backend/plugins/domain/auth/consts/consts.go similarity index 71% rename from backend/plugins/domain/auth/constants.go rename to backend/plugins/domain/auth/consts/consts.go index f49dd81f..be0c94a7 100644 --- a/backend/plugins/domain/auth/constants.go +++ b/backend/plugins/domain/auth/consts/consts.go @@ -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 +) diff --git a/backend/plugins/domain/auth/consts/errs.go b/backend/plugins/domain/auth/consts/errs.go new file mode 100644 index 00000000..cc2637ec --- /dev/null +++ b/backend/plugins/domain/auth/consts/errs.go @@ -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" +) diff --git a/backend/plugins/domain/auth/audit.go b/backend/plugins/domain/auth/controller/audit.go similarity index 83% rename from backend/plugins/domain/auth/audit.go rename to backend/plugins/domain/auth/controller/audit.go index f5755340..3ed895fa 100644 --- a/backend/plugins/domain/auth/audit.go +++ b/backend/plugins/domain/auth/controller/audit.go @@ -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(), diff --git a/backend/plugins/domain/auth/cap_handlers.go b/backend/plugins/domain/auth/controller/cap.go similarity index 50% rename from backend/plugins/domain/auth/cap_handlers.go rename to backend/plugins/domain/auth/controller/cap.go index d95987bd..925d86bb 100644 --- a/backend/plugins/domain/auth/cap_handlers.go +++ b/backend/plugins/domain/auth/controller/cap.go @@ -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 } diff --git a/backend/plugins/domain/auth/controller/cap_middleware.go b/backend/plugins/domain/auth/controller/cap_middleware.go new file mode 100644 index 00000000..810b9c7d --- /dev/null +++ b/backend/plugins/domain/auth/controller/cap_middleware.go @@ -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() + } +} diff --git a/backend/plugins/domain/auth/controller/controller.go b/backend/plugins/domain/auth/controller/controller.go new file mode 100644 index 00000000..3f01731c --- /dev/null +++ b/backend/plugins/domain/auth/controller/controller.go @@ -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) + } +} diff --git a/backend/plugins/domain/auth/controller/middleware.go b/backend/plugins/domain/auth/controller/middleware.go new file mode 100644 index 00000000..ea766950 --- /dev/null +++ b/backend/plugins/domain/auth/controller/middleware.go @@ -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() + } +} diff --git a/backend/plugins/domain/auth/controller/oauth.go b/backend/plugins/domain/auth/controller/oauth.go new file mode 100644 index 00000000..47647e74 --- /dev/null +++ b/backend/plugins/domain/auth/controller/oauth.go @@ -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()) +} diff --git a/backend/plugins/domain/auth/controller/user_info.go b/backend/plugins/domain/auth/controller/user_info.go new file mode 100644 index 00000000..62a35ac7 --- /dev/null +++ b/backend/plugins/domain/auth/controller/user_info.go @@ -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)), + ) +} diff --git a/backend/plugins/domain/auth/dao/auth_source.go b/backend/plugins/domain/auth/dao/auth_source.go new file mode 100644 index 00000000..0e7aff6a --- /dev/null +++ b/backend/plugins/domain/auth/dao/auth_source.go @@ -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 +} diff --git a/backend/plugins/domain/auth/cache.go b/backend/plugins/domain/auth/dao/cache.go similarity index 53% rename from backend/plugins/domain/auth/cache.go rename to backend/plugins/domain/auth/dao/cache.go index 0beda8ec..a2476651 100644 --- a/backend/plugins/domain/auth/cache.go +++ b/backend/plugins/domain/auth/dao/cache.go @@ -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() } diff --git a/backend/plugins/domain/auth/dao/dao.go b/backend/plugins/domain/auth/dao/dao.go new file mode 100644 index 00000000..9f09c43c --- /dev/null +++ b/backend/plugins/domain/auth/dao/dao.go @@ -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 +} diff --git a/backend/plugins/domain/auth/dao/external_account.go b/backend/plugins/domain/auth/dao/external_account.go new file mode 100644 index 00000000..d6b23236 --- /dev/null +++ b/backend/plugins/domain/auth/dao/external_account.go @@ -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 +} diff --git a/backend/plugins/domain/auth/dao/user_bridge.go b/backend/plugins/domain/auth/dao/user_bridge.go new file mode 100644 index 00000000..cb69c139 --- /dev/null +++ b/backend/plugins/domain/auth/dao/user_bridge.go @@ -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 +} diff --git a/backend/plugins/domain/auth/errs.go b/backend/plugins/domain/auth/errs.go deleted file mode 100644 index 08a4f0cb..00000000 --- a/backend/plugins/domain/auth/errs.go +++ /dev/null @@ -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" -) diff --git a/backend/plugins/domain/auth/facade.go b/backend/plugins/domain/auth/facade.go new file mode 100644 index 00000000..7c772892 --- /dev/null +++ b/backend/plugins/domain/auth/facade.go @@ -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) +} diff --git a/backend/plugins/domain/auth/handlers.go b/backend/plugins/domain/auth/handlers.go deleted file mode 100644 index 8f30cf67..00000000 --- a/backend/plugins/domain/auth/handlers.go +++ /dev/null @@ -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()) -} diff --git a/backend/plugins/domain/auth/middleware.go b/backend/plugins/domain/auth/middleware.go deleted file mode 100644 index 7f2d575a..00000000 --- a/backend/plugins/domain/auth/middleware.go +++ /dev/null @@ -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() - } -} diff --git a/backend/plugins/domain/auth/model/do/cached_token.go b/backend/plugins/domain/auth/model/do/cached_token.go new file mode 100644 index 00000000..cf497cec --- /dev/null +++ b/backend/plugins/domain/auth/model/do/cached_token.go @@ -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"` +} diff --git a/backend/plugins/domain/auth/model/do/cap_settings.go b/backend/plugins/domain/auth/model/do/cap_settings.go new file mode 100644 index 00000000..f82f852f --- /dev/null +++ b/backend/plugins/domain/auth/model/do/cap_settings.go @@ -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 +} diff --git a/backend/plugins/domain/auth/model/do/cap_settings_test.go b/backend/plugins/domain/auth/model/do/cap_settings_test.go new file mode 100644 index 00000000..a2692b0c --- /dev/null +++ b/backend/plugins/domain/auth/model/do/cap_settings_test.go @@ -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) + }) +} diff --git a/backend/plugins/domain/auth/model/do/oauth_state.go b/backend/plugins/domain/auth/model/do/oauth_state.go new file mode 100644 index 00000000..3cb70029 --- /dev/null +++ b/backend/plugins/domain/auth/model/do/oauth_state.go @@ -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 +} diff --git a/backend/plugins/domain/auth/model/do/oauth_state_test.go b/backend/plugins/domain/auth/model/do/oauth_state_test.go new file mode 100644 index 00000000..bc480356 --- /dev/null +++ b/backend/plugins/domain/auth/model/do/oauth_state_test.go @@ -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) +} diff --git a/backend/plugins/domain/auth/model/dto/auth_source.go b/backend/plugins/domain/auth/model/dto/auth_source.go new file mode 100644 index 00000000..f92dbc82 --- /dev/null +++ b/backend/plugins/domain/auth/model/dto/auth_source.go @@ -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"` +} diff --git a/backend/plugins/domain/auth/cap_models.go b/backend/plugins/domain/auth/model/dto/cap.go similarity index 65% rename from backend/plugins/domain/auth/cap_models.go rename to backend/plugins/domain/auth/model/dto/cap.go index 4a21d0f0..450dc501 100644 --- a/backend/plugins/domain/auth/cap_models.go +++ b/backend/plugins/domain/auth/model/dto/cap.go @@ -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"` -} diff --git a/backend/plugins/domain/auth/model/dto/user_info.go b/backend/plugins/domain/auth/model/dto/user_info.go new file mode 100644 index 00000000..795ee22e --- /dev/null +++ b/backend/plugins/domain/auth/model/dto/user_info.go @@ -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 +} diff --git a/backend/plugins/domain/auth/model/entity/auth_source.go b/backend/plugins/domain/auth/model/entity/auth_source.go new file mode 100644 index 00000000..9ebd2f11 --- /dev/null +++ b/backend/plugins/domain/auth/model/entity/auth_source.go @@ -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 = "" +} diff --git a/backend/plugins/domain/auth/model/entity/auth_source_test.go b/backend/plugins/domain/auth/model/entity/auth_source_test.go new file mode 100644 index 00000000..bde666db --- /dev/null +++ b/backend/plugins/domain/auth/model/entity/auth_source_test.go @@ -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()) + }) +} diff --git a/backend/plugins/domain/auth/model/entity/external_account.go b/backend/plugins/domain/auth/model/entity/external_account.go new file mode 100644 index 00000000..66d1b967 --- /dev/null +++ b/backend/plugins/domain/auth/model/entity/external_account.go @@ -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" +} diff --git a/backend/plugins/domain/auth/models.go b/backend/plugins/domain/auth/models.go deleted file mode 100644 index 00bb87aa..00000000 --- a/backend/plugins/domain/auth/models.go +++ /dev/null @@ -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 -} diff --git a/backend/plugins/domain/auth/plugin.go b/backend/plugins/domain/auth/plugin.go index 8a1c81bd..913665b2 100644 --- a/backend/plugins/domain/auth/plugin.go +++ b/backend/plugins/domain/auth/plugin.go @@ -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 diff --git a/backend/plugins/domain/auth/repository.go b/backend/plugins/domain/auth/repository.go deleted file mode 100644 index 1310085f..00000000 --- a/backend/plugins/domain/auth/repository.go +++ /dev/null @@ -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 -} diff --git a/backend/plugins/domain/auth/service.go b/backend/plugins/domain/auth/service.go deleted file mode 100644 index e26b488f..00000000 --- a/backend/plugins/domain/auth/service.go +++ /dev/null @@ -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 -} diff --git a/backend/plugins/domain/auth/service/auth_registry.go b/backend/plugins/domain/auth/service/auth_registry.go new file mode 100644 index 00000000..77a4f769 --- /dev/null +++ b/backend/plugins/domain/auth/service/auth_registry.go @@ -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 +} diff --git a/backend/plugins/domain/auth/service/auth_service.go b/backend/plugins/domain/auth/service/auth_service.go new file mode 100644 index 00000000..202d480b --- /dev/null +++ b/backend/plugins/domain/auth/service/auth_service.go @@ -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, + } +} diff --git a/backend/plugins/domain/auth/cap_service.go b/backend/plugins/domain/auth/service/cap_service.go similarity index 50% rename from backend/plugins/domain/auth/cap_service.go rename to backend/plugins/domain/auth/service/cap_service.go index 9e7e0a41..210b785f 100644 --- a/backend/plugins/domain/auth/cap_service.go +++ b/backend/plugins/domain/auth/service/cap_service.go @@ -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 +} diff --git a/backend/plugins/domain/auth/service/cap_settings.go b/backend/plugins/domain/auth/service/cap_settings.go new file mode 100644 index 00000000..3054a789 --- /dev/null +++ b/backend/plugins/domain/auth/service/cap_settings.go @@ -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 +} diff --git a/backend/plugins/domain/auth/service/oauth_service.go b/backend/plugins/domain/auth/service/oauth_service.go new file mode 100644 index 00000000..8eee6021 --- /dev/null +++ b/backend/plugins/domain/auth/service/oauth_service.go @@ -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) +} diff --git a/backend/plugins/domain/auth/provider_cache.go b/backend/plugins/domain/auth/service/oidc_provider_cache.go similarity index 65% rename from backend/plugins/domain/auth/provider_cache.go rename to backend/plugins/domain/auth/service/oidc_provider_cache.go index 86465b7b..ea5cbada 100644 --- a/backend/plugins/domain/auth/provider_cache.go +++ b/backend/plugins/domain/auth/service/oidc_provider_cache.go @@ -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) -} diff --git a/backend/plugins/domain/auth/service/service.go b/backend/plugins/domain/auth/service/service.go new file mode 100644 index 00000000..2af980fc --- /dev/null +++ b/backend/plugins/domain/auth/service/service.go @@ -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, + } +} diff --git a/backend/plugins/domain/auth/service/session_service.go b/backend/plugins/domain/auth/service/session_service.go new file mode 100644 index 00000000..21a7d09d --- /dev/null +++ b/backend/plugins/domain/auth/service/session_service.go @@ -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 +} diff --git a/backend/plugins/domain/auth/session.go b/backend/plugins/domain/auth/session.go deleted file mode 100644 index 007b375c..00000000 --- a/backend/plugins/domain/auth/session.go +++ /dev/null @@ -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 -}