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