mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-03 15:06:36 +08:00
merge(wavelet): sync upstream changes
This commit is contained in:
@@ -21,6 +21,12 @@ import (
|
||||
// SystemConfig aliases model.SystemConfig for external compatibility.
|
||||
type SystemConfig = model.SystemConfig
|
||||
|
||||
// TaskExecution aliases model.TaskExecution for external compatibility.
|
||||
type TaskExecution = model.TaskExecution
|
||||
|
||||
// Schedule aliases model.Schedule for external compatibility.
|
||||
type Schedule = model.Schedule
|
||||
|
||||
//go:embed migrations/*/*.sql
|
||||
var adminMigrations embed.FS
|
||||
|
||||
@@ -132,12 +138,17 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
ctx.Router().RegisterWhitelist("/robots.txt")
|
||||
|
||||
// 2. Register Background Tasks
|
||||
const defaultCleanupRetry = 3
|
||||
ctx.Task().Register(service.LogDBSwitchTask, &service.LogDBSwitchHandler{}, extpoints.WithTaskMeta(service.LogDBSwitchMeta))
|
||||
ctx.Task().Register(service.SystemCleanupTask, &service.SystemCleanupHandler{}, extpoints.WithTaskMeta(service.SystemCleanupMeta), extpoints.WithTaskRetry(defaultCleanupRetry))
|
||||
|
||||
// 2.1 Register Cron Schedule
|
||||
ctx.Schedule().RegisterCron("0 3 * * *", service.SystemCleanupTask, nil)
|
||||
|
||||
// 3. Register Settings Schemas
|
||||
ctx.Settings().Register(extpoints.SettingSchema{
|
||||
Key: "admin.system_cleanup_cron",
|
||||
Default: "0 4 * * *",
|
||||
Default: "0 3 * * *",
|
||||
Description: "Cron expression for nightly system logs and expired tokens cleanup",
|
||||
Type: "string",
|
||||
Category: "maintenance",
|
||||
|
||||
@@ -36,11 +36,13 @@ func TestAdminPluginUnit(t *testing.T) {
|
||||
// Verify tasks
|
||||
_, ok := ctx.Tasks().Get("logs:db_switch")
|
||||
require.True(t, ok)
|
||||
_, ok = ctx.Tasks().Get("system:cleanup")
|
||||
require.True(t, ok)
|
||||
|
||||
// Verify settings
|
||||
setting, ok := ctx.Settings().Get("admin.system_cleanup_cron")
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "0 4 * * *", setting.Default)
|
||||
assert.Equal(t, "0 3 * * *", setting.Default)
|
||||
|
||||
provider, err := core.Inject[contracts.PublicConfigProvider](ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/admin/model"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
// SystemCleanupTask 系统定期垃圾清理任务标识
|
||||
SystemCleanupTask = "system:cleanup"
|
||||
// TaskTypeSystemCleanup 系统定期垃圾清理管理类型
|
||||
TaskTypeSystemCleanup = "system_cleanup"
|
||||
taskQueueDefault = "default"
|
||||
)
|
||||
|
||||
// SystemCleanupMeta describes the system-wide cleanup task metadata.
|
||||
var SystemCleanupMeta = contracts.TaskMetaDTO{
|
||||
Type: TaskTypeSystemCleanup,
|
||||
AsynqTask: SystemCleanupTask,
|
||||
Name: "系统垃圾清理",
|
||||
DisplayName: "系统垃圾清理",
|
||||
Description: "定期清理过期任务执行记录,并通过领域事件广播触发各业务域自治清理(临时文件、历史推送等)",
|
||||
Category: "maintenance",
|
||||
SupportsTime: false,
|
||||
MaxRetry: 3,
|
||||
Queue: taskQueueDefault,
|
||||
Retryable: true,
|
||||
}
|
||||
|
||||
// SystemCleanupHandler handles the system-wide garbage cleanup task.
|
||||
type SystemCleanupHandler struct{}
|
||||
|
||||
// Execute executes system cleanup: clears old task executions and emits EventTopicSystemCleanup.
|
||||
func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*contracts.TaskResultDTO, error) {
|
||||
db := GetDB(ctx)
|
||||
if db == nil {
|
||||
return nil, errors.New("database service not available")
|
||||
}
|
||||
|
||||
// 1. 清理自身域(admin 域)的过期任务执行记录(7 天前)
|
||||
var deletedExecutions int64
|
||||
sevenDaysAgo := time.Now().Add(-7 * 24 * time.Hour)
|
||||
res := db.Where("created_at < ?", sevenDaysAgo).Delete(&model.TaskExecution{})
|
||||
if err := res.Error; err != nil {
|
||||
logger.WarnF(ctx, "清理过期任务执行日志失败: %v", err)
|
||||
} else {
|
||||
deletedExecutions = res.RowsAffected
|
||||
logger.InfoF(ctx, "已清理 7 天前任务执行日志,共 %d 条", deletedExecutions)
|
||||
}
|
||||
|
||||
// 2. 广播 EventTopicSystemCleanup 领域事件,由各业务域插件(upload, msg_gateway, user 等)自治执行各自的清理逻辑
|
||||
nowStr := time.Now().Format(time.RFC3339)
|
||||
if err := EmitEvent(ctx, contracts.EventTopicSystemCleanup, contracts.SystemCleanupEvent{
|
||||
TriggeredAt: nowStr,
|
||||
}); err != nil {
|
||||
logger.WarnF(ctx, "广播系统清理领域事件失败: %v", err)
|
||||
}
|
||||
|
||||
msg := fmt.Sprintf("系统垃圾清理完成,已清理过期任务执行日志 %d 条,并已广播领域清理事件", deletedExecutions)
|
||||
logger.InfoF(ctx, "%s", msg)
|
||||
return &contracts.TaskResultDTO{Message: msg}, nil
|
||||
}
|
||||
@@ -0,0 +1,78 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service_test
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/admin/model"
|
||||
"Wavelet/plugins/domain/admin/service"
|
||||
"context"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func TestSystemCleanupHandler_Execute(t *testing.T) {
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, sqliteDB.AutoMigrate(&model.TaskExecution{}))
|
||||
|
||||
service.SetDBService(&testDBService{db: sqliteDB})
|
||||
defer service.ResetServices()
|
||||
|
||||
now := time.Now()
|
||||
oldTime := now.Add(-10 * 24 * time.Hour)
|
||||
recentTime := now.Add(-1 * time.Hour)
|
||||
|
||||
// Seed old execution (should be cleaned)
|
||||
oldExec := model.TaskExecution{
|
||||
ID: 1,
|
||||
TaskID: "task-old",
|
||||
TaskType: "sample_task",
|
||||
Status: "success",
|
||||
CreatedAt: oldTime,
|
||||
}
|
||||
require.NoError(t, sqliteDB.Create(&oldExec).Error)
|
||||
|
||||
// Seed recent execution (should remain)
|
||||
recentExec := model.TaskExecution{
|
||||
ID: 2,
|
||||
TaskID: "task-recent",
|
||||
TaskType: "sample_task",
|
||||
Status: "success",
|
||||
CreatedAt: recentTime,
|
||||
}
|
||||
require.NoError(t, sqliteDB.Create(&recentExec).Error)
|
||||
|
||||
var eventFired atomic.Bool
|
||||
service.SetEventEmitter(func(ctx context.Context, topic string, payload any) error {
|
||||
if topic == contracts.EventTopicSystemCleanup {
|
||||
eventFired.Store(true)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
handler := &service.SystemCleanupHandler{}
|
||||
res, err := handler.Execute(context.Background(), nil)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, res)
|
||||
assert.Contains(t, res.Message, "系统垃圾清理完成")
|
||||
|
||||
// Verify old execution is deleted and recent remains
|
||||
var count int64
|
||||
sqliteDB.Model(&model.TaskExecution{}).Count(&count)
|
||||
assert.Equal(t, int64(1), count)
|
||||
|
||||
var remaining model.TaskExecution
|
||||
sqliteDB.First(&remaining)
|
||||
assert.Equal(t, uint64(2), remaining.ID)
|
||||
|
||||
assert.True(t, eventFired.Load(), "EventTopicSystemCleanup must be emitted")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
+4
-4
@@ -1,7 +1,7 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
package auth
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/response"
|
||||
@@ -14,12 +14,12 @@ import (
|
||||
|
||||
func TestVerifyMiddlewareMissingTokenIsBadRequest(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
restore := InstallTestRuntimeSettings(RuntimeSettings{LoginEnabled: true})
|
||||
restore := InstallCapTestRuntimeSettings(CapRuntimeSettings{LoginEnabled: true})
|
||||
t.Cleanup(restore)
|
||||
|
||||
engine := gin.New()
|
||||
engine.Use(response.ErrorHandlerMiddleware())
|
||||
engine.POST("/register", VerifyMiddleware(GetDefaultManager(), "register"), func(c *gin.Context) {
|
||||
engine.POST("/register", VerifyCaptchaMiddleware(GetDefaultCapManager(), "register"), func(c *gin.Context) {
|
||||
c.Status(http.StatusOK)
|
||||
})
|
||||
|
||||
@@ -27,6 +27,6 @@ func TestVerifyMiddlewareMissingTokenIsBadRequest(t *testing.T) {
|
||||
rec := httptest.NewRecorder()
|
||||
engine.ServeHTTP(rec, req)
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Errorf("VerifyMiddleware() status = %d, want %d", rec.Code, http.StatusBadRequest)
|
||||
t.Errorf("VerifyCaptchaMiddleware() status = %d, want %d", rec.Code, http.StatusBadRequest)
|
||||
}
|
||||
}
|
||||
@@ -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 cap
|
||||
// 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=cap.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 := GetDefaultManager()
|
||||
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=cap.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 := GetDefaultManager()
|
||||
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,572 +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
|
||||
}
|
||||
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"`
|
||||
}
|
||||
+7
-12
@@ -1,22 +1,23 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
// Package dto provides data transfer objects and views for the auth plugin.
|
||||
package dto
|
||||
|
||||
import (
|
||||
"Wavelet/plugins/domain/cap/pow"
|
||||
"Wavelet/plugins/domain/auth/pow"
|
||||
)
|
||||
|
||||
// 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"`
|
||||
}
|
||||
|
||||
// configRecord maps the columns selected from the system config table.
|
||||
type configRecord 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
|
||||
}
|
||||
@@ -0,0 +1,108 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package auth_test
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/plugins/domain/auth"
|
||||
"Wavelet/plugins/infra/cache_memory"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-contrib/sessions/cookie"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestOAuthRateLimiting(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
ctx := core.NewContext(context.Background())
|
||||
ctx.Config().SetSource(core.NewMapSource(nil))
|
||||
require.NoError(t, ctx.Config().Resolve())
|
||||
|
||||
testDB := setupTestDB(t)
|
||||
require.NoError(t, database.New(database.WithDB(testDB)).Apply(ctx))
|
||||
require.NoError(t, cache_memory.New().Apply(ctx))
|
||||
require.NoError(t, auth.New().Apply(ctx))
|
||||
|
||||
// Create an active OIDC source
|
||||
authSrc := auth.AuthSource{
|
||||
ID: 1,
|
||||
Name: "google",
|
||||
Type: "oidc",
|
||||
DisplayName: "Google",
|
||||
ClientID: "client-id-123",
|
||||
ClientSecret: "client-secret-456",
|
||||
OpenIDDiscoveryURL: "https://accounts.google.com",
|
||||
IsActive: true,
|
||||
}
|
||||
require.NoError(t, testDB.Create(&authSrc).Error)
|
||||
|
||||
router := gin.New()
|
||||
router.Use(response.ErrorHandlerMiddleware())
|
||||
store := cookie.NewStore([]byte("test-session-secret-123"))
|
||||
router.Use(sessions.Sessions("wavelet_session_id", store))
|
||||
router.Use(func(c *gin.Context) {
|
||||
c.Request = c.Request.WithContext(core.WithAppContext(c.Request.Context(), ctx.Root()))
|
||||
c.Next()
|
||||
})
|
||||
|
||||
for _, rd := range ctx.Router().Routes() {
|
||||
handlers := make([]gin.HandlerFunc, 0, len(rd.Middlewares)+len(rd.Handlers))
|
||||
for _, m := range rd.Middlewares {
|
||||
if h, ok := m.(gin.HandlerFunc); ok {
|
||||
handlers = append(handlers, h)
|
||||
} else if fn, ok := m.(func(*gin.Context)); ok {
|
||||
handlers = append(handlers, fn)
|
||||
}
|
||||
}
|
||||
for _, raw := range rd.Handlers {
|
||||
if h, ok := raw.(gin.HandlerFunc); ok {
|
||||
handlers = append(handlers, h)
|
||||
} else if fn, ok := raw.(func(*gin.Context)); ok {
|
||||
handlers = append(handlers, fn)
|
||||
}
|
||||
}
|
||||
router.Handle(rd.Method, rd.Path, handlers...)
|
||||
}
|
||||
|
||||
// 10 state slots are allowed per session (oauthStateLimitMax = 10)
|
||||
// We'll simulate 10 requests with the same cookie
|
||||
var cookies []*http.Cookie
|
||||
for i := 1; i <= 10; i++ {
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/oauth/login?source=google", nil)
|
||||
for _, ck := range cookies {
|
||||
req.AddCookie(ck)
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
if len(w.Result().Cookies()) > 0 {
|
||||
cookies = w.Result().Cookies()
|
||||
}
|
||||
}
|
||||
|
||||
// 11th request for the same session should be rate limited
|
||||
{
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/oauth/login?source=google", nil)
|
||||
for _, ck := range cookies {
|
||||
req.AddCookie(ck)
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
|
||||
var resp map[string]any
|
||||
err := json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "请求授权过于频繁,请稍后重试", resp["error_msg"])
|
||||
}
|
||||
}
|
||||
@@ -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,33 +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 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)
|
||||
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)
|
||||
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
|
||||
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)
|
||||
|
||||
// 2.1 Register Public / Auth Whitelist Endpoints
|
||||
// 2. Register Public / Auth Whitelist Endpoints
|
||||
publicEndpoints := []string{
|
||||
"/api/v1/oauth/sources",
|
||||
"/api/v1/oauth/login",
|
||||
@@ -125,49 +155,61 @@ 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)
|
||||
ctrl.RegisterRoutes(ctx.Router())
|
||||
|
||||
// 4. Register Settings Schemas
|
||||
const (
|
||||
settingTypeInteger = "integer"
|
||||
settingCategorySecurity = "security"
|
||||
)
|
||||
|
||||
ctx.Settings().Register(extpoints.SettingSchema{
|
||||
Key: "auth.session_age",
|
||||
Default: 86400 * 7,
|
||||
Description: "Default session lifetime in seconds",
|
||||
Type: "integer",
|
||||
Category: "security",
|
||||
Type: settingTypeInteger,
|
||||
Category: settingCategorySecurity,
|
||||
})
|
||||
ctx.Settings().Register(extpoints.SettingSchema{
|
||||
Key: "auth.login_rate_limit_max_attempts",
|
||||
Default: 5,
|
||||
Description: "Max login failure attempts before temporary IP lock",
|
||||
Type: "integer",
|
||||
Category: "security",
|
||||
Type: settingTypeInteger,
|
||||
Category: settingCategorySecurity,
|
||||
})
|
||||
ctx.Settings().Register(extpoints.SettingSchema{
|
||||
Key: "cap.login_enabled",
|
||||
Default: false,
|
||||
Description: "Whether to require CAPTCHA verification for user login",
|
||||
Type: "boolean",
|
||||
Category: settingCategorySecurity,
|
||||
})
|
||||
ctx.Settings().Register(extpoints.SettingSchema{
|
||||
Key: "cap.challenge_count",
|
||||
Default: 1,
|
||||
Description: "Number of PoW puzzle challenges to solve",
|
||||
Type: settingTypeInteger,
|
||||
Category: settingCategorySecurity,
|
||||
})
|
||||
|
||||
// 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) {
|
||||
svc.CapSettings.Invalidate()
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -62,6 +62,13 @@ func hashToken(token string) string {
|
||||
return hex.EncodeToString(h.Sum(nil))
|
||||
}
|
||||
|
||||
type testSystemConfig struct {
|
||||
Key string `gorm:"primaryKey"`
|
||||
Value string
|
||||
}
|
||||
|
||||
func (testSystemConfig) TableName() string { return "w_system_configs" }
|
||||
|
||||
func setupTestDB(t *testing.T) *gorm.DB {
|
||||
t.Helper()
|
||||
dbPath := filepath.Join(t.TempDir(), "auth_test.db")
|
||||
@@ -73,6 +80,7 @@ func setupTestDB(t *testing.T) *gorm.DB {
|
||||
&testAccessToken{},
|
||||
&auth.AuthSource{},
|
||||
&auth.ExternalAccount{},
|
||||
&testSystemConfig{},
|
||||
))
|
||||
|
||||
return testDB
|
||||
@@ -153,4 +161,25 @@ func TestAuthPluginUnit(t *testing.T) {
|
||||
current, err := authSvc.GetCurrentUser(userCtx)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, user.ID, current.ID)
|
||||
|
||||
// Test CaptchaService injection
|
||||
capSvc, err := core.Inject[contracts.CaptchaService](ctx)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, capSvc)
|
||||
assert.NotNil(t, capSvc.ChallengeHandler())
|
||||
assert.NotNil(t, capSvc.RedeemHandler())
|
||||
assert.NotNil(t, capSvc.VerifyMiddleware("login"))
|
||||
|
||||
// Verify CAPTCHA routes registered
|
||||
var foundChallenge, foundRedeem bool
|
||||
for _, rd := range ctx.Router().Routes() {
|
||||
if rd.Path == "/api/v1/cap/challenge" {
|
||||
foundChallenge = true
|
||||
}
|
||||
if rd.Path == "/api/v1/cap/redeem" {
|
||||
foundRedeem = true
|
||||
}
|
||||
}
|
||||
assert.True(t, foundChallenge, "expected /api/v1/cap/challenge route")
|
||||
assert.True(t, foundRedeem, "expected /api/v1/cap/redeem route")
|
||||
}
|
||||
|
||||
@@ -1,211 +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
|
||||
)
|
||||
|
||||
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 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
|
||||
}
|
||||
|
||||
// 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,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,194 @@
|
||||
// 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/model/dto"
|
||||
"Wavelet/plugins/domain/auth/pow"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// CaptchaManager orchestrates challenge generation and solution validation.
|
||||
type CaptchaManager struct {
|
||||
secret []byte
|
||||
store pow.Store
|
||||
settingsMgr *CapSettingsManager
|
||||
}
|
||||
|
||||
// NewCaptchaManager creates a new CAPTCHA Manager.
|
||||
func NewCaptchaManager(secret []byte, store pow.Store, settingsMgr *CapSettingsManager) *CaptchaManager {
|
||||
return &CaptchaManager{
|
||||
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 := m.settingsMgr.Current(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
challengeConfig := pow.ChallengeConfig{
|
||||
Count: settings.ChallengeCount,
|
||||
Size: settings.ChallengeSize,
|
||||
Difficulty: settings.ChallengeDifficulty,
|
||||
Expires: settings.ChallengeTTL,
|
||||
}
|
||||
return pow.GenerateChallenge(m.secret, challengeConfig, scope)
|
||||
}
|
||||
|
||||
// Redeem verifies PoW solutions and returns a one-time redeem token.
|
||||
func (m *CaptchaManager) Redeem(ctx context.Context, token string, solutions []int, scope string) (*dto.RedeemResponse, error) {
|
||||
sigHex := pow.JwtSigHex(token)
|
||||
if sigHex == "" {
|
||||
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 &dto.RedeemResponse{Success: false, Error: err.Error()}, nil //nolint:nilerr // validation errors returned as response
|
||||
}
|
||||
|
||||
now := time.Now().UnixNano() / int64(time.Millisecond)
|
||||
nonceTTL := time.Duration(payload.Expires-now) * time.Millisecond
|
||||
if nonceTTL < time.Second {
|
||||
nonceTTL = time.Second
|
||||
}
|
||||
|
||||
set, err := m.store.SetNX(ctx, nonceKey, "1", nonceTTL)
|
||||
if err != nil {
|
||||
return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrNonceStoreFailed}, err
|
||||
}
|
||||
if !set {
|
||||
return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrAlreadyRedeemed}, nil
|
||||
}
|
||||
|
||||
settings, err := m.settingsMgr.Current(ctx)
|
||||
if err != nil {
|
||||
return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrSettingsLoad}, err
|
||||
}
|
||||
|
||||
id := pow.RandomHex(consts.RedeemTokenIDLength)
|
||||
verToken := pow.RandomHex(consts.RedeemVerTokenLength)
|
||||
verHashBytes := sha256.Sum256([]byte(verToken))
|
||||
verHashHex := hex.EncodeToString(verHashBytes[:])
|
||||
|
||||
tokenKey := "cap:token:" + id + ":" + verHashHex
|
||||
tokenExpires := time.Now().Add(settings.TokenTTL)
|
||||
storeVal := strconv.FormatInt(tokenExpires.UnixNano(), 10) + "|" + scope
|
||||
|
||||
if err := m.store.Set(ctx, tokenKey, storeVal, settings.TokenTTL); err != nil {
|
||||
return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrTokenStoreFailed}, err
|
||||
}
|
||||
|
||||
return &dto.RedeemResponse{
|
||||
Success: true,
|
||||
Token: id + ":" + verToken,
|
||||
Expires: tokenExpires.UnixNano() / int64(time.Millisecond),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// VerifyToken validates and consumes the redeem token (single-use).
|
||||
func (m *CaptchaManager) VerifyToken(ctx context.Context, token, expectedScope string) (bool, error) {
|
||||
if token == "" {
|
||||
return false, nil
|
||||
}
|
||||
parts := strings.Split(token, ":")
|
||||
if len(parts) != consts.TokenPartsCount {
|
||||
return false, nil
|
||||
}
|
||||
id := parts[0]
|
||||
verToken := parts[1]
|
||||
|
||||
verHashBytes := sha256.Sum256([]byte(verToken))
|
||||
verHashHex := hex.EncodeToString(verHashBytes[:])
|
||||
|
||||
tokenKey := "cap:token:" + id + ":" + verHashHex
|
||||
|
||||
if m.store == nil {
|
||||
return false, nil
|
||||
}
|
||||
val, exists, err := m.store.GetAndDelete(ctx, tokenKey)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if !exists {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
valParts := strings.Split(val, "|")
|
||||
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 failure
|
||||
}
|
||||
tokenScope := valParts[1]
|
||||
|
||||
if expectedScope != "" && tokenScope != expectedScope {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
if time.Now().UnixNano() > expNano {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// CaptchaServiceImpl implements contracts.CaptchaService.
|
||||
type CaptchaServiceImpl struct {
|
||||
manager *CaptchaManager
|
||||
verifyMiddleware func(scope string) any
|
||||
challengeHandler any
|
||||
redeemHandler any
|
||||
}
|
||||
|
||||
// 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,
|
||||
}
|
||||
}
|
||||
|
||||
// VerifyMiddleware returns the captcha verification middleware.
|
||||
func (s *CaptchaServiceImpl) VerifyMiddleware(scope string) any {
|
||||
if s.verifyMiddleware != nil {
|
||||
return s.verifyMiddleware(scope)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ChallengeHandler returns the challenge HTTP handler.
|
||||
func (s *CaptchaServiceImpl) ChallengeHandler() any {
|
||||
return s.challengeHandler
|
||||
}
|
||||
|
||||
// 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
|
||||
}
|
||||
@@ -19,7 +19,7 @@ import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/core/extpoints"
|
||||
"Wavelet/plugins/domain/admin"
|
||||
"Wavelet/plugins/domain/message_gateway"
|
||||
"Wavelet/plugins/domain/msg_gateway"
|
||||
"Wavelet/plugins/domain/user"
|
||||
)
|
||||
|
||||
@@ -122,7 +122,7 @@ func TestRoutesMountedBeforeAuthServiceAreGuarded(t *testing.T) {
|
||||
core.WithPlugins(
|
||||
dbProvider(),
|
||||
user.New(),
|
||||
message_gateway.New(),
|
||||
msg_gateway.New(),
|
||||
authProvider(),
|
||||
),
|
||||
)
|
||||
@@ -147,7 +147,7 @@ func TestAuthConsumersDeclareAuthDependency(t *testing.T) {
|
||||
deps []reflect.Type
|
||||
}{
|
||||
{"user", user.New().Inject()},
|
||||
{"message_gateway", message_gateway.New().Inject()},
|
||||
{"msg_gateway", msg_gateway.New().Inject()},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
assert.Contains(t, tc.deps, want,
|
||||
@@ -194,7 +194,7 @@ func TestAuthGuardFailsClosed(t *testing.T) {
|
||||
prefix string
|
||||
}{
|
||||
{"user", user.New().Apply, http.MethodPost, "/api/v1/user/change-password"},
|
||||
{"message_gateway", message_gateway.New().Apply, http.MethodGet, "/api/v1/message-gateway"},
|
||||
{"msg_gateway", msg_gateway.New().Apply, http.MethodGet, "/api/v1/message-gateway"},
|
||||
{"admin", admin.New().Apply, http.MethodGet, "/api/v1/admin"},
|
||||
}
|
||||
|
||||
|
||||
@@ -1,24 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package cap 提供人机验证中间件
|
||||
package cap
|
||||
|
||||
// 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 cap
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/response"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// VerifyMiddleware returns a Gin middleware that checks and consumes the X-Cap-Token header.
|
||||
func VerifyMiddleware(mgr *Manager, scope string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if !ProtectionEnabled(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,111 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package cap provides the proof-of-work (PoW) CAPTCHA verification domain plugin for Cordis.
|
||||
package cap
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/core/extpoints"
|
||||
"reflect"
|
||||
)
|
||||
|
||||
// Plugin implements core.Plugin to provide CAPTCHA generation, validation, and route protection.
|
||||
type Plugin struct{}
|
||||
|
||||
// New creates a new cap domain plugin.
|
||||
func New() *Plugin {
|
||||
return &Plugin{}
|
||||
}
|
||||
|
||||
// Name returns the unique identifier for the cap domain plugin.
|
||||
func (p *Plugin) Name() string {
|
||||
return "cap"
|
||||
}
|
||||
|
||||
// Inject declares required dependencies for the cap domain plugin.
|
||||
func (p *Plugin) Inject() []reflect.Type {
|
||||
return []reflect.Type{
|
||||
reflect.TypeFor[contracts.DBService](),
|
||||
}
|
||||
}
|
||||
|
||||
// Manifest returns the plugin metadata.
|
||||
func (p *Plugin) Manifest() core.Manifest {
|
||||
return core.Manifest{
|
||||
Name: "cap",
|
||||
Version: "1.0.0",
|
||||
Description: "Proof-of-work CAPTCHA challenge and verification domain plugin",
|
||||
Author: "Wavelet Team",
|
||||
}
|
||||
}
|
||||
|
||||
type capAppConfig struct {
|
||||
SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"`
|
||||
}
|
||||
|
||||
// DeclareConfig declares configuration bindings for the cap plugin.
|
||||
func (p *Plugin) DeclareConfig() []core.ConfigBinding {
|
||||
return []core.ConfigBinding{
|
||||
{Prefix: "app", Target: &capAppConfig{}},
|
||||
}
|
||||
}
|
||||
|
||||
// Apply registers the cap routes and settings into the Context.
|
||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
var cfg capAppConfig
|
||||
if err := ctx.Config().Bind("app", &cfg); err == nil && cfg.SessionSecret != "" {
|
||||
SetSecret([]byte(cfg.SessionSecret))
|
||||
}
|
||||
|
||||
core.Bind[contracts.DBService](ctx, setDBService)
|
||||
ctx.OnDispose(func() error {
|
||||
setDBService(nil)
|
||||
return nil
|
||||
})
|
||||
|
||||
// Listen to system config changed events to invalidate cached settings
|
||||
ctx.Events().On(contracts.EventTopicConfigChanged, func(_ any) {
|
||||
InvalidateRuntimeSettings()
|
||||
})
|
||||
|
||||
core.Provide[contracts.CaptchaService](ctx, captchaService{})
|
||||
|
||||
// Register HTTP Routes
|
||||
capGroup := ctx.Router().Group("/api/v1/cap")
|
||||
{
|
||||
capGroup.GET("/challenge", Challenge)
|
||||
capGroup.POST("/challenge", Challenge)
|
||||
capGroup.POST("/redeem", Redeem)
|
||||
}
|
||||
ctx.Router().RegisterWhitelist("/api/v1/cap/challenge", "/api/v1/cap/redeem")
|
||||
|
||||
// Register Settings Schemas
|
||||
ctx.Settings().Register(extpoints.SettingSchema{
|
||||
Key: "cap.login_enabled",
|
||||
Default: false,
|
||||
Description: "Whether to require CAPTCHA verification for user login",
|
||||
Type: "boolean",
|
||||
Category: "security",
|
||||
})
|
||||
ctx.Settings().Register(extpoints.SettingSchema{
|
||||
Key: "cap.challenge_count",
|
||||
Default: 1,
|
||||
Description: "Number of PoW puzzle challenges to solve",
|
||||
Type: "integer",
|
||||
Category: "security",
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
type captchaService struct{}
|
||||
|
||||
func (captchaService) VerifyMiddleware(scope string) any {
|
||||
return VerifyMiddleware(GetDefaultManager(), scope)
|
||||
}
|
||||
|
||||
func (captchaService) ChallengeHandler() any { return Challenge }
|
||||
|
||||
func (captchaService) RedeemHandler() any { return Redeem }
|
||||
@@ -1,55 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
)
|
||||
|
||||
func TestApplyProvidesCaptchaService(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
if err := New().Apply(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
svc, err := core.Inject[contracts.CaptchaService](ctx)
|
||||
if err != nil || svc == nil {
|
||||
t.Fatalf("Inject CaptchaService: svc=%v err=%v", svc, err)
|
||||
}
|
||||
if svc.ChallengeHandler() == nil || svc.RedeemHandler() == nil {
|
||||
t.Fatal("handlers must be non-nil")
|
||||
}
|
||||
if svc.VerifyMiddleware("login") == nil {
|
||||
t.Fatal("VerifyMiddleware(login) must be non-nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyRegistersUnversionedCapRoutes(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
if err := New().Apply(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want := map[string]bool{
|
||||
"GET /api/v1/cap/challenge": false,
|
||||
"POST /api/v1/cap/challenge": false,
|
||||
"POST /api/v1/cap/redeem": false,
|
||||
}
|
||||
for _, rd := range ctx.Router().Routes() {
|
||||
key := rd.Method + " " + rd.Path
|
||||
if _, ok := want[key]; ok {
|
||||
want[key] = true
|
||||
}
|
||||
if key == "POST /api/cap/challenge" || key == "POST /api/cap/redeem" {
|
||||
t.Errorf("legacy route must not exist: %s", key)
|
||||
}
|
||||
}
|
||||
for key, ok := range want {
|
||||
if !ok {
|
||||
t.Errorf("missing route %s", key)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,56 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"context"
|
||||
"sync"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var (
|
||||
dbMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
)
|
||||
|
||||
// setDBService caches the DBService contract used by the persistence layer.
|
||||
func setDBService(s contracts.DBService) {
|
||||
dbMu.Lock()
|
||||
defer dbMu.Unlock()
|
||||
dbSvc = s
|
||||
}
|
||||
|
||||
// getDB resolves a GORM handle from the request/app context, then the Bind fallback.
|
||||
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
|
||||
}
|
||||
|
||||
// loadRuntimeSettings reads the CAPTCHA owned rows from the system config table.
|
||||
func loadRuntimeSettings(ctx context.Context) (RuntimeSettings, error) {
|
||||
var records []configRecord
|
||||
db := getDB(ctx)
|
||||
if db == nil {
|
||||
return parseRuntimeSettings(nil), nil
|
||||
}
|
||||
if err := db.Table("w_system_configs").Where("key IN ?", runtimeConfigKeys).Find(&records).Error; err != nil {
|
||||
return RuntimeSettings{}, err
|
||||
}
|
||||
configs := make(map[string]string, len(records))
|
||||
for _, r := range records {
|
||||
configs[r.Key] = r.Value
|
||||
}
|
||||
return parseRuntimeSettings(configs), nil
|
||||
}
|
||||
@@ -1,185 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strconv"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"golang.org/x/sync/singleflight"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultChallengeCount = 1
|
||||
defaultChallengeSize = 32
|
||||
defaultChallengeDifficulty = 4
|
||||
defaultChallengeTTL = 10 * time.Minute
|
||||
defaultTokenTTL = 20 * time.Minute
|
||||
)
|
||||
|
||||
// RuntimeSettings is the parsed CAPTCHA runtime configuration loaded from system_configs.
|
||||
type RuntimeSettings 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 runtimeConfigKeys = []string{
|
||||
ConfigKeyCapLoginEnabled,
|
||||
ConfigKeyCapChallengeCount,
|
||||
ConfigKeyCapChallengeSize,
|
||||
ConfigKeyCapChallengeDifficulty,
|
||||
ConfigKeyCapChallengeTTL,
|
||||
ConfigKeyCapTokenTTL,
|
||||
}
|
||||
|
||||
var runtimeConfigKeySet = func() map[string]struct{} {
|
||||
set := make(map[string]struct{}, len(runtimeConfigKeys))
|
||||
for _, key := range runtimeConfigKeys {
|
||||
set[key] = struct{}{}
|
||||
}
|
||||
return set
|
||||
}()
|
||||
|
||||
type runtimeSettingsStore struct {
|
||||
snapshot atomic.Pointer[RuntimeSettings]
|
||||
loadGroup singleflight.Group
|
||||
}
|
||||
|
||||
var settingsStore = &runtimeSettingsStore{}
|
||||
|
||||
// IsRuntimeConfigKey reports whether a system config key affects CAPTCHA runtime settings.
|
||||
func IsRuntimeConfigKey(key string) bool {
|
||||
_, ok := runtimeConfigKeySet[key]
|
||||
return ok
|
||||
}
|
||||
|
||||
// CurrentSettings returns the cached CAPTCHA runtime settings snapshot.
|
||||
func CurrentSettings(ctx context.Context) (RuntimeSettings, error) {
|
||||
return settingsStore.current(ctx)
|
||||
}
|
||||
|
||||
// ProtectionEnabled reports whether CAPTCHA verification is required for protected routes.
|
||||
func ProtectionEnabled(ctx context.Context) bool {
|
||||
settings, err := CurrentSettings(ctx)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return settings.LoginEnabled
|
||||
}
|
||||
|
||||
// InvalidateRuntimeSettings drops the in-process CAPTCHA settings snapshot.
|
||||
func InvalidateRuntimeSettings() {
|
||||
settingsStore.snapshot.Store(nil)
|
||||
}
|
||||
|
||||
// ResetRuntimeSettingsForTest clears the CAPTCHA runtime snapshot.
|
||||
func ResetRuntimeSettingsForTest() {
|
||||
InvalidateRuntimeSettings()
|
||||
}
|
||||
|
||||
// InstallTestRuntimeSettings installs a fixed snapshot for unit tests.
|
||||
func InstallTestRuntimeSettings(settings RuntimeSettings) func() {
|
||||
snapshot := settings
|
||||
settingsStore.snapshot.Store(&snapshot)
|
||||
return InvalidateRuntimeSettings
|
||||
}
|
||||
|
||||
func (s *runtimeSettingsStore) current(ctx context.Context) (RuntimeSettings, error) {
|
||||
s.ensureInvalidationListener()
|
||||
|
||||
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 := loadRuntimeSettings(ctx)
|
||||
if loadErr != nil {
|
||||
return RuntimeSettings{}, loadErr
|
||||
}
|
||||
|
||||
s.snapshot.Store(&settings)
|
||||
return settings, nil
|
||||
})
|
||||
if err != nil {
|
||||
return RuntimeSettings{}, err
|
||||
}
|
||||
|
||||
settings, ok := loaded.(RuntimeSettings)
|
||||
if !ok {
|
||||
return RuntimeSettings{}, errors.New("cap runtime settings loader returned unexpected type")
|
||||
}
|
||||
return settings, nil
|
||||
}
|
||||
|
||||
func parseRuntimeSettings(configs map[string]string) RuntimeSettings {
|
||||
settings := RuntimeSettings{
|
||||
ChallengeCount: defaultChallengeCount,
|
||||
ChallengeSize: defaultChallengeSize,
|
||||
ChallengeDifficulty: defaultChallengeDifficulty,
|
||||
ChallengeTTL: defaultChallengeTTL,
|
||||
TokenTTL: defaultTokenTTL,
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
func (s *runtimeSettingsStore) ensureInvalidationListener() {}
|
||||
@@ -1,182 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package cap provides CAPTCHA and proof-of-work (PoW) verification services.
|
||||
package cap
|
||||
|
||||
import (
|
||||
"Wavelet/plugins/domain/cap/pow"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
redeemTokenIDLength = 8 // 兑换 Token ID 字节长度
|
||||
redeemVerTokenLength = 15 // 兑换验证 Token 字节长度
|
||||
tokenPartsCount = 2 // 兑换 Token 由两部分组成
|
||||
valuePartsCount = 2 // 存储值由 scope 和过期时间组成
|
||||
)
|
||||
|
||||
// Manager orchestrates challenge generation and solution validation.
|
||||
type Manager struct {
|
||||
secret []byte
|
||||
store pow.Store
|
||||
}
|
||||
|
||||
// NewManager creates a new CAPTCHA Manager.
|
||||
func NewManager(secret []byte, store pow.Store) *Manager {
|
||||
return &Manager{
|
||||
secret: secret,
|
||||
store: store,
|
||||
}
|
||||
}
|
||||
|
||||
// Generate creates a challenge response.
|
||||
func (m *Manager) Generate(ctx context.Context, scope string) (*pow.ChallengeResponse, error) {
|
||||
settings, err := CurrentSettings(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
challengeConfig := pow.ChallengeConfig{
|
||||
Count: settings.ChallengeCount,
|
||||
Size: settings.ChallengeSize,
|
||||
Difficulty: settings.ChallengeDifficulty,
|
||||
Expires: settings.ChallengeTTL,
|
||||
}
|
||||
return pow.GenerateChallenge(m.secret, challengeConfig, scope)
|
||||
}
|
||||
|
||||
// Redeem verifies PoW solutions and returns a one-time redeem token.
|
||||
func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, scope string) (*RedeemResponse, error) {
|
||||
sigHex := pow.JwtSigHex(token)
|
||||
if sigHex == "" {
|
||||
return &RedeemResponse{Success: false, Error: 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
|
||||
}
|
||||
|
||||
now := time.Now().UnixNano() / int64(time.Millisecond)
|
||||
nonceTTL := time.Duration(payload.Expires-now) * time.Millisecond
|
||||
if nonceTTL < time.Second {
|
||||
nonceTTL = time.Second
|
||||
}
|
||||
|
||||
set, err := m.store.SetNX(ctx, nonceKey, "1", nonceTTL)
|
||||
if err != nil {
|
||||
return &RedeemResponse{Success: false, Error: redeemErrNonceStoreFailed}, err
|
||||
}
|
||||
if !set {
|
||||
return &RedeemResponse{Success: false, Error: redeemErrAlreadyRedeemed}, nil
|
||||
}
|
||||
|
||||
settings, err := CurrentSettings(ctx)
|
||||
if err != nil {
|
||||
return &RedeemResponse{Success: false, Error: redeemErrSettingsLoad}, err
|
||||
}
|
||||
|
||||
id := pow.RandomHex(redeemTokenIDLength)
|
||||
verToken := pow.RandomHex(redeemVerTokenLength)
|
||||
verHashBytes := sha256.Sum256([]byte(verToken))
|
||||
verHashHex := hex.EncodeToString(verHashBytes[:])
|
||||
|
||||
tokenKey := "cap:token:" + id + ":" + verHashHex
|
||||
tokenExpires := time.Now().Add(settings.TokenTTL)
|
||||
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 &RedeemResponse{
|
||||
Success: true,
|
||||
Token: id + ":" + verToken,
|
||||
Expires: tokenExpires.UnixNano() / int64(time.Millisecond),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// VerifyToken validates and consumes the redeem token (single-use).
|
||||
func (m *Manager) VerifyToken(ctx context.Context, token, expectedScope string) (bool, error) {
|
||||
if token == "" {
|
||||
return false, nil
|
||||
}
|
||||
parts := strings.Split(token, ":")
|
||||
if len(parts) != tokenPartsCount {
|
||||
return false, nil
|
||||
}
|
||||
id := parts[0]
|
||||
verToken := parts[1]
|
||||
|
||||
verHashBytes := sha256.Sum256([]byte(verToken))
|
||||
verHashHex := hex.EncodeToString(verHashBytes[:])
|
||||
|
||||
tokenKey := "cap:token:" + id + ":" + verHashHex
|
||||
|
||||
val, exists, err := sGetAndDelete(ctx, m.store, tokenKey)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if !exists {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
valParts := strings.Split(val, "|")
|
||||
if len(valParts) != 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
|
||||
}
|
||||
tokenScope := valParts[1]
|
||||
|
||||
if expectedScope != "" && tokenScope != expectedScope {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
if time.Now().UnixNano() > expNano {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
|
||||
var (
|
||||
defaultManagerMu sync.RWMutex
|
||||
defaultManager *Manager
|
||||
)
|
||||
|
||||
// SetSecret sets the shared secret used by the default manager.
|
||||
func SetSecret(secret []byte) {
|
||||
defaultManagerMu.Lock()
|
||||
defer defaultManagerMu.Unlock()
|
||||
if len(secret) > 0 {
|
||||
store := pow.NewMemoryStore(1 * time.Minute)
|
||||
defaultManager = NewManager(secret, store)
|
||||
}
|
||||
}
|
||||
|
||||
// GetDefaultManager yields the global singleton CAPTCHA manager.
|
||||
func GetDefaultManager() *Manager {
|
||||
defaultManagerMu.RLock()
|
||||
defer defaultManagerMu.RUnlock()
|
||||
return defaultManager
|
||||
}
|
||||
@@ -9,9 +9,10 @@ import (
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/plugins/domain/admin"
|
||||
"Wavelet/plugins/domain/auth"
|
||||
"Wavelet/plugins/domain/message_gateway"
|
||||
"Wavelet/plugins/domain/msg_gateway"
|
||||
"Wavelet/plugins/domain/risk_control"
|
||||
"Wavelet/plugins/domain/system"
|
||||
"Wavelet/plugins/domain/upload"
|
||||
"Wavelet/plugins/domain/user"
|
||||
"Wavelet/plugins/infra/cache"
|
||||
"Wavelet/plugins/infra/logger"
|
||||
@@ -23,6 +24,7 @@ import (
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -47,13 +49,16 @@ func setupTestDB(t *testing.T) *gorm.DB {
|
||||
&user.AccessToken{},
|
||||
&auth.AuthSource{},
|
||||
&auth.ExternalAccount{},
|
||||
&message_gateway.MessageChannel{},
|
||||
&message_gateway.MessageBinding{},
|
||||
&message_gateway.MessagePairingCode{},
|
||||
&msg_gateway.MessageChannel{},
|
||||
&msg_gateway.MessageBinding{},
|
||||
&msg_gateway.MessagePairingCode{},
|
||||
&admin.SystemConfig{},
|
||||
&message_gateway.PushChannel{},
|
||||
&message_gateway.PushEvent{},
|
||||
&message_gateway.PushHistory{},
|
||||
&admin.TaskExecution{},
|
||||
&msg_gateway.PushChannel{},
|
||||
&msg_gateway.PushEvent{},
|
||||
&msg_gateway.PushHistory{},
|
||||
&upload.Upload{},
|
||||
&upload.UploadStat{},
|
||||
))
|
||||
|
||||
db.SetDB(testDB)
|
||||
@@ -262,15 +267,15 @@ func TestMessageGatewayPlugin(t *testing.T) {
|
||||
require.NoError(t, cache.New().Apply(ctx))
|
||||
require.NoError(t, logger.New().Apply(ctx))
|
||||
|
||||
p := message_gateway.New()
|
||||
assert.Equal(t, "message_gateway", p.Name())
|
||||
assert.Equal(t, "message_gateway", p.Manifest().Name)
|
||||
p := msg_gateway.New()
|
||||
assert.Equal(t, "msg_gateway", p.Name())
|
||||
assert.Equal(t, "msg_gateway", p.Manifest().Name)
|
||||
require.NoError(t, p.Apply(ctx))
|
||||
|
||||
// 1. Migrations
|
||||
entry, ok := ctx.Migrations().Get("message_gateway")
|
||||
entry, ok := ctx.Migrations().Get("msg_gateway")
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "message_gateway", entry.PluginID)
|
||||
assert.Equal(t, "msg_gateway", entry.PluginID)
|
||||
|
||||
// 2. Routes
|
||||
routes := ctx.Router().Routes()
|
||||
@@ -287,24 +292,24 @@ func TestMessageGatewayPlugin(t *testing.T) {
|
||||
assert.True(t, hasBindings)
|
||||
|
||||
// 3. Tasks & Schedules
|
||||
taskDef, ok := ctx.Tasks().Get("message_gateway:push_notification")
|
||||
taskDef, ok := ctx.Tasks().Get("msg_gateway:push_notification")
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, 3, taskDef.Retry)
|
||||
|
||||
schedDef, ok := ctx.Schedules().Get("message_gateway:cleanup_pairing_codes")
|
||||
schedDef, ok := ctx.Schedules().Get("msg_gateway:cleanup_pairing_codes")
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "*/10 * * * *", schedDef.Spec)
|
||||
|
||||
// 4. EventBus Trigger
|
||||
var receivedEvent message_gateway.PushNotificationEvent
|
||||
var receivedEvent msg_gateway.PushNotificationEvent
|
||||
var eventFired bool
|
||||
ctx.Events().On("notification:push", func(c context.Context, e message_gateway.PushNotificationEvent) error {
|
||||
ctx.Events().On("notification:push", func(c context.Context, e msg_gateway.PushNotificationEvent) error {
|
||||
eventFired = true
|
||||
receivedEvent = e
|
||||
return nil
|
||||
})
|
||||
|
||||
err := ctx.Events().Emit(context.Background(), "notification:push", message_gateway.PushNotificationEvent{
|
||||
err := ctx.Events().Emit(context.Background(), "notification:push", msg_gateway.PushNotificationEvent{
|
||||
UserID: 99,
|
||||
Channel: "telegram",
|
||||
Title: "System Alert",
|
||||
@@ -317,7 +322,7 @@ func TestMessageGatewayPlugin(t *testing.T) {
|
||||
assert.Equal(t, "System Alert", receivedEvent.Title)
|
||||
|
||||
// 5. Settings
|
||||
schema, ok := ctx.Settings().Get("message_gateway.pairing_code_expiry_minutes")
|
||||
schema, ok := ctx.Settings().Get("msg_gateway.pairing_code_expiry_minutes")
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, 15, schema.Default)
|
||||
}
|
||||
@@ -392,7 +397,7 @@ func TestAdminPlugin(t *testing.T) {
|
||||
// 3. Settings
|
||||
schema, ok := ctx.Settings().Get("admin.system_cleanup_cron")
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "0 4 * * *", schema.Default)
|
||||
assert.Equal(t, "0 3 * * *", schema.Default)
|
||||
|
||||
provider, err := core.Inject[contracts.PublicConfigProvider](ctx)
|
||||
require.NoError(t, err)
|
||||
@@ -480,7 +485,7 @@ func TestAllDomainPluginsCombined(t *testing.T) {
|
||||
// Apply Domain plugins
|
||||
require.NoError(t, auth.New().Apply(ctx))
|
||||
require.NoError(t, user.New().Apply(ctx))
|
||||
require.NoError(t, message_gateway.New().Apply(ctx))
|
||||
require.NoError(t, msg_gateway.New().Apply(ctx))
|
||||
require.NoError(t, risk_control.New().Apply(ctx))
|
||||
require.NoError(t, admin.New().Apply(ctx))
|
||||
|
||||
@@ -532,3 +537,125 @@ func TestAllDomainPluginsCombined(t *testing.T) {
|
||||
// Clean shutdown
|
||||
require.NoError(t, ctx.Dispose())
|
||||
}
|
||||
|
||||
func TestSystemCleanupEventDrivenCoordination(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
ctx.Config().SetSource(core.NewMapSource(nil))
|
||||
require.NoError(t, ctx.Config().Resolve())
|
||||
testDB := setupTestDB(t)
|
||||
|
||||
// Apply Infra & Domain plugins
|
||||
require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx))
|
||||
require.NoError(t, cache.New().Apply(ctx))
|
||||
require.NoError(t, logger.New().Apply(ctx))
|
||||
require.NoError(t, storage.New().Apply(ctx))
|
||||
|
||||
require.NoError(t, user.New().Apply(ctx))
|
||||
require.NoError(t, msg_gateway.New().Apply(ctx))
|
||||
require.NoError(t, upload.New().Apply(ctx))
|
||||
require.NoError(t, admin.New().Apply(ctx))
|
||||
|
||||
// 1. Seed old and recent task executions (admin domain)
|
||||
now := time.Now()
|
||||
oldTime := now.Add(-10 * 24 * time.Hour)
|
||||
recentTime := now.Add(-1 * time.Hour)
|
||||
|
||||
oldExec := admin.TaskExecution{
|
||||
ID: 101,
|
||||
TaskID: "task-old-exec",
|
||||
TaskType: "test_task",
|
||||
Status: "success",
|
||||
CreatedAt: oldTime,
|
||||
}
|
||||
recentExec := admin.TaskExecution{
|
||||
ID: 102,
|
||||
TaskID: "task-recent-exec",
|
||||
TaskType: "test_task",
|
||||
Status: "success",
|
||||
CreatedAt: recentTime,
|
||||
}
|
||||
require.NoError(t, testDB.Create(&oldExec).Error)
|
||||
require.NoError(t, testDB.Model(&oldExec).UpdateColumn("created_at", oldTime).Error)
|
||||
require.NoError(t, testDB.Create(&recentExec).Error)
|
||||
|
||||
// 2. Seed old and recent push histories (msg_gateway domain, 30 days retention)
|
||||
oldHistoryTime := now.Add(-40 * 24 * time.Hour)
|
||||
oldHistory := msg_gateway.PushHistory{
|
||||
EventKey: "login",
|
||||
Channel: "telegram",
|
||||
Target: "123",
|
||||
Title: "Old login",
|
||||
Content: "Old content",
|
||||
Level: "info",
|
||||
Status: "success",
|
||||
CreatedAt: oldHistoryTime,
|
||||
}
|
||||
recentHistory := msg_gateway.PushHistory{
|
||||
EventKey: "login",
|
||||
Channel: "telegram",
|
||||
Target: "123",
|
||||
Title: "Recent login",
|
||||
Content: "Recent content",
|
||||
Level: "info",
|
||||
Status: "success",
|
||||
CreatedAt: recentTime,
|
||||
}
|
||||
require.NoError(t, testDB.Create(&oldHistory).Error)
|
||||
require.NoError(t, testDB.Model(&oldHistory).UpdateColumn("created_at", oldHistoryTime).Error)
|
||||
require.NoError(t, testDB.Create(&recentHistory).Error)
|
||||
|
||||
// 3. Seed old pending upload and recent pending upload (upload domain)
|
||||
oldUpload := upload.Upload{
|
||||
ID: 901,
|
||||
UserID: 1,
|
||||
FileName: "old.png",
|
||||
FilePath: "uploads/old.png",
|
||||
FileSize: 100,
|
||||
Status: upload.UploadStatusPending,
|
||||
CreatedAt: now.Add(-2 * time.Hour),
|
||||
}
|
||||
recentUpload := upload.Upload{
|
||||
ID: 902,
|
||||
UserID: 1,
|
||||
FileName: "recent.png",
|
||||
FilePath: "uploads/recent.png",
|
||||
FileSize: 100,
|
||||
Status: upload.UploadStatusPending,
|
||||
CreatedAt: now.Add(-10 * time.Minute),
|
||||
}
|
||||
require.NoError(t, testDB.Create(&oldUpload).Error)
|
||||
require.NoError(t, testDB.Model(&oldUpload).UpdateColumn("created_at", now.Add(-2*time.Hour)).Error)
|
||||
require.NoError(t, testDB.Create(&recentUpload).Error)
|
||||
|
||||
// Dispatch admin system cleanup task handler
|
||||
taskDef, ok := ctx.Tasks().Get("system:cleanup")
|
||||
require.True(t, ok, "system:cleanup task must be registered in admin")
|
||||
require.NotNil(t, taskDef.Handler)
|
||||
|
||||
type resultExecutor interface {
|
||||
Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error)
|
||||
}
|
||||
handler, ok := taskDef.Handler.(resultExecutor)
|
||||
require.True(t, ok)
|
||||
|
||||
res, err := handler.Execute(context.Background(), nil)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, res)
|
||||
|
||||
// Verify Admin cleanup: old task execution deleted, recent remains
|
||||
var execCount int64
|
||||
testDB.Model(&admin.TaskExecution{}).Count(&execCount)
|
||||
assert.Equal(t, int64(1), execCount)
|
||||
|
||||
// Verify MsgGateway cleanup: old push history deleted, recent remains
|
||||
var historyCount int64
|
||||
testDB.Model(&msg_gateway.PushHistory{}).Count(&historyCount)
|
||||
assert.Equal(t, int64(1), historyCount)
|
||||
|
||||
// Verify Upload cleanup: old pending upload deleted, recent remains
|
||||
var uploadCount int64
|
||||
testDB.Model(&upload.Upload{}).Count(&uploadCount)
|
||||
assert.Equal(t, int64(1), uploadCount)
|
||||
|
||||
require.NoError(t, ctx.Dispose())
|
||||
}
|
||||
|
||||
@@ -1,51 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package message_gateway_test
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/message_gateway"
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type mockDBService struct {
|
||||
db *gorm.DB
|
||||
}
|
||||
|
||||
func (m *mockDBService) GORM() *gorm.DB {
|
||||
return m.db
|
||||
}
|
||||
|
||||
func (m *mockDBService) DB(ctx context.Context) *gorm.DB {
|
||||
return m.db.WithContext(ctx)
|
||||
}
|
||||
|
||||
func (m *mockDBService) Named(_ string) *gorm.DB {
|
||||
return m.db
|
||||
}
|
||||
|
||||
func TestUpsertPairingCode_ReusesUnexpired(t *testing.T) {
|
||||
testDB, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
message_gateway.SetDBServiceForTest(&mockDBService{db: testDB})
|
||||
defer func() {
|
||||
message_gateway.SetDBServiceForTest(nil)
|
||||
cleanup()
|
||||
}()
|
||||
ctx := context.Background()
|
||||
first, err := message_gateway.UpsertPairingCode(ctx, 1, "tg-1", "ABCD1234", time.Now().Add(15*time.Minute))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
second, err := message_gateway.UpsertPairingCode(ctx, 1, "tg-1", "ZZZZ9999", time.Now().Add(15*time.Minute))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if first.Code != second.Code || first.Code != "ABCD1234" {
|
||||
t.Fatalf("reuse failed: %+v %+v", first, second)
|
||||
}
|
||||
}
|
||||
@@ -1,46 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package model
|
||||
|
||||
// Field is one admin form field.
|
||||
type Field struct {
|
||||
Key string `json:"key"`
|
||||
Type string `json:"type"`
|
||||
Required bool `json:"required"`
|
||||
}
|
||||
|
||||
// Definition describes a channel type form.
|
||||
type Definition struct {
|
||||
Type string `json:"type"`
|
||||
Fields []Field `json:"fields"`
|
||||
}
|
||||
|
||||
// ChannelDTO represents a channel for admin consumption.
|
||||
type ChannelDTO struct {
|
||||
ID uint64 `json:"id,string"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
OwnerScope string `json:"owner_scope"`
|
||||
OwnerID *uint64 `json:"owner_id,string,omitempty"`
|
||||
Enabled bool `json:"enabled"`
|
||||
Credentials map[string]string `json:"credentials"`
|
||||
Extra map[string]string `json:"extra"`
|
||||
}
|
||||
|
||||
// CreateChannelRequest is admin create payload.
|
||||
type CreateChannelRequest struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
Credentials map[string]string `json:"credentials"`
|
||||
Extra map[string]string `json:"extra"`
|
||||
}
|
||||
|
||||
// UpdateChannelRequest is admin update payload.
|
||||
type UpdateChannelRequest struct {
|
||||
Name string `json:"name"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
Credentials map[string]string `json:"credentials"`
|
||||
Extra map[string]string `json:"extra"`
|
||||
}
|
||||
@@ -1,242 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package model defines the domain entities, DTOs, and schemas for message_gateway.
|
||||
package model
|
||||
|
||||
import (
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Channel type and scope constants.
|
||||
const (
|
||||
ChannelTypeTelegram = "telegram"
|
||||
ChannelTypeQQ = "qq"
|
||||
MessageChannelTypeTelegram = "telegram"
|
||||
MessageChannelTypeQQ = "qq"
|
||||
MessageOwnerScopeSystem = "system"
|
||||
|
||||
TypeCustom = "custom"
|
||||
TypeEmail = "email"
|
||||
TypeTelegram = "telegram"
|
||||
)
|
||||
|
||||
// Capability describes what an adapter can send and receive.
|
||||
type Capability struct {
|
||||
Text bool
|
||||
Image bool
|
||||
File bool
|
||||
Reply bool
|
||||
Group bool
|
||||
}
|
||||
|
||||
// ChannelConfig is the decrypted runtime config passed to a factory.
|
||||
type ChannelConfig struct {
|
||||
ID uint64
|
||||
Type string
|
||||
Name string
|
||||
Credentials map[string]string
|
||||
Extra map[string]string
|
||||
}
|
||||
|
||||
// Recipient is the outbound destination on a platform.
|
||||
type Recipient struct {
|
||||
ChatID string
|
||||
PlatformUserID string
|
||||
}
|
||||
|
||||
// Attachment is a downloaded inbound file sitting on local disk.
|
||||
type Attachment struct {
|
||||
Path string
|
||||
FileName string
|
||||
MIME string
|
||||
Error string
|
||||
}
|
||||
|
||||
// InboundMessage is a normalized private-chat message.
|
||||
type InboundMessage struct {
|
||||
ChannelID uint64
|
||||
ChannelType string
|
||||
PlatformUserID string
|
||||
ChatID string
|
||||
MessageID string
|
||||
Text string
|
||||
Attachments []Attachment
|
||||
BindingUserID *uint64
|
||||
}
|
||||
|
||||
// OutboundMessage is a reply or probe send.
|
||||
type OutboundMessage struct {
|
||||
Text string
|
||||
ReplyToID string
|
||||
Attachments []Attachment
|
||||
}
|
||||
|
||||
// MessageChannel is an admin-configured messaging adapter.
|
||||
type MessageChannel struct {
|
||||
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
|
||||
Type string `json:"type" gorm:"size:32;not null"`
|
||||
Name string `json:"name" gorm:"size:64;not null"`
|
||||
OwnerScope string `json:"owner_scope" gorm:"size:32;not null;default:'system'"`
|
||||
OwnerID *uint64 `json:"owner_id,omitempty"`
|
||||
Credentials string `json:"credentials" gorm:"type:text;not null"`
|
||||
Extra string `json:"extra" gorm:"type:text"`
|
||||
Enabled bool `json:"enabled" gorm:"default:false;not null"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (MessageChannel) TableName() string {
|
||||
return "w_message_channels"
|
||||
}
|
||||
|
||||
// MessageBinding maps a platform user to a Wavelet user on one channel.
|
||||
type MessageBinding struct {
|
||||
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
|
||||
ChannelID uint64 `json:"channel_id" gorm:"not null;index"`
|
||||
PlatformUserID string `json:"platform_user_id" gorm:"size:128;not null;index"`
|
||||
UserID uint64 `json:"user_id" gorm:"not null;index"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (MessageBinding) TableName() string {
|
||||
return "w_message_bindings"
|
||||
}
|
||||
|
||||
// MessagePairingCode is a one-time bind code.
|
||||
type MessagePairingCode struct {
|
||||
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
|
||||
Code string `json:"code" gorm:"size:32;uniqueIndex;not null"`
|
||||
ChannelID uint64 `json:"channel_id" gorm:"not null;index"`
|
||||
PlatformUserID string `json:"platform_user_id" gorm:"size:128;not null;index"`
|
||||
UserID uint64 `json:"user_id" gorm:"not null;index"`
|
||||
ExpiresAt time.Time `json:"expires_at" gorm:"not null;index"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (MessagePairingCode) TableName() string {
|
||||
return "w_message_pairing_codes"
|
||||
}
|
||||
|
||||
// PushChannel 消息通道模型
|
||||
type PushChannel struct {
|
||||
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
|
||||
Name string `json:"name" gorm:"size:100;not null"`
|
||||
Description string `json:"description" gorm:"size:255"`
|
||||
Type string `json:"type" gorm:"size:50;not null;index"`
|
||||
URL string `json:"url" gorm:"type:text"`
|
||||
Token string `json:"token" gorm:"type:text"`
|
||||
Other string `json:"other" gorm:"type:text"`
|
||||
Enabled bool `json:"enabled" gorm:"index;not null;default:true"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
|
||||
}
|
||||
|
||||
// TableName 指定 GORM 表名
|
||||
func (PushChannel) TableName() string {
|
||||
return "w_push_channels"
|
||||
}
|
||||
|
||||
// Validate 验证与标准化字段
|
||||
func (c *PushChannel) Validate() error {
|
||||
c.Name = strings.TrimSpace(c.Name)
|
||||
if c.Name == "" {
|
||||
return errors.New(errs.ErrChannelNameRequired)
|
||||
}
|
||||
c.Type = strings.TrimSpace(c.Type)
|
||||
if c.Type == "" {
|
||||
return errors.New(errs.ErrChannelTypeRequired)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// PushEvent 系统通知事件模型
|
||||
type PushEvent struct {
|
||||
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
|
||||
EventKey string `json:"event_key" gorm:"uniqueIndex;size:80;not null"`
|
||||
Name string `json:"name" gorm:"size:100;not null"`
|
||||
TaskType string `json:"task_type" gorm:"size:100;index;not null;default:''"`
|
||||
Channels []string `json:"channels" gorm:"type:text;serializer:json;not null"`
|
||||
Targets []string `json:"targets" gorm:"type:text;serializer:json;not null"`
|
||||
Template string `json:"template" gorm:"type:text;not null"`
|
||||
Enabled bool `json:"enabled" gorm:"index;not null;default:false"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
|
||||
}
|
||||
|
||||
// TableName 指定 GORM 表名
|
||||
func (PushEvent) TableName() string {
|
||||
return "w_push_events"
|
||||
}
|
||||
|
||||
// Validate 验证 PushEvent 实体字段
|
||||
func (e *PushEvent) Validate() error {
|
||||
e.EventKey = strings.TrimSpace(e.EventKey)
|
||||
if e.EventKey == "" {
|
||||
return errors.New(errs.ErrEventKeyRequired)
|
||||
}
|
||||
e.Name = strings.TrimSpace(e.Name)
|
||||
if e.Name == "" {
|
||||
return errors.New(errs.ErrNameRequired)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// PushHistory 推送日志/历史实体
|
||||
type PushHistory struct {
|
||||
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
|
||||
EventKey string `json:"event_key" gorm:"size:80;not null;index"`
|
||||
Channel string `json:"channel" gorm:"size:50;not null;index"`
|
||||
Target string `json:"target" gorm:"size:255;not null"`
|
||||
Title string `json:"title" gorm:"size:255;not null"`
|
||||
Content string `json:"content" gorm:"type:text;not null"`
|
||||
Level string `json:"level" gorm:"size:20;not null;default:'INFO'"`
|
||||
Status string `json:"status" gorm:"size:20;not null;index"`
|
||||
ErrorMsg string `json:"error_msg" gorm:"type:text"`
|
||||
Payload string `json:"payload" gorm:"type:text"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
|
||||
}
|
||||
|
||||
// TableName 指定 GORM 表名
|
||||
func (PushHistory) TableName() string {
|
||||
return "w_push_histories"
|
||||
}
|
||||
|
||||
// BindRequest is the user bind body.
|
||||
type BindRequest struct {
|
||||
ChannelID string `json:"channel_id"`
|
||||
Code string `json:"code"`
|
||||
}
|
||||
|
||||
// BindingDTO is a user-facing binding row.
|
||||
type BindingDTO struct {
|
||||
ID uint64 `json:"id,string"`
|
||||
UserID uint64 `json:"user_id,string"`
|
||||
ChannelID uint64 `json:"channel_id,string"`
|
||||
ChannelName string `json:"channel_name"`
|
||||
ChannelType string `json:"channel_type"`
|
||||
PlatformUserID string `json:"platform_user_id"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
// PublicChannelDTO is an enabled channel a user can bind to.
|
||||
type PublicChannelDTO struct {
|
||||
ID uint64 `json:"id,string"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
}
|
||||
|
||||
// PushNotificationEvent defines the payload for eventbus notification trigger.
|
||||
type PushNotificationEvent struct {
|
||||
UserID uint64 `json:"user_id"`
|
||||
Channel string `json:"channel"`
|
||||
Title string `json:"title"`
|
||||
Content string `json:"content"`
|
||||
Metadata map[string]any `json:"metadata,omitempty"`
|
||||
}
|
||||
@@ -1,33 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package message_gateway
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestGenerateCode_AlphabetAndLength(t *testing.T) {
|
||||
code, err := GenerateCode()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(code) != 8 {
|
||||
t.Fatalf("len=%d", len(code))
|
||||
}
|
||||
for _, r := range code {
|
||||
if !strings.ContainsRune(CodeAlphabet, r) {
|
||||
t.Fatalf("bad rune %q", r)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeAndFormat(t *testing.T) {
|
||||
if got := NormalizeCode("ab-cd-ef-gh"); got != "ABCDEFGH" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
if got := FormatCode("ABCDEFGH"); got != "ABCD-EFGH" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
}
|
||||
@@ -1,359 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package message_gateway provides the Bot gateway, multi-channel notification dispatching, and asynchronous push worker domain plugin for Cordis.
|
||||
package message_gateway
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/core/extpoints"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/message_gateway/channels/qq"
|
||||
"Wavelet/plugins/domain/message_gateway/channels/telegram"
|
||||
"Wavelet/plugins/domain/message_gateway/handler"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/message_gateway/repository"
|
||||
"Wavelet/plugins/domain/message_gateway/service"
|
||||
"context"
|
||||
"embed"
|
||||
"reflect"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
//go:embed migrations/*/*.sql
|
||||
var mgMigrations embed.FS
|
||||
|
||||
// Option configures the message_gateway plugin.
|
||||
type Option func(*Plugin)
|
||||
|
||||
// WithAutoStartRunner enables automatic bot runner startup in the background.
|
||||
func WithAutoStartRunner(enable bool) Option {
|
||||
return func(p *Plugin) {
|
||||
p.autoStartRunner = enable
|
||||
}
|
||||
}
|
||||
|
||||
// Plugin implements core.Plugin to provide Bot gateway and notification dispatch domain services.
|
||||
type Plugin struct {
|
||||
autoStartRunner bool
|
||||
cancelRunner context.CancelFunc
|
||||
}
|
||||
|
||||
// New creates a new message_gateway domain plugin.
|
||||
func New(opts ...Option) *Plugin {
|
||||
p := &Plugin{}
|
||||
for _, opt := range opts {
|
||||
if opt != nil {
|
||||
opt(p)
|
||||
}
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
// Name returns the unique identifier for the message_gateway domain plugin.
|
||||
func (p *Plugin) Name() string {
|
||||
return "message_gateway"
|
||||
}
|
||||
|
||||
// Inject declares required dependencies for the message_gateway domain plugin.
|
||||
func (p *Plugin) Inject() []reflect.Type {
|
||||
return []reflect.Type{
|
||||
reflect.TypeFor[contracts.DBService](),
|
||||
// AuthService is captured as a middleware value in Apply, so it cannot
|
||||
// be late-bound with core.When like the other services below; the
|
||||
// kernel must mount auth first or the routes get a pass-through guard.
|
||||
reflect.TypeFor[contracts.AuthService](),
|
||||
}
|
||||
}
|
||||
|
||||
// Manifest returns the plugin metadata.
|
||||
func (p *Plugin) Manifest() core.Manifest {
|
||||
return core.Manifest{
|
||||
Name: "message_gateway",
|
||||
Version: "1.0.0",
|
||||
Description: "Bot gateway, multi-channel notification push, and async worker dispatch plugin",
|
||||
Author: "Wavelet Team",
|
||||
}
|
||||
}
|
||||
|
||||
type mgAppConfig struct {
|
||||
SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"`
|
||||
}
|
||||
|
||||
// DeclareConfig declares configuration bindings for the message_gateway plugin.
|
||||
func (p *Plugin) DeclareConfig() []core.ConfigBinding {
|
||||
return []core.ConfigBinding{
|
||||
{Prefix: "app", Target: &mgAppConfig{}},
|
||||
}
|
||||
}
|
||||
|
||||
// Apply registers message_gateway migrations, routes, tasks, schedules, events, and settings into the Context.
|
||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
var cfg mgAppConfig
|
||||
if err := ctx.Config().Bind("app", &cfg); err == nil && cfg.SessionSecret != "" {
|
||||
service.SetCredentialSecret(cfg.SessionSecret)
|
||||
}
|
||||
core.Bind[contracts.DBService](ctx, repository.SetDBService)
|
||||
core.Bind[contracts.CacheService](ctx, func(cache contracts.CacheService) {
|
||||
repository.SetCacheService(cache)
|
||||
service.SetCacheService(cache)
|
||||
})
|
||||
core.Bind[contracts.TaskService](ctx, service.SetTaskService)
|
||||
core.Bind[contracts.UserService](ctx, service.SetUserService)
|
||||
ctx.OnDispose(func() error {
|
||||
repository.SetDBService(nil)
|
||||
repository.SetCacheService(nil)
|
||||
service.SetCacheService(nil)
|
||||
service.SetTaskService(nil)
|
||||
service.SetUserService(nil)
|
||||
return nil
|
||||
})
|
||||
|
||||
// 0. Resolve auth service for middleware (via IoC, not direct import)
|
||||
denyAuth := ginutil.AuthUnavailable()
|
||||
loginMW := denyAuth
|
||||
adminMW := denyAuth
|
||||
if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil {
|
||||
if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok {
|
||||
loginMW = mw
|
||||
}
|
||||
if mw, ok := authSvc.RequireAdminMiddleware().(gin.HandlerFunc); ok {
|
||||
adminMW = mw
|
||||
}
|
||||
}
|
||||
|
||||
// 1. Register migrations
|
||||
ctx.Migrations().Register("message_gateway", mgMigrations)
|
||||
|
||||
// 2. Register User HTTP Routes
|
||||
handler.RegisterUserRoutes(ctx.Router().Group("/api/v1"), loginMW)
|
||||
|
||||
// 3. Register Admin Message Gateway HTTP Routes
|
||||
handler.RegisterAdminRoutes(ctx.Router().Group("/api/v1/admin"), loginMW, adminMW)
|
||||
|
||||
// 4. Register Admin Push HTTP Routes
|
||||
handler.RegisterAdminPushRoutes(ctx.Router().Group("/api/v1/admin"), loginMW, adminMW)
|
||||
|
||||
service.Register(model.MessageChannelTypeTelegram, telegram.New)
|
||||
service.Register(model.MessageChannelTypeQQ, qq.New)
|
||||
|
||||
const defaultTaskRetry = 3
|
||||
pushHandler := &service.PushHandler{}
|
||||
|
||||
// 5. Register background tasks
|
||||
ctx.Task().Register("message_gateway:push_notification", func(c context.Context, payload []byte) error {
|
||||
return pushHandler.Execute(c, payload)
|
||||
},
|
||||
extpoints.WithTaskType("push_notification"),
|
||||
extpoints.WithTaskName("消息网关推送通知"),
|
||||
extpoints.WithTaskDescription("异步执行系统通知的多渠道派发与推送"),
|
||||
extpoints.WithTaskCategory("push"),
|
||||
extpoints.WithTaskRetry(defaultTaskRetry),
|
||||
extpoints.WithTaskQueue("default"),
|
||||
extpoints.WithTaskRetryable(true),
|
||||
)
|
||||
|
||||
ctx.Task().Register(service.SendNotificationTask, func(c context.Context, payload []byte) error {
|
||||
return pushHandler.Execute(c, payload)
|
||||
}, extpoints.WithTaskMeta(service.SendNotificationMeta), extpoints.WithTaskRetry(defaultTaskRetry))
|
||||
|
||||
ctx.Task().Register(service.TaskDispatchBotMsg, &service.BotDispatchHandler{},
|
||||
extpoints.WithTaskMeta(service.BotDispatchMeta))
|
||||
|
||||
ctx.Task().Register("message_gateway:cleanup_pairing_codes", func(c context.Context, _ []byte) error {
|
||||
return repository.DeleteExpiredPairingCodes(c)
|
||||
},
|
||||
extpoints.WithTaskType("cleanup_pairing_codes"),
|
||||
extpoints.WithTaskName("清理过期配对码"),
|
||||
extpoints.WithTaskDescription("定时清理已过期的平台 Bot 配对码"),
|
||||
extpoints.WithTaskCategory("messaging"),
|
||||
extpoints.WithTaskRetry(defaultTaskRetry),
|
||||
extpoints.WithTaskQueue("default"),
|
||||
extpoints.WithTaskRetryable(true),
|
||||
)
|
||||
|
||||
// 6. Register Cron Schedules
|
||||
ctx.Schedule().RegisterCron("*/10 * * * *", "message_gateway:cleanup_pairing_codes", map[string]any{"action": "cleanup"})
|
||||
|
||||
// 7. Register EventBus listeners for decoupled push triggers
|
||||
ctx.Events().On("notification:push", func(c context.Context, e model.PushNotificationEvent) error {
|
||||
meta := model.EventMetadata{
|
||||
Key: "eventbus:" + e.Channel,
|
||||
Name: e.Title,
|
||||
DefaultTemplate: model.NotificationMessage{
|
||||
Title: e.Title,
|
||||
Content: e.Content,
|
||||
Level: model.DefaultLevelInfo,
|
||||
Ext: e.Metadata,
|
||||
},
|
||||
Description: "EventBus triggered notification",
|
||||
}
|
||||
service.DefaultTrigger.Trigger(c, meta, map[string]any{
|
||||
"user.id": e.UserID,
|
||||
"title": e.Title,
|
||||
"content": e.Content,
|
||||
})
|
||||
return nil
|
||||
})
|
||||
|
||||
// 8. Register task completed event listener
|
||||
ctx.Events().On(contracts.EventTopicTaskCompleted, func(c context.Context, e contracts.TaskCompletedEvent) error {
|
||||
service.HandleTaskCompleted(c, e)
|
||||
return nil
|
||||
})
|
||||
|
||||
// 9. Register built-in domain events and provide PushRegistry
|
||||
service.RegisterCustomEvents()
|
||||
core.Provide[contracts.PushRegistry](ctx, service.PushRegistryAdapter{})
|
||||
|
||||
// 10. Register Settings Schemas
|
||||
ctx.Settings().Register(extpoints.SettingSchema{
|
||||
Key: "message_gateway.pairing_code_expiry_minutes",
|
||||
Default: 15,
|
||||
Description: "Expiry duration for bot pairing codes in minutes",
|
||||
Type: "integer",
|
||||
Category: "messaging",
|
||||
})
|
||||
ctx.Settings().Register(extpoints.SettingSchema{
|
||||
Key: "message_gateway.max_bindings_per_user",
|
||||
Default: 5,
|
||||
Description: "Maximum platform bot bindings per user",
|
||||
Type: "integer",
|
||||
Category: "messaging",
|
||||
})
|
||||
|
||||
// 11. Optional runner start & lifecycle
|
||||
if p.autoStartRunner {
|
||||
runnerCtx, cancel := context.WithCancel(ctx.GoContext())
|
||||
p.cancelRunner = cancel
|
||||
util.Go(func() {
|
||||
_ = service.Start(runnerCtx)
|
||||
})
|
||||
}
|
||||
|
||||
ctx.OnDispose(func() error {
|
||||
if p.cancelRunner != nil {
|
||||
p.cancelRunner()
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Re-exported constants.
|
||||
const (
|
||||
CodeAlphabet = service.CodeAlphabet
|
||||
CodeLength = service.CodeLength
|
||||
)
|
||||
|
||||
// MessageChannel is an alias for model.MessageChannel.
|
||||
type MessageChannel = model.MessageChannel
|
||||
|
||||
// MessageBinding is an alias for model.MessageBinding.
|
||||
type MessageBinding = model.MessageBinding
|
||||
|
||||
// MessagePairingCode is an alias for model.MessagePairingCode.
|
||||
type MessagePairingCode = model.MessagePairingCode
|
||||
|
||||
// PushChannel is an alias for model.PushChannel.
|
||||
type PushChannel = model.PushChannel
|
||||
|
||||
// PushEvent is an alias for model.PushEvent.
|
||||
type PushEvent = model.PushEvent
|
||||
|
||||
// PushHistory is an alias for model.PushHistory.
|
||||
type PushHistory = model.PushHistory
|
||||
|
||||
// PushNotificationEvent is an alias for model.PushNotificationEvent.
|
||||
type PushNotificationEvent = model.PushNotificationEvent
|
||||
|
||||
// ChannelConfig is an alias for model.ChannelConfig.
|
||||
type ChannelConfig = model.ChannelConfig
|
||||
|
||||
// Capability is an alias for model.Capability.
|
||||
type Capability = model.Capability
|
||||
|
||||
// Recipient is an alias for model.Recipient.
|
||||
type Recipient = model.Recipient
|
||||
|
||||
// Attachment is an alias for model.Attachment.
|
||||
type Attachment = model.Attachment
|
||||
|
||||
// InboundMessage is an alias for model.InboundMessage.
|
||||
type InboundMessage = model.InboundMessage
|
||||
|
||||
// OutboundMessage is an alias for model.OutboundMessage.
|
||||
type OutboundMessage = model.OutboundMessage
|
||||
|
||||
// BindingDTO is an alias for model.BindingDTO.
|
||||
type BindingDTO = model.BindingDTO
|
||||
|
||||
// PublicChannelDTO is an alias for model.PublicChannelDTO.
|
||||
type PublicChannelDTO = model.PublicChannelDTO
|
||||
|
||||
// Definition is an alias for model.Definition.
|
||||
type Definition = model.Definition
|
||||
|
||||
// ChannelDTO is an alias for model.ChannelDTO.
|
||||
type ChannelDTO = model.ChannelDTO
|
||||
|
||||
// CreateChannelRequest is an alias for model.CreateChannelRequest.
|
||||
type CreateChannelRequest = model.CreateChannelRequest
|
||||
|
||||
// UpdateChannelRequest is an alias for model.UpdateChannelRequest.
|
||||
type UpdateChannelRequest = model.UpdateChannelRequest
|
||||
|
||||
// PushDefinition is an alias for model.PushDefinition.
|
||||
type PushDefinition = model.PushDefinition
|
||||
|
||||
// PushField is an alias for model.PushField.
|
||||
type PushField = model.PushField
|
||||
|
||||
// NotificationMessage is an alias for model.NotificationMessage.
|
||||
type NotificationMessage = model.NotificationMessage
|
||||
|
||||
// EventMetadata is an alias for model.EventMetadata.
|
||||
type EventMetadata = model.EventMetadata
|
||||
|
||||
// SendPayload is an alias for model.SendPayload.
|
||||
type SendPayload = model.SendPayload
|
||||
|
||||
// Handler is an alias for service.Handler.
|
||||
type Handler = service.Handler
|
||||
|
||||
// Factory is an alias for service.Factory.
|
||||
type Factory = service.Factory
|
||||
|
||||
// Channel is an alias for service.Channel.
|
||||
type Channel = service.Channel
|
||||
|
||||
// Runner is an alias for service.Runner.
|
||||
type Runner = service.Runner
|
||||
|
||||
// EventTrigger is an alias for service.EventTrigger.
|
||||
type EventTrigger = service.EventTrigger
|
||||
|
||||
// PushHandler is an alias for service.PushHandler.
|
||||
type PushHandler = service.PushHandler
|
||||
|
||||
// Re-exported variables and functions.
|
||||
var (
|
||||
SetDBServiceForTest = repository.SetDBServiceForTest
|
||||
UpsertPairingCode = repository.UpsertPairingCode
|
||||
Register = service.Register
|
||||
Lookup = service.Lookup
|
||||
GenerateCode = service.GenerateCode
|
||||
NormalizeCode = service.NormalizeCode
|
||||
FormatCode = service.FormatCode
|
||||
Start = service.Start
|
||||
Stop = service.Stop
|
||||
GlobalRunner = service.GlobalRunner
|
||||
DefaultTrigger = service.DefaultTrigger
|
||||
SyncEvents = service.SyncEvents
|
||||
AdminLogin = service.AdminLogin
|
||||
HandleAdminLoggedIn = service.HandleAdminLoggedIn
|
||||
)
|
||||
@@ -1,104 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/util"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/smtp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func init() {
|
||||
Register("email", &EmailPusher{})
|
||||
}
|
||||
|
||||
// EmailPusher 极简 SMTP 邮件推送实现 (静态、解耦)
|
||||
type EmailPusher struct{}
|
||||
|
||||
// sanitizeEmailHeader removes CR/LF bytes so untrusted values cannot inject
|
||||
// additional email headers (email header injection).
|
||||
func sanitizeEmailHeader(v string) string {
|
||||
v = strings.ReplaceAll(v, "\r", "")
|
||||
v = strings.ReplaceAll(v, "\n", "")
|
||||
return v
|
||||
}
|
||||
|
||||
// Send 发送邮件
|
||||
func (p *EmailPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, ext map[string]any) (string, error) {
|
||||
if cfg.URL == "" || cfg.Key == "" || cfg.Secret == "" {
|
||||
return "", errors.New("email: SMTP configuration (url, key, secret) is incomplete")
|
||||
}
|
||||
if target == "" {
|
||||
return "", errors.New("email: target email address is required")
|
||||
}
|
||||
|
||||
title := bodyTitle(body)
|
||||
content := bodyContent(body, "<p><b>%s</b>: %v</p>", "")
|
||||
|
||||
// 邮件头和体
|
||||
from := cfg.Key
|
||||
to := target
|
||||
|
||||
// 如果 ext 中指定了 from_name,我们在 From 头部包含它
|
||||
fromName := "System Notification"
|
||||
if ext != nil {
|
||||
if fn, ok := ext["from_name"].(string); ok && fn != "" {
|
||||
fromName = fn
|
||||
}
|
||||
}
|
||||
|
||||
subjectHeader := fmt.Sprintf("Subject: %s\r\n", sanitizeEmailHeader(title))
|
||||
fromHeader := fmt.Sprintf("From: %s <%s>\r\n", sanitizeEmailHeader(fromName), sanitizeEmailHeader(from))
|
||||
toHeader := fmt.Sprintf("To: %s\r\n", sanitizeEmailHeader(to))
|
||||
mimeHeader := "MIME-version: 1.0;\r\nContent-Type: text/html; charset=\"UTF-8\";\r\n\r\n"
|
||||
|
||||
// 拼装完整的邮件报文
|
||||
// 简单的 HTML 正文渲染
|
||||
htmlBody := fmt.Sprintf(`<html><body><h2>%s</h2><div>%s</div></body></html>`, title, content)
|
||||
msg := []byte(fromHeader + toHeader + subjectHeader + mimeHeader + htmlBody + "\r\n")
|
||||
|
||||
// 解析 Host 和 Port
|
||||
host, port, err := net.SplitHostPort(cfg.URL)
|
||||
if err != nil {
|
||||
host = cfg.URL
|
||||
port = "25" // 默认 SMTP 端口
|
||||
}
|
||||
|
||||
auth := smtp.PlainAuth("", cfg.Key, cfg.Secret, host)
|
||||
|
||||
// 异步超时处理
|
||||
errChan := make(chan error, 1)
|
||||
util.Go(func() {
|
||||
errChan <- smtp.SendMail(host+":"+port, auth, from, []string{to}, msg)
|
||||
})
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return "", ctx.Err()
|
||||
case err := <-errChan:
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("email: send smtp mail failed: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// ValidateConfig 校验邮件 SMTP 配置
|
||||
func (p *EmailPusher) ValidateConfig(cfg Config) error {
|
||||
if cfg.URL == "" {
|
||||
return errors.New("SMTP host:port is required")
|
||||
}
|
||||
if cfg.Key == "" {
|
||||
return errors.New("SMTP username is required")
|
||||
}
|
||||
if cfg.Secret == "" {
|
||||
return errors.New("SMTP password is required")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -1,26 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestSanitizeEmailHeader(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
want string
|
||||
}{
|
||||
{"plain", "System Notification", "System Notification"},
|
||||
{"crlf stripped", "alert\r\nBcc: attacker@example.com", "alertBcc: attacker@example.com"},
|
||||
{"cr stripped", "a\rb", "ab"},
|
||||
{"lf stripped", "a\nb", "ab"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := sanitizeEmailHeader(tt.input); got != tt.want {
|
||||
t.Errorf("sanitizeEmailHeader(%q) = %q, want %q", tt.input, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,256 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/httppool"
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
func init() {
|
||||
Register("lark", &LarkPusher{})
|
||||
}
|
||||
|
||||
const (
|
||||
msgTypeInteractive = "interactive"
|
||||
)
|
||||
|
||||
// LarkPusher 飞书 Webhook 机器人推送实现
|
||||
type LarkPusher struct{}
|
||||
|
||||
type larkTextContent struct {
|
||||
Text string `json:"text"`
|
||||
}
|
||||
|
||||
type larkCardHeaderTitle struct {
|
||||
Content string `json:"content"`
|
||||
Tag string `json:"tag"`
|
||||
}
|
||||
|
||||
type larkCardHeader struct {
|
||||
Template string `json:"template"` // "blue", "orange", "red" etc.
|
||||
Title larkCardHeaderTitle `json:"title"`
|
||||
}
|
||||
|
||||
type larkCardElementText struct {
|
||||
Content string `json:"content"`
|
||||
Tag string `json:"tag"` // "lark_md"
|
||||
}
|
||||
|
||||
type larkCardElement struct {
|
||||
Tag string `json:"tag"` // "div"
|
||||
Text larkCardElementText `json:"text"`
|
||||
}
|
||||
|
||||
type larkCardContent struct {
|
||||
Header larkCardHeader `json:"header"`
|
||||
Elements []larkCardElement `json:"elements"`
|
||||
}
|
||||
|
||||
type larkMessageRequest struct {
|
||||
MessageType string `json:"msg_type"`
|
||||
Timestamp string `json:"timestamp,omitempty"`
|
||||
Sign string `json:"sign,omitempty"`
|
||||
Content larkTextContent `json:"content,omitempty"`
|
||||
Card *larkCardContent `json:"card,omitempty"`
|
||||
}
|
||||
|
||||
type larkMessageResponse struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
|
||||
// Send 执行飞书消息发送
|
||||
//
|
||||
//nolint:nestif,cyclop
|
||||
func (p *LarkPusher) Send(ctx context.Context, cfg Config, _ string, body map[string]any, template string, _ map[string]any) (string, error) {
|
||||
if cfg.URL == "" {
|
||||
return "", errors.New("lark: URL is required")
|
||||
}
|
||||
|
||||
var req larkMessageRequest
|
||||
|
||||
// 1. 如果有自定义模板,我们尝试进行解析
|
||||
if template != "" {
|
||||
rendered := ParseTemplate(template, body)
|
||||
|
||||
// 尝试解析原生的 Lark Card
|
||||
var customCard larkCardContent
|
||||
var rawMap map[string]any
|
||||
_ = json.Unmarshal([]byte(rendered), &rawMap)
|
||||
|
||||
if rawMap != nil && rawMap["elements"] != nil {
|
||||
// 如果包含 elements 字段,说明是用户定制的原生飞书卡片 JSON
|
||||
if err := json.Unmarshal([]byte(rendered), &customCard); err == nil {
|
||||
req.MessageType = msgTypeInteractive
|
||||
req.Card = &customCard
|
||||
} else {
|
||||
req.MessageType = "text"
|
||||
req.Content.Text = rendered
|
||||
}
|
||||
} else {
|
||||
// 说明配置的是系统统一通知消息 of JSON 模板:{"title": "...", "content": "...", "level": "..."}
|
||||
type larkNotificationMessage struct {
|
||||
Title string `json:"title"`
|
||||
Content string `json:"content"`
|
||||
Level string `json:"level"`
|
||||
}
|
||||
var msg larkNotificationMessage
|
||||
if err := json.Unmarshal([]byte(rendered), &msg); err == nil && (msg.Title != "" || msg.Content != "") {
|
||||
title := msg.Title
|
||||
if title == "" {
|
||||
title = defaultTitle
|
||||
}
|
||||
content := msg.Content
|
||||
level := strings.ToUpper(msg.Level)
|
||||
if level == "" {
|
||||
level = levelInfo
|
||||
}
|
||||
|
||||
headerColor := "blue"
|
||||
switch level {
|
||||
case "IMPORTANT":
|
||||
headerColor = "orange"
|
||||
case "CRITICAL":
|
||||
headerColor = "red"
|
||||
}
|
||||
|
||||
req.MessageType = msgTypeInteractive
|
||||
req.Card = &larkCardContent{
|
||||
Header: larkCardHeader{
|
||||
Template: headerColor,
|
||||
Title: larkCardHeaderTitle{
|
||||
Content: title,
|
||||
Tag: "plain_text",
|
||||
},
|
||||
},
|
||||
Elements: []larkCardElement{
|
||||
{
|
||||
Tag: "div",
|
||||
Text: larkCardElementText{
|
||||
Content: content,
|
||||
Tag: "lark_md",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
} else {
|
||||
// 兜底:如果无法按 JSON 解析出结构化字段,当做普通文本发送
|
||||
req.MessageType = "text"
|
||||
req.Content.Text = rendered
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// 2. 如果无模板,默认生成一个精美的飞书互动卡片
|
||||
title := bodyTitle(body)
|
||||
content := bodyContent(body, "**%s**: %v", "\n")
|
||||
level := bodyLevel(body)
|
||||
|
||||
// 根据级别确定飞书卡片头部的背景色模板
|
||||
headerColor := "blue"
|
||||
switch level {
|
||||
case "IMPORTANT":
|
||||
headerColor = "orange"
|
||||
case "CRITICAL":
|
||||
headerColor = "red"
|
||||
}
|
||||
|
||||
req.MessageType = msgTypeInteractive
|
||||
req.Card = &larkCardContent{
|
||||
Header: larkCardHeader{
|
||||
Template: headerColor,
|
||||
Title: larkCardHeaderTitle{
|
||||
Content: title,
|
||||
Tag: "plain_text",
|
||||
},
|
||||
},
|
||||
Elements: []larkCardElement{
|
||||
{
|
||||
Tag: "div",
|
||||
Text: larkCardElementText{
|
||||
Content: content,
|
||||
Tag: "lark_md",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// 3. 计算签名 (如果配置了 secret)
|
||||
if cfg.Secret != "" {
|
||||
timestamp := time.Now().Unix()
|
||||
sign, err := larkSign(cfg.Secret, timestamp)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("lark: sign failed: %w", err)
|
||||
}
|
||||
req.Timestamp = strconv.FormatInt(timestamp, 10)
|
||||
req.Sign = sign
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("lark: marshal request failed: %w", err)
|
||||
}
|
||||
|
||||
// 4. 发送 POST 请求
|
||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, cfg.URL, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("lark: create http request failed: %w", err)
|
||||
}
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
client := httppool.NewClient(defaultHTTPClientTimeout)
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("lark: http request failed: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("lark: http status %s", resp.Status)
|
||||
}
|
||||
|
||||
var res larkMessageResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&res); err != nil {
|
||||
return "", fmt.Errorf("lark: decode response failed: %w", err)
|
||||
}
|
||||
|
||||
if res.Code != 0 {
|
||||
return "", fmt.Errorf("lark: send message failed, code %d: %s", res.Code, res.Msg)
|
||||
}
|
||||
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// ValidateConfig 校验飞书配置
|
||||
func (p *LarkPusher) ValidateConfig(cfg Config) error {
|
||||
if cfg.URL == "" {
|
||||
return errors.New("webhook URL is required")
|
||||
}
|
||||
if !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") {
|
||||
return errors.New("webhook URL must start with http:// or https://")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func larkSign(secret string, timestamp int64) (string, error) {
|
||||
stringToSign := fmt.Sprintf("%v", timestamp) + "\n" + secret
|
||||
h := hmac.New(sha256.New, []byte(stringToSign))
|
||||
_, err := h.Write(nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64.StdEncoding.EncodeToString(h.Sum(nil)), nil
|
||||
}
|
||||
@@ -1,140 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/httppool"
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func init() {
|
||||
Register("telegram", &TelegramPusher{})
|
||||
}
|
||||
|
||||
// TelegramPusher Telegram 机器人推送实现
|
||||
type TelegramPusher struct{}
|
||||
|
||||
type telegramMessageRequest struct {
|
||||
ChatID string `json:"chat_id"`
|
||||
Text string `json:"text"`
|
||||
ParseMode string `json:"parse_mode,omitempty"`
|
||||
}
|
||||
|
||||
type telegramErrorResponse struct {
|
||||
Ok bool `json:"ok"`
|
||||
ErrorCode int `json:"error_code"`
|
||||
Description string `json:"description"`
|
||||
}
|
||||
|
||||
// Send 执行 Telegram 消息发送
|
||||
func (p *TelegramPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, template string, _ map[string]any) (string, error) {
|
||||
if cfg.Secret == "" {
|
||||
return "", errors.New("telegram: Bot Token (Secret) is required")
|
||||
}
|
||||
|
||||
chatID := target
|
||||
if chatID == "" {
|
||||
chatID = cfg.Key // Use default chat ID (Key) if target is blank
|
||||
}
|
||||
if chatID == "" {
|
||||
return "", errors.New("telegram: chat_id (target or default Key) is required")
|
||||
}
|
||||
|
||||
baseURL := cfg.URL
|
||||
if baseURL == "" {
|
||||
baseURL = "https://api.telegram.org"
|
||||
}
|
||||
baseURL = strings.TrimSuffix(baseURL, "/")
|
||||
|
||||
title := bodyTitle(body)
|
||||
content := bodyContent(body, "<b>%s</b>: %v", "\n")
|
||||
level := bodyLevel(body)
|
||||
|
||||
var text string
|
||||
if template != "" {
|
||||
text = ParseTemplate(template, body)
|
||||
} else {
|
||||
text = fmt.Sprintf("<b>[%s] %s</b>\n\n%s", escapeHTML(level), escapeHTML(title), escapeHTML(content))
|
||||
}
|
||||
|
||||
// Try sending with HTML parse mode
|
||||
err := p.sendMessage(ctx, baseURL, cfg.Secret, chatID, text, "HTML")
|
||||
if err != nil {
|
||||
// Fallback: send as plain text without parse mode
|
||||
plainText := text
|
||||
if template == "" {
|
||||
plainText = fmt.Sprintf("[%s] %s\n\n%s", level, title, content)
|
||||
}
|
||||
fallbackErr := p.sendMessage(ctx, baseURL, cfg.Secret, chatID, plainText, "")
|
||||
if fallbackErr != nil {
|
||||
return "", fmt.Errorf("telegram: send message failed (fallback also failed): %w (original HTML error: %w)", fallbackErr, err)
|
||||
}
|
||||
}
|
||||
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// ValidateConfig 校验 Telegram 配置
|
||||
func (p *TelegramPusher) ValidateConfig(cfg Config) error {
|
||||
if cfg.Secret == "" {
|
||||
return errors.New("bot Token (Secret) is required")
|
||||
}
|
||||
if cfg.URL != "" {
|
||||
if !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") {
|
||||
return errors.New("API base URL must start with http:// or https://")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *TelegramPusher) sendMessage(ctx context.Context, baseURL, token, chatID, text, parseMode string) error {
|
||||
apiURL := fmt.Sprintf("%s/bot%s/sendMessage", baseURL, token)
|
||||
|
||||
reqPayload := telegramMessageRequest{
|
||||
ChatID: chatID,
|
||||
Text: text,
|
||||
ParseMode: parseMode,
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(reqPayload)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal request failed: %w", err)
|
||||
}
|
||||
|
||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return fmt.Errorf("create http request failed: %w", err)
|
||||
}
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
client := httppool.NewClient(defaultHTTPClientTimeout)
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil {
|
||||
return fmt.Errorf("http request failed: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
var errRes telegramErrorResponse
|
||||
if decodeErr := json.NewDecoder(resp.Body).Decode(&errRes); decodeErr == nil {
|
||||
return fmt.Errorf("http status %d: %s", resp.StatusCode, errRes.Description)
|
||||
}
|
||||
return fmt.Errorf("http status %s", resp.Status)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func escapeHTML(s string) string {
|
||||
s = strings.ReplaceAll(s, "&", "&")
|
||||
s = strings.ReplaceAll(s, "<", "<")
|
||||
s = strings.ReplaceAll(s, ">", ">")
|
||||
return s
|
||||
}
|
||||
@@ -1,116 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestTelegramPusher_Send(t *testing.T) {
|
||||
t.Run("successful send with HTML parse mode", func(t *testing.T) {
|
||||
var receivedReq telegramMessageRequest
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
assert.Equal(t, "/botmy-token/sendMessage", r.URL.Path)
|
||||
assert.Equal(t, http.MethodPost, r.Method)
|
||||
assert.Equal(t, "application/json", r.Header.Get("Content-Type"))
|
||||
|
||||
err := json.NewDecoder(r.Body).Decode(&receivedReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(`{"ok": true}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
pusher := &TelegramPusher{}
|
||||
cfg := Config{
|
||||
Channel: "telegram",
|
||||
URL: server.URL,
|
||||
Secret: "my-token",
|
||||
}
|
||||
body := map[string]any{
|
||||
"title": "Alert",
|
||||
"content": "Host down",
|
||||
"level": "CRITICAL",
|
||||
}
|
||||
_, err := pusher.Send(context.Background(), cfg, "123456", body, "", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, "123456", receivedReq.ChatID)
|
||||
assert.Contains(t, receivedReq.Text, "[CRITICAL] Alert")
|
||||
assert.Contains(t, receivedReq.Text, "Host down")
|
||||
assert.Equal(t, "HTML", receivedReq.ParseMode)
|
||||
})
|
||||
|
||||
t.Run("fallback to plain text on HTML error", func(t *testing.T) {
|
||||
var requests []*telegramMessageRequest
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var req telegramMessageRequest
|
||||
err := json.NewDecoder(r.Body).Decode(&req)
|
||||
require.NoError(t, err)
|
||||
requests = append(requests, &req)
|
||||
|
||||
if len(requests) == 1 {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
_, _ = w.Write([]byte(`{"ok": false, "error_code": 400, "description": "Bad Request: can't parse entities"}`))
|
||||
} else {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(`{"ok": true}`))
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
pusher := &TelegramPusher{}
|
||||
cfg := Config{
|
||||
Channel: "telegram",
|
||||
URL: server.URL,
|
||||
Secret: "my-token",
|
||||
}
|
||||
body := map[string]any{
|
||||
"title": "Alert & Info",
|
||||
"content": "A < B comparison",
|
||||
"level": "INFO",
|
||||
}
|
||||
_, err := pusher.Send(context.Background(), cfg, "123456", body, "", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Len(t, requests, 2)
|
||||
assert.Equal(t, "HTML", requests[0].ParseMode)
|
||||
assert.Equal(t, "", requests[1].ParseMode)
|
||||
assert.Contains(t, requests[1].Text, "[INFO] Alert & Info")
|
||||
assert.Contains(t, requests[1].Text, "A < B comparison")
|
||||
})
|
||||
|
||||
t.Run("validation error", func(t *testing.T) {
|
||||
pusher := &TelegramPusher{}
|
||||
cfg := Config{
|
||||
Channel: "telegram",
|
||||
URL: "https://api.telegram.org",
|
||||
}
|
||||
err := pusher.ValidateConfig(cfg)
|
||||
assert.Error(t, err)
|
||||
|
||||
cfg = Config{
|
||||
Channel: "telegram",
|
||||
URL: "ftp://api.telegram.org",
|
||||
Secret: "token",
|
||||
}
|
||||
err = pusher.ValidateConfig(cfg)
|
||||
assert.Error(t, err)
|
||||
|
||||
cfg = Config{
|
||||
Channel: "telegram",
|
||||
Secret: "token",
|
||||
}
|
||||
err = pusher.ValidateConfig(cfg)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
}
|
||||
@@ -1,111 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"maps"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ParseTemplate parses template strings by replacing {{placeholder}} structures with values from body.
|
||||
// It is a single-pass parser designed for high performance and low allocations.
|
||||
func ParseTemplate(template string, body map[string]any) string {
|
||||
var buf strings.Builder
|
||||
buf.Grow(len(template))
|
||||
|
||||
i := 0
|
||||
for {
|
||||
pos := strings.Index(template[i:], "{{")
|
||||
if pos == -1 {
|
||||
buf.WriteString(template[i:])
|
||||
break
|
||||
}
|
||||
// Write prefix
|
||||
buf.WriteString(template[i : i+pos])
|
||||
i += pos + 2 // skip "{{"
|
||||
|
||||
endPos := strings.Index(template[i:], "}}")
|
||||
if endPos == -1 {
|
||||
// Unbalanced "{{"
|
||||
buf.WriteString("{{")
|
||||
buf.WriteString(template[i:])
|
||||
break
|
||||
}
|
||||
key := template[i : i+endPos]
|
||||
if val, ok := body[key]; ok {
|
||||
buf.WriteString(formatValue(val))
|
||||
} else {
|
||||
// Keep the placeholder if key not found
|
||||
buf.WriteString("{{")
|
||||
buf.WriteString(key)
|
||||
buf.WriteString("}}")
|
||||
}
|
||||
i += endPos + 2 // skip "}}"
|
||||
}
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
func formatValue(v any) string {
|
||||
if v == nil {
|
||||
return ""
|
||||
}
|
||||
switch val := v.(type) {
|
||||
case string:
|
||||
return val
|
||||
case []byte:
|
||||
return string(val)
|
||||
case int:
|
||||
return strconv.Itoa(val)
|
||||
case int32:
|
||||
return strconv.FormatInt(int64(val), 10)
|
||||
case int64:
|
||||
return strconv.FormatInt(val, 10)
|
||||
case float64:
|
||||
return strconv.FormatFloat(val, 'f', -1, 64)
|
||||
case bool:
|
||||
return strconv.FormatBool(val)
|
||||
default:
|
||||
// If it's a map, slice, or struct, marshal it to JSON.
|
||||
b, err := json.Marshal(v)
|
||||
if err == nil {
|
||||
return string(b)
|
||||
}
|
||||
return fmt.Sprintf("%v", v)
|
||||
}
|
||||
}
|
||||
|
||||
// bodyTitle returns the notification title, falling back to the default.
|
||||
func bodyTitle(body map[string]any) string {
|
||||
if t, ok := body["title"].(string); ok && t != "" {
|
||||
return t
|
||||
}
|
||||
return defaultTitle
|
||||
}
|
||||
|
||||
// bodyContent returns the notification body, rendering every entry with format
|
||||
// (a "%s … %v" pair) and joining them with sep when no content field is given.
|
||||
// Entries render in sorted key order so identical bodies always produce
|
||||
// identical text.
|
||||
func bodyContent(body map[string]any, format, sep string) string {
|
||||
if c, ok := body["content"].(string); ok && c != "" {
|
||||
return c
|
||||
}
|
||||
parts := make([]string, 0, len(body))
|
||||
for _, k := range slices.Sorted(maps.Keys(body)) {
|
||||
parts = append(parts, fmt.Sprintf(format, k, body[k]))
|
||||
}
|
||||
return strings.Join(parts, sep)
|
||||
}
|
||||
|
||||
// bodyLevel returns the upper-cased notification level, falling back to INFO.
|
||||
func bodyLevel(body map[string]any) string {
|
||||
if l, ok := body["level"].(string); ok && l != "" {
|
||||
return strings.ToUpper(l)
|
||||
}
|
||||
return levelInfo
|
||||
}
|
||||
@@ -1,38 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package message_gateway
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type stubChannel struct{}
|
||||
|
||||
func (stubChannel) Type() string { return "stub" }
|
||||
func (stubChannel) Connect(context.Context) error {
|
||||
return nil
|
||||
}
|
||||
func (stubChannel) Disconnect(context.Context) error { return nil }
|
||||
func (stubChannel) Send(context.Context, Recipient, OutboundMessage) error {
|
||||
return nil
|
||||
}
|
||||
func (stubChannel) Capabilities() Capability { return Capability{Text: true} }
|
||||
|
||||
func TestRegisterLookup(t *testing.T) {
|
||||
Register("stub", func(ChannelConfig, Handler) (Channel, error) {
|
||||
return stubChannel{}, nil
|
||||
})
|
||||
fn, ok := Lookup("stub")
|
||||
if !ok {
|
||||
t.Fatal("expected factory")
|
||||
}
|
||||
ch, err := fn(ChannelConfig{}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if ch.Type() != "stub" {
|
||||
t.Fatalf("type=%s", ch.Type())
|
||||
}
|
||||
}
|
||||
@@ -1,131 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository_test
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_gateway/repository"
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// stubDBService satisfies contracts.DBService over a test database handle.
|
||||
type stubDBService struct{ db *gorm.DB }
|
||||
|
||||
func (s stubDBService) GORM() *gorm.DB { return s.db }
|
||||
|
||||
func (s stubDBService) DB(_ context.Context) *gorm.DB { return s.db }
|
||||
|
||||
func (s stubDBService) Named(_ string) *gorm.DB { return s.db }
|
||||
|
||||
// TestFindUserByFieldRecordRejectsUnlistedColumns pins the column allow-list. The
|
||||
// lookup column is interpolated into SQL, so an unlisted name must be refused before
|
||||
// any query is built rather than trusted because call sites happen to pass literals.
|
||||
func TestFindUserByFieldRecordRejectsUnlistedColumns(t *testing.T) {
|
||||
db, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
if err := db.Table("w_users").Create(map[string]any{"id": 77, "username": "seeded"}).Error; err != nil {
|
||||
t.Fatalf("seed user failed: %v", err)
|
||||
}
|
||||
|
||||
repository.SetDBServiceForTest(stubDBService{db: db})
|
||||
t.Cleanup(func() { repository.SetDBServiceForTest(nil) })
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
user, err := repository.FindUserByFieldRecord(ctx, "username", "seeded")
|
||||
if err != nil {
|
||||
t.Fatalf("allowlisted lookup by username failed: %v", err)
|
||||
}
|
||||
if user.ID != 77 {
|
||||
t.Errorf("allowlisted lookup returned ID %d, want 77", user.ID)
|
||||
}
|
||||
if _, err := repository.FindUserByFieldRecord(ctx, "id", uint64(77)); err != nil {
|
||||
t.Errorf("allowlisted lookup by id failed: %v", err)
|
||||
}
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
field string
|
||||
}{
|
||||
{"tautology injection", `username = '' OR 1=1 --`},
|
||||
{"stacked statement", "id; DROP TABLE w_users"},
|
||||
{"column outside allow-list", "password"},
|
||||
{"empty field", ""},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
if _, err := repository.FindUserByFieldRecord(ctx, tc.field, "seeded"); !errors.Is(err, errs.ErrUnsupportedUserLookupField) {
|
||||
t.Errorf("%s: got err %v, want ErrUnsupportedUserLookupField", tc.name, err)
|
||||
}
|
||||
}
|
||||
|
||||
var remaining int64
|
||||
if err := db.Table("w_users").Count(&remaining).Error; err != nil || remaining != 1 {
|
||||
t.Fatalf("w_users damaged by rejected lookups: count=%d err=%v", remaining, err)
|
||||
}
|
||||
}
|
||||
|
||||
// smtpTestValues are the four system-config rows the built-in email channel reads.
|
||||
var smtpTestValues = map[string]string{
|
||||
"smtp_host": "mail.example.test",
|
||||
"smtp_port": "465",
|
||||
"smtp_username": "notify@example.test",
|
||||
"smtp_password": "s3cret-value",
|
||||
}
|
||||
|
||||
// TestLoadSMTPConfigRecordMapsEveryKey guards the single-query rewrite: every field
|
||||
// must still be filled from its own row.
|
||||
func TestLoadSMTPConfigRecordMapsEveryKey(t *testing.T) {
|
||||
db, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
keys := make([]string, 0, len(smtpTestValues))
|
||||
for key := range smtpTestValues {
|
||||
keys = append(keys, key)
|
||||
}
|
||||
if err := db.Table("w_system_configs").Where("key IN ?", keys).Delete(map[string]any{}).Error; err != nil {
|
||||
t.Fatalf("clear smtp rows: %v", err)
|
||||
}
|
||||
for _, key := range keys {
|
||||
row := map[string]any{"key": key, "value": smtpTestValues[key], "type": "system"}
|
||||
if err := db.Table("w_system_configs").Create(row).Error; err != nil {
|
||||
t.Fatalf("seed %s: %v", key, err)
|
||||
}
|
||||
}
|
||||
|
||||
repository.SetDBServiceForTest(stubDBService{db: db})
|
||||
t.Cleanup(func() { repository.SetDBServiceForTest(nil) })
|
||||
|
||||
cfg, err := repository.LoadSMTPConfigRecord(context.Background())
|
||||
if err != nil {
|
||||
t.Fatalf("LoadSMTPConfigRecord: %v", err)
|
||||
}
|
||||
if cfg.Host != smtpTestValues["smtp_host"] || cfg.Port != smtpTestValues["smtp_port"] ||
|
||||
cfg.Username != smtpTestValues["smtp_username"] || cfg.Password != smtpTestValues["smtp_password"] {
|
||||
t.Errorf("got %+v, want every SMTP field mapped from its own row", cfg)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadSMTPConfigRecordSurfacesReadFailure pins the actual defect: a read that
|
||||
// fails used to be discarded, returning four blank strings that callers could only
|
||||
// interpret as "SMTP was never configured", so the notification was dropped silently.
|
||||
func TestLoadSMTPConfigRecordSurfacesReadFailure(t *testing.T) {
|
||||
bare, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatalf("open bare sqlite: %v", err)
|
||||
}
|
||||
|
||||
repository.SetDBServiceForTest(stubDBService{db: bare})
|
||||
t.Cleanup(func() { repository.SetDBServiceForTest(nil) })
|
||||
|
||||
if _, err := repository.LoadSMTPConfigRecord(context.Background()); err == nil {
|
||||
t.Fatal("LoadSMTPConfigRecord returned nil error although the config table cannot be read")
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,404 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package service implements domain business logic and channel runners for message_gateway.
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/message_gateway/repository"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
"unicode"
|
||||
)
|
||||
|
||||
// Handler processes one inbound message.
|
||||
type Handler func(ctx context.Context, msg model.InboundMessage) error
|
||||
|
||||
// Factory constructs a Channel from decrypted config.
|
||||
type Factory func(cfg model.ChannelConfig, onInbound Handler) (Channel, error)
|
||||
|
||||
// Channel is one connected messaging adapter.
|
||||
type Channel interface {
|
||||
Type() string
|
||||
Connect(ctx context.Context) error
|
||||
Disconnect(ctx context.Context) error
|
||||
Send(ctx context.Context, to model.Recipient, msg model.OutboundMessage) error
|
||||
Capabilities() model.Capability
|
||||
}
|
||||
|
||||
var (
|
||||
factoriesMu sync.RWMutex
|
||||
factories = map[string]Factory{}
|
||||
)
|
||||
|
||||
// Register stores a channel factory under typ.
|
||||
func Register(typ string, fn Factory) {
|
||||
factoriesMu.Lock()
|
||||
defer factoriesMu.Unlock()
|
||||
factories[typ] = fn
|
||||
}
|
||||
|
||||
// Lookup returns a previously registered factory.
|
||||
func Lookup(typ string) (Factory, bool) {
|
||||
factoriesMu.RLock()
|
||||
defer factoriesMu.RUnlock()
|
||||
fn, ok := factories[typ]
|
||||
return fn, ok
|
||||
}
|
||||
|
||||
// CodeAlphabet excludes easily confused runes 0/O/1/I.
|
||||
const CodeAlphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
|
||||
|
||||
// CodeLength is the raw pairing code size.
|
||||
const CodeLength = 8
|
||||
|
||||
// GenerateCode returns an 8-character pairing code.
|
||||
func GenerateCode() (string, error) {
|
||||
buf := make([]byte, CodeLength)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
out := make([]byte, CodeLength)
|
||||
for i, b := range buf {
|
||||
out[i] = CodeAlphabet[int(b)%len(CodeAlphabet)]
|
||||
}
|
||||
return string(out), nil
|
||||
}
|
||||
|
||||
// NormalizeCode strips separators and uppercases.
|
||||
func NormalizeCode(s string) string {
|
||||
var b strings.Builder
|
||||
for _, r := range s {
|
||||
if r == '-' || unicode.IsSpace(r) {
|
||||
continue
|
||||
}
|
||||
b.WriteRune(unicode.ToUpper(r))
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
// FormatCode renders ABCD-EFGH.
|
||||
func FormatCode(s string) string {
|
||||
s = NormalizeCode(s)
|
||||
if len(s) != CodeLength {
|
||||
return s
|
||||
}
|
||||
return s[:4] + "-" + s[4:]
|
||||
}
|
||||
|
||||
var (
|
||||
credentialSecretMu sync.RWMutex
|
||||
credentialSecret string
|
||||
)
|
||||
|
||||
// SetCredentialSecret sets the secret used to derive CredentialKey.
|
||||
func SetCredentialSecret(secret string) {
|
||||
credentialSecretMu.Lock()
|
||||
defer credentialSecretMu.Unlock()
|
||||
credentialSecret = secret
|
||||
}
|
||||
|
||||
// CredentialKey is AES-256 hex derived from the session secret.
|
||||
func CredentialKey() string {
|
||||
credentialSecretMu.RLock()
|
||||
secret := credentialSecret
|
||||
credentialSecretMu.RUnlock()
|
||||
sum := sha256.Sum256([]byte(secret))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
// EncryptCredentials encrypts a credential map as JSON.
|
||||
func EncryptCredentials(creds map[string]string) (string, error) {
|
||||
if creds == nil {
|
||||
creds = map[string]string{}
|
||||
}
|
||||
raw, err := json.Marshal(creds)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return util.Encrypt(CredentialKey(), string(raw))
|
||||
}
|
||||
|
||||
// DecryptCredentials decrypts a credential map.
|
||||
func DecryptCredentials(ciphertext string) (map[string]string, error) {
|
||||
if ciphertext == "" {
|
||||
return map[string]string{}, nil
|
||||
}
|
||||
plain, err := util.Decrypt(CredentialKey(), ciphertext)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var out map[string]string
|
||||
if err := json.Unmarshal([]byte(plain), &out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if out == nil {
|
||||
out = map[string]string{}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// ParseExtra decodes optional extra JSON into a string map.
|
||||
func ParseExtra(raw string) map[string]string {
|
||||
if raw == "" {
|
||||
return map[string]string{}
|
||||
}
|
||||
var out map[string]string
|
||||
if err := json.Unmarshal([]byte(raw), &out); err != nil || out == nil {
|
||||
return map[string]string{}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// EncodeExtra encodes extra fields as JSON.
|
||||
func EncodeExtra(extra map[string]string) string {
|
||||
if extra == nil {
|
||||
return ""
|
||||
}
|
||||
raw, err := json.Marshal(extra)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return string(raw)
|
||||
}
|
||||
|
||||
// Runner manages lifecycle for long-lived channel adapters (WebSocket, long-polling, etc.).
|
||||
type Runner struct {
|
||||
mu sync.Mutex
|
||||
running bool
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
// GlobalRunner is the default global runner instance.
|
||||
var GlobalRunner = &Runner{}
|
||||
|
||||
// Start starts all background long-lived channel runners.
|
||||
func Start(ctx context.Context) error {
|
||||
GlobalRunner.mu.Lock()
|
||||
defer GlobalRunner.mu.Unlock()
|
||||
|
||||
if GlobalRunner.running {
|
||||
return nil
|
||||
}
|
||||
|
||||
runCtx, cancel := context.WithCancel(ctx)
|
||||
GlobalRunner.cancel = cancel
|
||||
GlobalRunner.running = true
|
||||
|
||||
logger.InfoF(runCtx, "[MessageGateway] Starting bot channel runners...")
|
||||
return nil
|
||||
}
|
||||
|
||||
// Stop stops the channel runner.
|
||||
func Stop() {
|
||||
GlobalRunner.mu.Lock()
|
||||
defer GlobalRunner.mu.Unlock()
|
||||
|
||||
if !GlobalRunner.running {
|
||||
return
|
||||
}
|
||||
|
||||
if GlobalRunner.cancel != nil {
|
||||
GlobalRunner.cancel()
|
||||
}
|
||||
GlobalRunner.running = false
|
||||
}
|
||||
|
||||
// Cordis contract singletons consumed by service layer.
|
||||
var (
|
||||
cacheMu sync.RWMutex
|
||||
cacheSvc contracts.CacheService
|
||||
taskMu sync.RWMutex
|
||||
taskSvc contracts.TaskService
|
||||
userMu sync.RWMutex
|
||||
userSvc contracts.UserService
|
||||
)
|
||||
|
||||
// SetCacheService sets the cache service.
|
||||
func SetCacheService(s contracts.CacheService) {
|
||||
cacheMu.Lock()
|
||||
defer cacheMu.Unlock()
|
||||
cacheSvc = s
|
||||
}
|
||||
|
||||
// SetTaskService sets the task service.
|
||||
func SetTaskService(s contracts.TaskService) {
|
||||
taskMu.Lock()
|
||||
defer taskMu.Unlock()
|
||||
taskSvc = s
|
||||
}
|
||||
|
||||
// SetUserService sets the user service.
|
||||
func SetUserService(s contracts.UserService) {
|
||||
userMu.Lock()
|
||||
defer userMu.Unlock()
|
||||
userSvc = s
|
||||
}
|
||||
|
||||
// GetCache resolves the cache service for the context.
|
||||
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
|
||||
}
|
||||
|
||||
// GetTaskService returns the task service.
|
||||
func GetTaskService(ctx context.Context) contracts.TaskService {
|
||||
if s, err := core.InjectFrom[contracts.TaskService](ctx); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
taskMu.RLock()
|
||||
defer taskMu.RUnlock()
|
||||
return taskSvc
|
||||
}
|
||||
|
||||
// GetUserService resolves the user service for the context.
|
||||
func GetUserService(ctx context.Context) contracts.UserService {
|
||||
if s, err := core.InjectFrom[contracts.UserService](ctx); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
userMu.RLock()
|
||||
s := userSvc
|
||||
userMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
|
||||
// BindChannel consumes a pairing code and binds the platform identity to the user.
|
||||
func BindChannel(ctx context.Context, userID uint64, req model.BindRequest) (model.BindingDTO, error) {
|
||||
channelID, err := strconv.ParseUint(strings.TrimSpace(req.ChannelID), 10, 64)
|
||||
if err != nil || channelID == 0 {
|
||||
return model.BindingDTO{}, errs.ErrChannelIDRequired
|
||||
}
|
||||
code := NormalizeCode(req.Code)
|
||||
if code == "" {
|
||||
return model.BindingDTO{}, errs.ErrCodeInvalid
|
||||
}
|
||||
pairing, err := repository.GetPairingCode(ctx, code)
|
||||
if err != nil {
|
||||
if errors.Is(err, errs.ErrRecordNotFound) {
|
||||
return model.BindingDTO{}, errs.ErrCodeInvalid
|
||||
}
|
||||
return model.BindingDTO{}, err
|
||||
}
|
||||
if !pairing.ExpiresAt.After(time.Now()) {
|
||||
return model.BindingDTO{}, errs.ErrCodeInvalid
|
||||
}
|
||||
if pairing.ChannelID != channelID {
|
||||
return model.BindingDTO{}, errs.ErrChannelMismatch
|
||||
}
|
||||
ch, err := repository.GetMessageChannel(ctx, channelID)
|
||||
if err != nil {
|
||||
if errors.Is(err, errs.ErrRecordNotFound) {
|
||||
return model.BindingDTO{}, errs.ErrCodeInvalid
|
||||
}
|
||||
return model.BindingDTO{}, err
|
||||
}
|
||||
if !ch.Enabled {
|
||||
return model.BindingDTO{}, errs.ErrChannelDisabled
|
||||
}
|
||||
|
||||
existing, err := repository.GetBindingByChannelPlatform(ctx, channelID, pairing.PlatformUserID)
|
||||
if err != nil && !errors.Is(err, errs.ErrRecordNotFound) {
|
||||
return model.BindingDTO{}, err
|
||||
}
|
||||
if err == nil && existing != nil {
|
||||
if existing.UserID != userID {
|
||||
return model.BindingDTO{}, errs.ErrPlatformAlreadyBound
|
||||
}
|
||||
_ = repository.DeletePairingCode(ctx, pairing.Code)
|
||||
return ToBindingDTO(existing, ch), nil
|
||||
}
|
||||
|
||||
row := &model.MessageBinding{
|
||||
UserID: userID,
|
||||
ChannelID: channelID,
|
||||
PlatformUserID: pairing.PlatformUserID,
|
||||
}
|
||||
if err := repository.CreateMessageBinding(ctx, row); err != nil {
|
||||
return model.BindingDTO{}, err
|
||||
}
|
||||
if err := repository.DeletePairingCode(ctx, pairing.Code); err != nil {
|
||||
return model.BindingDTO{}, err
|
||||
}
|
||||
return ToBindingDTO(row, ch), nil
|
||||
}
|
||||
|
||||
// ListEnabledPublicChannels returns the channels a user may bind to.
|
||||
func ListEnabledPublicChannels(ctx context.Context) ([]model.PublicChannelDTO, error) {
|
||||
rows, err := repository.ListEnabledMessageChannels(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]model.PublicChannelDTO, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
out = append(out, model.PublicChannelDTO{ID: row.ID, Name: row.Name, Type: row.Type})
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// ListUserBindings returns the binding rows of one user enriched with channel info.
|
||||
func ListUserBindings(ctx context.Context, userID uint64) ([]model.BindingDTO, error) {
|
||||
rows, err := repository.ListBindingsByUser(ctx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]model.BindingDTO, 0, len(rows))
|
||||
for i := range rows {
|
||||
ch, err := repository.GetMessageChannel(ctx, rows[i].ChannelID)
|
||||
if err != nil {
|
||||
out = append(out, ToBindingDTO(&rows[i], nil))
|
||||
continue
|
||||
}
|
||||
out = append(out, ToBindingDTO(&rows[i], ch))
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// UnbindChannel removes a binding owned by the given user.
|
||||
func UnbindChannel(ctx context.Context, userID, bindingID uint64) error {
|
||||
row, err := repository.GetMessageBinding(ctx, bindingID)
|
||||
if err != nil {
|
||||
if errors.Is(err, errs.ErrRecordNotFound) {
|
||||
return errs.ErrBindingNotFound
|
||||
}
|
||||
return err
|
||||
}
|
||||
if row.UserID != userID {
|
||||
return errs.ErrBindingForbidden
|
||||
}
|
||||
return repository.DeleteMessageBinding(ctx, bindingID)
|
||||
}
|
||||
|
||||
// ToBindingDTO projects a binding row and its optional channel onto the user DTO.
|
||||
func ToBindingDTO(row *model.MessageBinding, ch *model.MessageChannel) model.BindingDTO {
|
||||
dto := model.BindingDTO{
|
||||
ID: row.ID,
|
||||
UserID: row.UserID,
|
||||
ChannelID: row.ChannelID,
|
||||
PlatformUserID: row.PlatformUserID,
|
||||
CreatedAt: row.CreatedAt,
|
||||
}
|
||||
if ch != nil {
|
||||
dto.ChannelName = ch.Name
|
||||
dto.ChannelType = ch.Type
|
||||
}
|
||||
return dto
|
||||
}
|
||||
+11
-10
@@ -7,8 +7,9 @@ package qq
|
||||
import (
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/message_gateway/service"
|
||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/do"
|
||||
"Wavelet/plugins/domain/msg_gateway/service"
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
@@ -33,7 +34,7 @@ type qqEvent struct {
|
||||
|
||||
// Adapter is an official QQ Bot C2C channel.
|
||||
type Adapter struct {
|
||||
cfg model.ChannelConfig
|
||||
cfg do.ChannelConfig
|
||||
onInbound service.Handler
|
||||
api openapi.OpenAPI
|
||||
tokenSrc oauth2.TokenSource
|
||||
@@ -43,7 +44,7 @@ type Adapter struct {
|
||||
}
|
||||
|
||||
// New constructs a QQ adapter.
|
||||
func New(cfg model.ChannelConfig, onInbound service.Handler) (service.Channel, error) {
|
||||
func New(cfg do.ChannelConfig, onInbound service.Handler) (service.Channel, error) {
|
||||
if strings.TrimSpace(cfg.Credentials["app_id"]) == "" || strings.TrimSpace(cfg.Credentials["app_secret"]) == "" {
|
||||
return nil, fmt.Errorf("qq: app_id and app_secret are required")
|
||||
}
|
||||
@@ -51,11 +52,11 @@ func New(cfg model.ChannelConfig, onInbound service.Handler) (service.Channel, e
|
||||
}
|
||||
|
||||
// Type returns qq.
|
||||
func (a *Adapter) Type() string { return model.ChannelTypeQQ }
|
||||
func (a *Adapter) Type() string { return consts.ChannelTypeQQ }
|
||||
|
||||
// Capabilities reports C2C text/media support.
|
||||
func (a *Adapter) Capabilities() model.Capability {
|
||||
return model.Capability{Text: true, Image: true, File: true, Reply: true}
|
||||
func (a *Adapter) Capabilities() do.Capability {
|
||||
return do.Capability{Text: true, Image: true, File: true, Reply: true}
|
||||
}
|
||||
|
||||
// Connect starts the official WebSocket session (C2C intent).
|
||||
@@ -128,7 +129,7 @@ func (a *Adapter) Disconnect(_ context.Context) error {
|
||||
}
|
||||
|
||||
// Send posts a C2C text reply.
|
||||
func (a *Adapter) Send(ctx context.Context, to model.Recipient, msg model.OutboundMessage) error {
|
||||
func (a *Adapter) Send(ctx context.Context, to do.Recipient, msg do.OutboundMessage) error {
|
||||
a.mu.Lock()
|
||||
api := a.api
|
||||
a.mu.Unlock()
|
||||
@@ -152,9 +153,9 @@ func (a *Adapter) handleEvent(ctx context.Context, ev qqEvent) {
|
||||
if disconnected || a.onInbound == nil {
|
||||
return
|
||||
}
|
||||
_ = a.onInbound(ctx, model.InboundMessage{
|
||||
_ = a.onInbound(ctx, do.InboundMessage{
|
||||
ChannelID: a.cfg.ID,
|
||||
ChannelType: model.ChannelTypeQQ,
|
||||
ChannelType: consts.ChannelTypeQQ,
|
||||
PlatformUserID: ev.UserID,
|
||||
ChatID: ev.UserID,
|
||||
MessageID: ev.MessageID,
|
||||
+5
-5
@@ -4,14 +4,14 @@
|
||||
package qq
|
||||
|
||||
import (
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/do"
|
||||
"context"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestHandleEvent_DropsNonC2C(t *testing.T) {
|
||||
var got int
|
||||
a := &Adapter{onInbound: func(ctx context.Context, msg model.InboundMessage) error {
|
||||
a := &Adapter{onInbound: func(_ context.Context, _ do.InboundMessage) error {
|
||||
got++
|
||||
return nil
|
||||
}}
|
||||
@@ -22,8 +22,8 @@ func TestHandleEvent_DropsNonC2C(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestHandleEvent_C2CText(t *testing.T) {
|
||||
var got model.InboundMessage
|
||||
a := &Adapter{cfg: model.ChannelConfig{ID: 3}, onInbound: func(ctx context.Context, msg model.InboundMessage) error {
|
||||
var got do.InboundMessage
|
||||
a := &Adapter{cfg: do.ChannelConfig{ID: 3}, onInbound: func(_ context.Context, msg do.InboundMessage) error {
|
||||
got = msg
|
||||
return nil
|
||||
}}
|
||||
@@ -34,7 +34,7 @@ func TestHandleEvent_C2CText(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestNew_RequiresCreds(t *testing.T) {
|
||||
_, err := New(model.ChannelConfig{}, nil)
|
||||
_, err := New(do.ChannelConfig{}, nil)
|
||||
if err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
+17
-16
@@ -7,8 +7,9 @@ package telegram
|
||||
import (
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/message_gateway/service"
|
||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/do"
|
||||
"Wavelet/plugins/domain/msg_gateway/service"
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
@@ -22,13 +23,13 @@ import (
|
||||
|
||||
// Adapter is a Telegram private-chat channel.
|
||||
type Adapter struct {
|
||||
cfg model.ChannelConfig
|
||||
cfg do.ChannelConfig
|
||||
onInbound service.Handler
|
||||
bot *tele.Bot
|
||||
}
|
||||
|
||||
// New constructs a Telegram adapter. Call service.Register from the runner.
|
||||
func New(cfg model.ChannelConfig, onInbound service.Handler) (service.Channel, error) {
|
||||
func New(cfg do.ChannelConfig, onInbound service.Handler) (service.Channel, error) {
|
||||
if strings.TrimSpace(cfg.Credentials["bot_token"]) == "" {
|
||||
return nil, fmt.Errorf("telegram: bot_token is required")
|
||||
}
|
||||
@@ -36,11 +37,11 @@ func New(cfg model.ChannelConfig, onInbound service.Handler) (service.Channel, e
|
||||
}
|
||||
|
||||
// Type returns telegram.
|
||||
func (a *Adapter) Type() string { return model.ChannelTypeTelegram }
|
||||
func (a *Adapter) Type() string { return consts.ChannelTypeTelegram }
|
||||
|
||||
// Capabilities reports private-chat media support.
|
||||
func (a *Adapter) Capabilities() model.Capability {
|
||||
return model.Capability{Text: true, Image: true, File: true, Reply: true}
|
||||
func (a *Adapter) Capabilities() do.Capability {
|
||||
return do.Capability{Text: true, Image: true, File: true, Reply: true}
|
||||
}
|
||||
|
||||
// longPollWindow is how long Telegram may hold a getUpdates call open before
|
||||
@@ -50,7 +51,7 @@ func (a *Adapter) Capabilities() model.Capability {
|
||||
const longPollWindow = 10 * time.Second
|
||||
|
||||
// buildTeleSettings assembles the telebot settings.
|
||||
func buildTeleSettings(cfg model.ChannelConfig) tele.Settings {
|
||||
func buildTeleSettings(cfg do.ChannelConfig) tele.Settings {
|
||||
pref := tele.Settings{
|
||||
Token: cfg.Credentials["bot_token"],
|
||||
Poller: &tele.LongPoller{Timeout: longPollWindow},
|
||||
@@ -99,7 +100,7 @@ func (a *Adapter) Disconnect(_ context.Context) error {
|
||||
}
|
||||
|
||||
// Send replies to a private chat.
|
||||
func (a *Adapter) Send(_ context.Context, to model.Recipient, msg model.OutboundMessage) error {
|
||||
func (a *Adapter) Send(_ context.Context, to do.Recipient, msg do.OutboundMessage) error {
|
||||
if a.bot == nil {
|
||||
return fmt.Errorf("telegram: not connected")
|
||||
}
|
||||
@@ -118,9 +119,9 @@ func (a *Adapter) handleTeleMessage(ctx context.Context, m *tele.Message) {
|
||||
if a.onInbound == nil {
|
||||
return
|
||||
}
|
||||
msg := model.InboundMessage{
|
||||
msg := do.InboundMessage{
|
||||
ChannelID: a.cfg.ID,
|
||||
ChannelType: model.ChannelTypeTelegram,
|
||||
ChannelType: consts.ChannelTypeTelegram,
|
||||
PlatformUserID: strconv.FormatInt(m.Sender.ID, 10),
|
||||
ChatID: strconv.FormatInt(m.Chat.ID, 10),
|
||||
MessageID: strconv.Itoa(m.ID),
|
||||
@@ -146,7 +147,7 @@ func (a *Adapter) handleTeleMessage(ctx context.Context, m *tele.Message) {
|
||||
// downloadMedia fetches message media into a scratch directory, returned so the
|
||||
// caller can remove it once the inbound handler no longer needs the paths.
|
||||
// An empty dir means nothing was downloaded.
|
||||
func (a *Adapter) downloadMedia(m *tele.Message) (string, []model.Attachment) {
|
||||
func (a *Adapter) downloadMedia(m *tele.Message) (string, []do.Attachment) {
|
||||
var files []*tele.File
|
||||
var names []string
|
||||
if m.Photo != nil {
|
||||
@@ -166,16 +167,16 @@ func (a *Adapter) downloadMedia(m *tele.Message) (string, []model.Attachment) {
|
||||
}
|
||||
dir, err := os.MkdirTemp("", "wg-tg-*")
|
||||
if err != nil {
|
||||
return "", []model.Attachment{{Error: err.Error()}}
|
||||
return "", []do.Attachment{{Error: err.Error()}}
|
||||
}
|
||||
out := make([]model.Attachment, 0, len(files))
|
||||
out := make([]do.Attachment, 0, len(files))
|
||||
for i, f := range files {
|
||||
path := filepath.Join(dir, names[i])
|
||||
if err := a.bot.Download(f, path); err != nil {
|
||||
out = append(out, model.Attachment{FileName: names[i], Error: err.Error()})
|
||||
out = append(out, do.Attachment{FileName: names[i], Error: err.Error()})
|
||||
continue
|
||||
}
|
||||
out = append(out, model.Attachment{Path: path, FileName: names[i]})
|
||||
out = append(out, do.Attachment{Path: path, FileName: names[i]})
|
||||
}
|
||||
return dir, out
|
||||
}
|
||||
+7
-7
@@ -4,7 +4,7 @@
|
||||
package telegram
|
||||
|
||||
import (
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/do"
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -16,7 +16,7 @@ import (
|
||||
// telebot 以 int(timeout/time.Second) 下发给 getUpdates。写成裸整数会被解释为
|
||||
// 纳秒,令 timeout=0,长轮询退化为对 Bot API 的空转轮询。
|
||||
func TestBuildTeleSettingsLongPollWindow(t *testing.T) {
|
||||
pref := buildTeleSettings(model.ChannelConfig{
|
||||
pref := buildTeleSettings(do.ChannelConfig{
|
||||
Credentials: map[string]string{"bot_token": "token"},
|
||||
Extra: map[string]string{"base_url": "https://tg.example.com/api/"},
|
||||
})
|
||||
@@ -35,7 +35,7 @@ func TestBuildTeleSettingsLongPollWindow(t *testing.T) {
|
||||
|
||||
func TestHandleUpdate_DropsGroups(t *testing.T) {
|
||||
var got int
|
||||
a := &Adapter{onInbound: func(ctx context.Context, msg model.InboundMessage) error {
|
||||
a := &Adapter{onInbound: func(_ context.Context, _ do.InboundMessage) error {
|
||||
got++
|
||||
return nil
|
||||
}}
|
||||
@@ -51,10 +51,10 @@ func TestHandleUpdate_DropsGroups(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestHandleUpdate_PrivateText(t *testing.T) {
|
||||
var got model.InboundMessage
|
||||
var got do.InboundMessage
|
||||
a := &Adapter{
|
||||
cfg: model.ChannelConfig{ID: 7, Type: "telegram"},
|
||||
onInbound: func(ctx context.Context, msg model.InboundMessage) error {
|
||||
cfg: do.ChannelConfig{ID: 7, Type: "telegram"},
|
||||
onInbound: func(_ context.Context, msg do.InboundMessage) error {
|
||||
got = msg
|
||||
return nil
|
||||
},
|
||||
@@ -71,7 +71,7 @@ func TestHandleUpdate_PrivateText(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestNew_RequiresToken(t *testing.T) {
|
||||
_, err := New(model.ChannelConfig{}, nil)
|
||||
_, err := New(do.ChannelConfig{}, nil)
|
||||
if err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package consts defines constants, error types, and identifiers for msg_gateway.
|
||||
package consts
|
||||
|
||||
// Bot Channel type and scope constants.
|
||||
const (
|
||||
ChannelTypeTelegram = "telegram"
|
||||
ChannelTypeQQ = "qq"
|
||||
MessageChannelTypeTelegram = "telegram"
|
||||
MessageChannelTypeQQ = "qq"
|
||||
MessageOwnerScopeSystem = "system"
|
||||
)
|
||||
|
||||
// Bot Task and Schedule identifier constants.
|
||||
const (
|
||||
TaskCleanupPairingCodes = "msg_gateway:cleanup_pairing_codes"
|
||||
TaskDispatchBotMsg = "msg_gateway:dispatch_bot_msg"
|
||||
TaskTypeDispatchBotMsg = "dispatch_bot_msg"
|
||||
)
|
||||
|
||||
// Pairing code constants.
|
||||
const (
|
||||
CodeAlphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
|
||||
CodeLength = 8
|
||||
)
|
||||
+8
-12
@@ -1,9 +1,7 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package errs defines error sentinels and user-facing error message constants
|
||||
// for the message_gateway plugin.
|
||||
package errs
|
||||
package consts
|
||||
|
||||
import "errors"
|
||||
|
||||
@@ -16,14 +14,15 @@ var (
|
||||
ErrBindingForbidden = errors.New("cannot unbind another user's binding")
|
||||
ErrChannelIDRequired = errors.New("channel_id is required")
|
||||
ErrChannelDisabled = errors.New("channel is not enabled")
|
||||
ErrChannelNotFound = errors.New("channel not found")
|
||||
ErrEventNotFound = errors.New("notification event not found")
|
||||
ErrUserNotFound = errors.New("user not found")
|
||||
ErrNoAdminUser = errors.New("no admin user found")
|
||||
ErrTaskServiceNotAvail = errors.New("task service not available")
|
||||
|
||||
// ErrRecordNotFound maps GORM's missing-row sentinel at the repository boundary so
|
||||
// ErrRecordNotFound maps GORM's missing-row sentinel at the DAO boundary so
|
||||
// upper layers never import gorm. Its text matches gorm.ErrRecordNotFound verbatim.
|
||||
ErrRecordNotFound = errors.New("record not found")
|
||||
|
||||
// ErrUnsupportedUserLookupField rejects a column name that the repository is not
|
||||
// allowed to interpolate into a WHERE clause.
|
||||
ErrUnsupportedUserLookupField = errors.New("unsupported user lookup field")
|
||||
)
|
||||
|
||||
// User-facing validation and error message constants.
|
||||
@@ -32,7 +31,7 @@ const (
|
||||
ErrTypeInvalid = "type must be telegram or qq"
|
||||
ErrTelegramTokenRequired = "telegram bot secret is required" //nolint:gosec // user-facing validation text
|
||||
ErrQQCredentialsRequired = "qq app id and secret are required" //nolint:gosec // user-facing validation text
|
||||
ErrChannelNotFound = "channel not found"
|
||||
ErrChannelNotFoundText = "channel not found"
|
||||
ErrChannelProbeFailed = "channel probe failed"
|
||||
ErrBotDispatchTextRequired = "message text is required"
|
||||
ErrBotChannelNotRegistered = "channel adapter is not registered"
|
||||
@@ -42,7 +41,6 @@ const (
|
||||
ErrInvalidBindingID = "invalid binding id"
|
||||
ErrInvalidChannelID = "invalid channel id"
|
||||
ErrInvalidEventID = "invalid event id"
|
||||
ErrEventNotFound = "notification event not found"
|
||||
ErrValidationFailed = "validation failed"
|
||||
|
||||
ErrMissingTelegramToken = "missing telegram bot token"
|
||||
@@ -63,8 +61,6 @@ const (
|
||||
ErrEventKeyOrTaskType = "either event_key or task_type must be provided"
|
||||
ErrUnsupportedEventKey = "unsupported built-in event key"
|
||||
ErrTaskServiceUnavailable = "task service not available"
|
||||
ErrUserNotFound = "user not found"
|
||||
ErrNoAdminUser = "no admin user found"
|
||||
|
||||
ErrPayloadRequired = "payload is required"
|
||||
ErrInvalidJSONFormat = "invalid json format"
|
||||
@@ -0,0 +1,43 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package consts
|
||||
|
||||
// Push notification channel type constants.
|
||||
const (
|
||||
TypeCustom = "custom"
|
||||
TypeEmail = "email"
|
||||
TypeTelegram = "telegram"
|
||||
ChannelCustom = "custom"
|
||||
ChannelEmail = "email"
|
||||
ChannelLark = "lark"
|
||||
ChannelDingTalk = "dingtalk"
|
||||
ChannelTelegram = "telegram"
|
||||
ChannelBark = "bark"
|
||||
ChannelDiscord = "discord"
|
||||
ChannelSlack = "slack"
|
||||
ChannelPushover = "pushover"
|
||||
)
|
||||
|
||||
// Push message template and payload keys.
|
||||
const (
|
||||
DefaultLevelInfo = "INFO"
|
||||
KeyTitle = "title"
|
||||
KeyContent = "content"
|
||||
KeyLevel = "level"
|
||||
|
||||
KeyURL = "url"
|
||||
KeyToken = "token"
|
||||
KeyOther = "other"
|
||||
|
||||
TypeText = "text"
|
||||
TypePassword = "password"
|
||||
TypeTextarea = "textarea"
|
||||
)
|
||||
|
||||
// Push task identifier constants.
|
||||
const (
|
||||
TaskPushNotification = "msg_gateway:push_notification"
|
||||
SendNotificationTask = "push:send"
|
||||
TaskTypeSendNotification = "send_notification"
|
||||
)
|
||||
+13
-18
@@ -1,14 +1,14 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package handler
|
||||
package controller
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_gateway/service"
|
||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||
"Wavelet/plugins/domain/msg_gateway/service"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
@@ -19,7 +19,7 @@ import (
|
||||
// @Tags admin-message-gateway
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]model.Definition}
|
||||
// @Success 200 {object} response.Any{data=[]do.Definition}
|
||||
// @Router /api/v1/admin/message-gateway/channels/definitions [get]
|
||||
func ListAdminChannelDefinitions(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(service.ListDefinitions()))
|
||||
@@ -31,7 +31,7 @@ func ListAdminChannelDefinitions(c *gin.Context) {
|
||||
// @Tags admin-message-gateway
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]model.ChannelDTO}
|
||||
// @Success 200 {object} response.Any{data=[]do.ChannelDTO}
|
||||
// @Router /api/v1/admin/message-gateway/channels [get]
|
||||
func ListAdminChannels(c *gin.Context) {
|
||||
rows, err := service.ListChannels(c.Request.Context())
|
||||
@@ -43,17 +43,12 @@ func ListAdminChannels(c *gin.Context) {
|
||||
}
|
||||
|
||||
func parseAdminChannelID(c *gin.Context) (uint64, bool) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, errs.ErrInvalidChannelID)
|
||||
return 0, false
|
||||
}
|
||||
return id, true
|
||||
return parseUint64Param(c, "id", consts.ErrInvalidChannelID)
|
||||
}
|
||||
|
||||
func handleAdminChannelError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
|
||||
if err.Error() == errs.ErrChannelNotFound {
|
||||
response.AbortNotFound(c, err.Error())
|
||||
if errors.Is(err, consts.ErrChannelNotFound) || err.Error() == consts.ErrChannelNotFoundText {
|
||||
response.AbortNotFound(c, consts.ErrChannelNotFoundText)
|
||||
return
|
||||
}
|
||||
fallback(c, err.Error())
|
||||
@@ -66,8 +61,8 @@ func handleAdminChannelError(c *gin.Context, err error, fallback func(c *gin.Con
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body model.CreateChannelRequest true "create body"
|
||||
// @Success 200 {object} response.Any{data=model.ChannelDTO}
|
||||
// @Param request body do.CreateChannelRequest true "create body"
|
||||
// @Success 200 {object} response.Any{data=do.ChannelDTO}
|
||||
// @Failure 400 {object} response.Any
|
||||
// @Router /api/v1/admin/message-gateway/channels [post]
|
||||
func CreateAdminChannel(c *gin.Context) {
|
||||
@@ -82,8 +77,8 @@ func CreateAdminChannel(c *gin.Context) {
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "channel id"
|
||||
// @Param request body model.UpdateChannelRequest true "update body"
|
||||
// @Success 200 {object} response.Any{data=model.ChannelDTO}
|
||||
// @Param request body do.UpdateChannelRequest true "update body"
|
||||
// @Success 200 {object} response.Any{data=do.ChannelDTO}
|
||||
// @Failure 400 {object} response.Any
|
||||
// @Failure 404 {object} response.Any
|
||||
// @Router /api/v1/admin/message-gateway/channels/{id} [patch]
|
||||
@@ -0,0 +1,69 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package controller
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"Wavelet/pkg/response"
|
||||
"context"
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// currentUser extracts the authenticated UserDTO from gin.Context.
|
||||
func currentUser(c *gin.Context) (*contracts.UserDTO, bool) {
|
||||
return ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
}
|
||||
|
||||
// parseUint64Param parses a uint64 URL path parameter.
|
||||
func parseUint64Param(c *gin.Context, paramName, errInvalid string) (uint64, bool) {
|
||||
id, err := strconv.ParseUint(c.Param(paramName), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, errInvalid)
|
||||
return 0, false
|
||||
}
|
||||
return id, true
|
||||
}
|
||||
|
||||
// handleJSONRequest binds a JSON body, executes the service handler, and writes the standard success envelope.
|
||||
func handleJSONRequest[Req any, Res any](c *gin.Context, handler func(ctx context.Context, req Req) (Res, error)) {
|
||||
var req Req
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
res, err := handler(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(res))
|
||||
}
|
||||
|
||||
// handleEntityUpdate resolves a path identifier and JSON body, executes the updater, and handles errors with onErr.
|
||||
func handleEntityUpdate[Req any, Res any](
|
||||
c *gin.Context,
|
||||
parseID func(*gin.Context) (uint64, bool),
|
||||
updater func(ctx context.Context, id uint64, req Req) (Res, error),
|
||||
onErr func(*gin.Context, error),
|
||||
) {
|
||||
id, ok := parseID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var req Req
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
dto, err := updater(c.Request.Context(), id, req)
|
||||
if err != nil {
|
||||
onErr(c, err)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(dto))
|
||||
}
|
||||
+15
-21
@@ -1,16 +1,15 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package handler
|
||||
package controller
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/message_gateway/service"
|
||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/do"
|
||||
"Wavelet/plugins/domain/msg_gateway/service"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
@@ -24,7 +23,7 @@ import (
|
||||
// @Success 200 {object} response.Any "通道配置定义列表"
|
||||
// @Router /api/v1/admin/push/channels/definitions [get]
|
||||
func ListPushChannelDefinitions(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(model.ListPushDefinitions()))
|
||||
c.JSON(http.StatusOK, response.OK(do.ListPushDefinitions()))
|
||||
}
|
||||
|
||||
// ListPushChannels 获取消息通道列表
|
||||
@@ -33,7 +32,7 @@ func ListPushChannelDefinitions(c *gin.Context) {
|
||||
// @Tags admin-push
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]model.PushChannel} "消息通道列表"
|
||||
// @Success 200 {object} response.Any{data=[]entity.PushChannel} "消息通道列表"
|
||||
// @Router /api/v1/admin/push/channels [get]
|
||||
func ListPushChannels(c *gin.Context) {
|
||||
channels, err := service.ListPushChannels(c.Request.Context())
|
||||
@@ -46,18 +45,13 @@ func ListPushChannels(c *gin.Context) {
|
||||
|
||||
// parsePushChannelID reads the path identifier of a push channel.
|
||||
func parsePushChannelID(c *gin.Context) (uint64, bool) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, errs.ErrInvalidChannelID)
|
||||
return 0, false
|
||||
}
|
||||
return id, true
|
||||
return parseUint64Param(c, "id", consts.ErrInvalidChannelID)
|
||||
}
|
||||
|
||||
// handlePushChannelNotFoundError maps a missing channel row to 404, others to fallback.
|
||||
func handlePushChannelNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
|
||||
if errors.Is(err, errs.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, errs.ErrChannelNotFound)
|
||||
if errors.Is(err, consts.ErrRecordNotFound) || errors.Is(err, consts.ErrChannelNotFound) || err.Error() == consts.ErrChannelNotFoundText {
|
||||
response.AbortNotFound(c, consts.ErrChannelNotFoundText)
|
||||
return
|
||||
}
|
||||
fallback(c, err.Error())
|
||||
@@ -70,8 +64,8 @@ func handlePushChannelNotFoundError(c *gin.Context, err error, fallback func(c *
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body model.CreatePushChannelRequest true "创建参数"
|
||||
// @Success 200 {object} response.Any{data=model.PushChannel} "创建成功"
|
||||
// @Param request body do.CreatePushChannelRequest true "创建参数"
|
||||
// @Success 200 {object} response.Any{data=entity.PushChannel} "创建成功"
|
||||
// @Router /api/v1/admin/push/channels [post]
|
||||
func CreatePushChannel(c *gin.Context) {
|
||||
handleJSONRequest(c, service.CreatePushChannel)
|
||||
@@ -85,8 +79,8 @@ func CreatePushChannel(c *gin.Context) {
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path uint64 true "通道ID"
|
||||
// @Param request body model.UpdatePushChannelRequest true "更新参数"
|
||||
// @Success 200 {object} response.Any{data=model.PushChannel} "更新成功"
|
||||
// @Param request body do.UpdatePushChannelRequest true "更新参数"
|
||||
// @Success 200 {object} response.Any{data=entity.PushChannel} "更新成功"
|
||||
// @Router /api/v1/admin/push/channels/{id} [put]
|
||||
func UpdatePushChannel(c *gin.Context) {
|
||||
handleEntityUpdate(c, parsePushChannelID, service.UpdatePushChannel, func(c *gin.Context, err error) {
|
||||
@@ -123,11 +117,11 @@ func DeletePushChannel(c *gin.Context) {
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body model.TestPushChannelRequest true "测试参数"
|
||||
// @Param request body do.TestPushChannelRequest true "测试参数"
|
||||
// @Success 200 {object} response.Any "测试触发成功"
|
||||
// @Router /api/v1/admin/push/channels/test [post]
|
||||
func TestPushChannel(c *gin.Context) {
|
||||
var req model.TestPushChannelRequest
|
||||
var req do.TestPushChannelRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
+16
-20
@@ -1,13 +1,13 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package handler
|
||||
package controller
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/message_gateway/service"
|
||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/do"
|
||||
"Wavelet/plugins/domain/msg_gateway/service"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
@@ -21,7 +21,7 @@ import (
|
||||
// @Tags admin-push
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]model.PushEvent} "通知事件列表"
|
||||
// @Success 200 {object} response.Any{data=[]entity.PushEvent} "通知事件列表"
|
||||
// @Router /api/v1/admin/push/events [get]
|
||||
func ListPushEvents(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
@@ -47,18 +47,13 @@ func ListBuiltInPushEvents(c *gin.Context) {
|
||||
|
||||
// parsePushEventID reads the path identifier of a push event.
|
||||
func parsePushEventID(c *gin.Context) (uint64, bool) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, errs.ErrInvalidEventID)
|
||||
return 0, false
|
||||
}
|
||||
return id, true
|
||||
return parseUint64Param(c, "id", consts.ErrInvalidEventID)
|
||||
}
|
||||
|
||||
// handlePushEventNotFoundError maps a missing event row to 404, others to fallback.
|
||||
func handlePushEventNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
|
||||
if errors.Is(err, errs.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, errs.ErrEventNotFound)
|
||||
if errors.Is(err, consts.ErrRecordNotFound) || errors.Is(err, consts.ErrEventNotFound) || err.Error() == consts.ErrEventNotFound.Error() {
|
||||
response.AbortNotFound(c, consts.ErrEventNotFound.Error())
|
||||
return
|
||||
}
|
||||
fallback(c, err.Error())
|
||||
@@ -71,8 +66,8 @@ func handlePushEventNotFoundError(c *gin.Context, err error, fallback func(c *gi
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body model.CreatePushEventRequest true "创建参数"
|
||||
// @Success 200 {object} response.Any{data=model.PushEvent} "创建成功"
|
||||
// @Param request body do.CreatePushEventRequest true "创建参数"
|
||||
// @Success 200 {object} response.Any{data=entity.PushEvent} "创建成功"
|
||||
// @Router /api/v1/admin/push/events [post]
|
||||
func CreatePushEvent(c *gin.Context) {
|
||||
handleJSONRequest(c, service.CreatePushEvent)
|
||||
@@ -108,7 +103,7 @@ func DeletePushEvent(c *gin.Context) {
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "事件 ID"
|
||||
// @Param request body model.UpdatePushEventRequest true "更新参数"
|
||||
// @Param request body do.UpdatePushEventRequest true "更新参数"
|
||||
// @Success 200 {object} response.Any{data=string} "修改成功"
|
||||
// @Router /api/v1/admin/push/events/{id} [put]
|
||||
func UpdatePushEvent(c *gin.Context) {
|
||||
@@ -117,7 +112,7 @@ func UpdatePushEvent(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
var req model.UpdatePushEventRequest
|
||||
var req do.UpdatePushEventRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
@@ -175,8 +170,9 @@ func ListPushHistories(c *gin.Context) {
|
||||
pageSize = 20
|
||||
}
|
||||
|
||||
total, results, err := service.ListPushHistories(c.Request.Context(), model.PushHistoryListFilter{
|
||||
total, results, err := service.ListPushHistories(c.Request.Context(), do.PushHistoryListFilter{
|
||||
EventKey: c.Query("event_key"),
|
||||
Channel: c.Query("channel"),
|
||||
Status: c.Query("status"),
|
||||
Page: page,
|
||||
PageSize: pageSize,
|
||||
@@ -199,11 +195,11 @@ func ListPushHistories(c *gin.Context) {
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body model.TestPushRequest true "测试请求体"
|
||||
// @Param request body do.TestPushRequest true "测试请求体"
|
||||
// @Success 200 {object} response.Any{data=string} "测试成功"
|
||||
// @Router /api/v1/admin/push/test [post]
|
||||
func TestPush(c *gin.Context) {
|
||||
var req model.TestPushRequest
|
||||
var req do.TestPushRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
+2
-1
@@ -1,7 +1,8 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package handler
|
||||
// Package controller provides HTTP endpoints for msg_gateway.
|
||||
package controller
|
||||
|
||||
import (
|
||||
"Wavelet/core/extpoints"
|
||||
+18
-69
@@ -1,81 +1,31 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package handler provides HTTP endpoints for message_gateway.
|
||||
package handler
|
||||
package controller
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/message_gateway/service"
|
||||
"context"
|
||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/do"
|
||||
"Wavelet/plugins/domain/msg_gateway/service"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func currentUser(c *gin.Context) (*contracts.UserDTO, bool) {
|
||||
return ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
}
|
||||
|
||||
// handleJSONRequest binds a JSON body, runs the service use case and writes the
|
||||
// standard success envelope; any service error surfaces as a bad request.
|
||||
func handleJSONRequest[Req any, Res any](c *gin.Context, handler func(ctx context.Context, req Req) (Res, error)) {
|
||||
var req Req
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
res, err := handler(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(res))
|
||||
}
|
||||
|
||||
// handleEntityUpdate resolves a path identifier plus JSON body, runs the updater
|
||||
// use case and writes the success envelope; error classification is delegated to onErr.
|
||||
func handleEntityUpdate[Req any, Res any](
|
||||
c *gin.Context,
|
||||
parseID func(*gin.Context) (uint64, bool),
|
||||
updater func(ctx context.Context, id uint64, req Req) (Res, error),
|
||||
onErr func(*gin.Context, error),
|
||||
) {
|
||||
id, ok := parseID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var req Req
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
dto, err := updater(c.Request.Context(), id, req)
|
||||
if err != nil {
|
||||
onErr(c, err)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(dto))
|
||||
}
|
||||
|
||||
// ListChannels lists enabled channels a user can bind.
|
||||
// @Summary List enabled messaging channels
|
||||
// @Description Returns enabled system bots the current user can pair with
|
||||
// @Tags message-gateway
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]model.PublicChannelDTO}
|
||||
// @Success 200 {object} response.Any{data=[]do.PublicChannelDTO}
|
||||
// @Failure 401 {object} response.Any
|
||||
// @Router /api/v1/message-gateway/channels [get]
|
||||
func ListChannels(c *gin.Context) {
|
||||
if user, ok := currentUser(c); !ok || user == nil {
|
||||
response.AbortUnauthorized(c, errs.ErrLoginRequired)
|
||||
response.AbortUnauthorized(c, consts.ErrLoginRequired)
|
||||
return
|
||||
}
|
||||
rows, err := service.ListEnabledPublicChannels(c.Request.Context())
|
||||
@@ -92,13 +42,13 @@ func ListChannels(c *gin.Context) {
|
||||
// @Tags message-gateway
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]model.BindingDTO}
|
||||
// @Success 200 {object} response.Any{data=[]do.BindingDTO}
|
||||
// @Failure 401 {object} response.Any
|
||||
// @Router /api/v1/message-gateway/bindings [get]
|
||||
func ListBindings(c *gin.Context) {
|
||||
user, ok := currentUser(c)
|
||||
if !ok || user == nil {
|
||||
response.AbortUnauthorized(c, errs.ErrLoginRequired)
|
||||
response.AbortUnauthorized(c, consts.ErrLoginRequired)
|
||||
return
|
||||
}
|
||||
rows, err := service.ListUserBindings(c.Request.Context(), user.ID)
|
||||
@@ -116,25 +66,25 @@ func ListBindings(c *gin.Context) {
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body model.BindRequest true "bind body"
|
||||
// @Success 200 {object} response.Any{data=model.BindingDTO}
|
||||
// @Param request body do.BindRequest true "bind body"
|
||||
// @Success 200 {object} response.Any{data=do.BindingDTO}
|
||||
// @Failure 400 {object} response.Any
|
||||
// @Failure 409 {object} response.Any
|
||||
// @Router /api/v1/message-gateway/bindings [post]
|
||||
func BindBinding(c *gin.Context) {
|
||||
user, ok := currentUser(c)
|
||||
if !ok || user == nil {
|
||||
response.AbortUnauthorized(c, errs.ErrLoginRequired)
|
||||
response.AbortUnauthorized(c, consts.ErrLoginRequired)
|
||||
return
|
||||
}
|
||||
var req model.BindRequest
|
||||
var req do.BindRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
dto, err := service.BindChannel(c.Request.Context(), user.ID, req)
|
||||
if err != nil {
|
||||
if errors.Is(err, errs.ErrPlatformAlreadyBound) {
|
||||
if errors.Is(err, consts.ErrPlatformAlreadyBound) {
|
||||
response.AbortConflict(c, err.Error())
|
||||
return
|
||||
}
|
||||
@@ -158,20 +108,19 @@ func BindBinding(c *gin.Context) {
|
||||
func UnbindBinding(c *gin.Context) {
|
||||
user, ok := currentUser(c)
|
||||
if !ok || user == nil {
|
||||
response.AbortUnauthorized(c, errs.ErrLoginRequired)
|
||||
response.AbortUnauthorized(c, consts.ErrLoginRequired)
|
||||
return
|
||||
}
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, errs.ErrInvalidBindingID)
|
||||
id, ok := parseUint64Param(c, "id", consts.ErrInvalidBindingID)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := service.UnbindChannel(c.Request.Context(), user.ID, id); err != nil {
|
||||
if errors.Is(err, errs.ErrBindingNotFound) {
|
||||
if errors.Is(err, consts.ErrBindingNotFound) {
|
||||
response.AbortNotFound(c, err.Error())
|
||||
return
|
||||
}
|
||||
if errors.Is(err, errs.ErrBindingForbidden) {
|
||||
if errors.Is(err, consts.ErrBindingForbidden) {
|
||||
response.AbortForbidden(c, err.Error())
|
||||
return
|
||||
}
|
||||
+30
-78
@@ -1,68 +1,20 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package repository provides data persistence for the message_gateway plugin.
|
||||
package repository
|
||||
package dao
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var (
|
||||
dbMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
)
|
||||
|
||||
// SetDBServiceForTest injects a DBService for tests. Production wiring must use Apply.
|
||||
func SetDBServiceForTest(s contracts.DBService) {
|
||||
SetDBService(s)
|
||||
}
|
||||
|
||||
// SetDBService sets the database service singleton.
|
||||
func SetDBService(s contracts.DBService) {
|
||||
dbMu.Lock()
|
||||
defer dbMu.Unlock()
|
||||
dbSvc = s
|
||||
}
|
||||
|
||||
// GetDB resolves the persistence handle for the current call, preferring an
|
||||
// explicitly injected *core.Context before falling back to the plugin singleton.
|
||||
func GetDB(ctx context.Context) *gorm.DB {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
}
|
||||
dbMu.RLock()
|
||||
s := dbSvc
|
||||
dbMu.RUnlock()
|
||||
if s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// mapNotFound translates GORM's missing-row sentinel into the plugin-level
|
||||
// errs.ErrRecordNotFound so the service and handler layers stay free of gorm imports.
|
||||
func mapNotFound(err error) error {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return errs.ErrRecordNotFound
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// CreateMessageChannel inserts a channel row.
|
||||
func CreateMessageChannel(ctx context.Context, ch *model.MessageChannel) error {
|
||||
func CreateMessageChannel(ctx context.Context, ch *entity.MessageChannel) error {
|
||||
if ch.ID == 0 {
|
||||
ch.ID = idgen.NextUint64ID()
|
||||
}
|
||||
@@ -70,13 +22,13 @@ func CreateMessageChannel(ctx context.Context, ch *model.MessageChannel) error {
|
||||
}
|
||||
|
||||
// UpdateMessageChannel saves a channel row.
|
||||
func UpdateMessageChannel(ctx context.Context, ch *model.MessageChannel) error {
|
||||
func UpdateMessageChannel(ctx context.Context, ch *entity.MessageChannel) error {
|
||||
return GetDB(ctx).Save(ch).Error
|
||||
}
|
||||
|
||||
// GetMessageChannel loads a channel by id.
|
||||
func GetMessageChannel(ctx context.Context, id uint64) (*model.MessageChannel, error) {
|
||||
var ch model.MessageChannel
|
||||
func GetMessageChannel(ctx context.Context, id uint64) (*entity.MessageChannel, error) {
|
||||
var ch entity.MessageChannel
|
||||
if err := GetDB(ctx).Where("id = ?", id).First(&ch).Error; err != nil {
|
||||
return nil, mapNotFound(err)
|
||||
}
|
||||
@@ -84,8 +36,8 @@ func GetMessageChannel(ctx context.Context, id uint64) (*model.MessageChannel, e
|
||||
}
|
||||
|
||||
// ListMessageChannels returns all channels newest first.
|
||||
func ListMessageChannels(ctx context.Context) ([]model.MessageChannel, error) {
|
||||
var rows []model.MessageChannel
|
||||
func ListMessageChannels(ctx context.Context) ([]entity.MessageChannel, error) {
|
||||
var rows []entity.MessageChannel
|
||||
if err := GetDB(ctx).Order("id DESC").Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -95,18 +47,18 @@ func ListMessageChannels(ctx context.Context) ([]model.MessageChannel, error) {
|
||||
// DeleteMessageChannel removes pairings, bindings, then the channel.
|
||||
func DeleteMessageChannel(ctx context.Context, id uint64) error {
|
||||
return GetDB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("channel_id = ?", id).Delete(&model.MessagePairingCode{}).Error; err != nil {
|
||||
if err := tx.Where("channel_id = ?", id).Delete(&entity.MessagePairingCode{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("channel_id = ?", id).Delete(&model.MessageBinding{}).Error; err != nil {
|
||||
if err := tx.Where("channel_id = ?", id).Delete(&entity.MessageBinding{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Delete(&model.MessageChannel{}, id).Error
|
||||
return tx.Delete(&entity.MessageChannel{}, id).Error
|
||||
})
|
||||
}
|
||||
|
||||
// CreateMessageBinding inserts a binding.
|
||||
func CreateMessageBinding(ctx context.Context, b *model.MessageBinding) error {
|
||||
func CreateMessageBinding(ctx context.Context, b *entity.MessageBinding) error {
|
||||
if b.ID == 0 {
|
||||
b.ID = idgen.NextUint64ID()
|
||||
}
|
||||
@@ -114,8 +66,8 @@ func CreateMessageBinding(ctx context.Context, b *model.MessageBinding) error {
|
||||
}
|
||||
|
||||
// GetBindingByChannelPlatform finds a binding for a platform user on a channel.
|
||||
func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*model.MessageBinding, error) {
|
||||
var b model.MessageBinding
|
||||
func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*entity.MessageBinding, error) {
|
||||
var b entity.MessageBinding
|
||||
err := GetDB(ctx).Where("channel_id = ? AND platform_user_id = ?", channelID, platformUserID).First(&b).Error
|
||||
if err != nil {
|
||||
return nil, mapNotFound(err)
|
||||
@@ -124,8 +76,8 @@ func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platform
|
||||
}
|
||||
|
||||
// ListBindingsByUser lists bindings for a Wavelet user.
|
||||
func ListBindingsByUser(ctx context.Context, userID uint64) ([]model.MessageBinding, error) {
|
||||
var rows []model.MessageBinding
|
||||
func ListBindingsByUser(ctx context.Context, userID uint64) ([]entity.MessageBinding, error) {
|
||||
var rows []entity.MessageBinding
|
||||
if err := GetDB(ctx).Where("user_id = ?", userID).Order("id DESC").Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -133,8 +85,8 @@ func ListBindingsByUser(ctx context.Context, userID uint64) ([]model.MessageBind
|
||||
}
|
||||
|
||||
// ListBindingsByChannel lists bindings on one messaging channel.
|
||||
func ListBindingsByChannel(ctx context.Context, channelID uint64) ([]model.MessageBinding, error) {
|
||||
var rows []model.MessageBinding
|
||||
func ListBindingsByChannel(ctx context.Context, channelID uint64) ([]entity.MessageBinding, error) {
|
||||
var rows []entity.MessageBinding
|
||||
if err := GetDB(ctx).Where("channel_id = ?", channelID).Order("id DESC").Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -142,8 +94,8 @@ func ListBindingsByChannel(ctx context.Context, channelID uint64) ([]model.Messa
|
||||
}
|
||||
|
||||
// GetMessageBinding loads a binding by id.
|
||||
func GetMessageBinding(ctx context.Context, id uint64) (*model.MessageBinding, error) {
|
||||
var b model.MessageBinding
|
||||
func GetMessageBinding(ctx context.Context, id uint64) (*entity.MessageBinding, error) {
|
||||
var b entity.MessageBinding
|
||||
if err := GetDB(ctx).Where("id = ?", id).First(&b).Error; err != nil {
|
||||
return nil, mapNotFound(err)
|
||||
}
|
||||
@@ -152,12 +104,12 @@ func GetMessageBinding(ctx context.Context, id uint64) (*model.MessageBinding, e
|
||||
|
||||
// DeleteMessageBinding deletes a binding by id.
|
||||
func DeleteMessageBinding(ctx context.Context, id uint64) error {
|
||||
return GetDB(ctx).Delete(&model.MessageBinding{}, id).Error
|
||||
return GetDB(ctx).Delete(&entity.MessageBinding{}, id).Error
|
||||
}
|
||||
|
||||
// UpsertPairingCode reuses an unexpired code for the same channel+platform user.
|
||||
func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*model.MessagePairingCode, error) {
|
||||
var existing model.MessagePairingCode
|
||||
func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*entity.MessagePairingCode, error) {
|
||||
var existing entity.MessagePairingCode
|
||||
err := GetDB(ctx).
|
||||
Where("channel_id = ? AND platform_user_id = ? AND expires_at > ?", channelID, platformUserID, time.Now()).
|
||||
First(&existing).Error
|
||||
@@ -167,7 +119,7 @@ func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, co
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
row := &model.MessagePairingCode{
|
||||
row := &entity.MessagePairingCode{
|
||||
Code: code,
|
||||
ChannelID: channelID,
|
||||
PlatformUserID: platformUserID,
|
||||
@@ -180,8 +132,8 @@ func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, co
|
||||
}
|
||||
|
||||
// GetPairingCode loads a pairing code by normalized code string.
|
||||
func GetPairingCode(ctx context.Context, code string) (*model.MessagePairingCode, error) {
|
||||
var row model.MessagePairingCode
|
||||
func GetPairingCode(ctx context.Context, code string) (*entity.MessagePairingCode, error) {
|
||||
var row entity.MessagePairingCode
|
||||
if err := GetDB(ctx).Where("code = ?", code).First(&row).Error; err != nil {
|
||||
return nil, mapNotFound(err)
|
||||
}
|
||||
@@ -190,17 +142,17 @@ func GetPairingCode(ctx context.Context, code string) (*model.MessagePairingCode
|
||||
|
||||
// DeletePairingCode removes a pairing code.
|
||||
func DeletePairingCode(ctx context.Context, code string) error {
|
||||
return GetDB(ctx).Where("code = ?", code).Delete(&model.MessagePairingCode{}).Error
|
||||
return GetDB(ctx).Where("code = ?", code).Delete(&entity.MessagePairingCode{}).Error
|
||||
}
|
||||
|
||||
// DeleteExpiredPairingCodes removes expired pairing rows.
|
||||
func DeleteExpiredPairingCodes(ctx context.Context) error {
|
||||
return GetDB(ctx).Where("expires_at <= ?", time.Now()).Delete(&model.MessagePairingCode{}).Error
|
||||
return GetDB(ctx).Where("expires_at <= ?", time.Now()).Delete(&entity.MessagePairingCode{}).Error
|
||||
}
|
||||
|
||||
// ListEnabledMessageChannels returns enabled channels.
|
||||
func ListEnabledMessageChannels(ctx context.Context) ([]model.MessageChannel, error) {
|
||||
var rows []model.MessageChannel
|
||||
func ListEnabledMessageChannels(ctx context.Context) ([]entity.MessageChannel, error) {
|
||||
var rows []entity.MessageChannel
|
||||
if err := GetDB(ctx).Where("enabled = ?", true).Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -0,0 +1,64 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package dao_test
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/msg_gateway/dao"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestBotDAO_ChannelAndBinding(t *testing.T) {
|
||||
_ = idgen.Init(1)
|
||||
db, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
require.NoError(t, db.AutoMigrate(&entity.MessageChannel{}, &entity.MessageBinding{}, &entity.MessagePairingCode{}))
|
||||
|
||||
dao.SetDBServiceForTest(stubDBService{db: db})
|
||||
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
ch := entity.MessageChannel{
|
||||
Name: "tg_bot",
|
||||
Type: "telegram",
|
||||
OwnerScope: "system",
|
||||
Credentials: "encrypted_token",
|
||||
Enabled: true,
|
||||
}
|
||||
require.NoError(t, dao.CreateMessageChannel(ctx, &ch))
|
||||
assert.NotZero(t, ch.ID)
|
||||
|
||||
code, err := dao.UpsertPairingCode(ctx, ch.ID, "tg_user_1", "ABCD1234", time.Now().Add(10*time.Minute))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "ABCD1234", code.Code)
|
||||
|
||||
// Reusing pairing code
|
||||
code2, err := dao.UpsertPairingCode(ctx, ch.ID, "tg_user_1", "XYZ9999", time.Now().Add(10*time.Minute))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "ABCD1234", code2.Code)
|
||||
|
||||
binding := entity.MessageBinding{
|
||||
UserID: 42,
|
||||
ChannelID: ch.ID,
|
||||
PlatformUserID: "tg_user_1",
|
||||
}
|
||||
require.NoError(t, dao.CreateMessageBinding(ctx, &binding))
|
||||
assert.NotZero(t, binding.ID)
|
||||
|
||||
bindings, err := dao.ListBindingsByUser(ctx, 42)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, bindings, 1)
|
||||
|
||||
require.NoError(t, dao.DeleteMessageChannel(ctx, ch.ID))
|
||||
_, err = dao.GetMessageChannel(ctx, ch.ID)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package dao provides database persistence and caching for the msg_gateway plugin.
|
||||
package dao
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/msg_gateway/consts"
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var (
|
||||
dbMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
cacheMu sync.RWMutex
|
||||
cacheSvc contracts.CacheService
|
||||
)
|
||||
|
||||
// SetDBService sets the database service singleton.
|
||||
func SetDBService(s contracts.DBService) {
|
||||
dbMu.Lock()
|
||||
defer dbMu.Unlock()
|
||||
dbSvc = s
|
||||
}
|
||||
|
||||
// SetDBServiceForTest injects a DBService for tests.
|
||||
func SetDBServiceForTest(s contracts.DBService) {
|
||||
SetDBService(s)
|
||||
}
|
||||
|
||||
// SetCacheService sets the cache service singleton.
|
||||
func SetCacheService(s contracts.CacheService) {
|
||||
cacheMu.Lock()
|
||||
defer cacheMu.Unlock()
|
||||
cacheSvc = s
|
||||
}
|
||||
|
||||
// GetDB resolves the persistence handle for the current call.
|
||||
func GetDB(ctx context.Context) *gorm.DB {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
}
|
||||
dbMu.RLock()
|
||||
s := dbSvc
|
||||
dbMu.RUnlock()
|
||||
if s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetCache resolves the cache service for the current call.
|
||||
func GetCache(ctx context.Context) contracts.CacheService {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
}
|
||||
cacheMu.RLock()
|
||||
s := cacheSvc
|
||||
cacheMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
|
||||
// mapNotFound translates GORM's missing-row sentinel into the plugin-level
|
||||
// consts.ErrRecordNotFound so the service and controller layers stay free of gorm imports.
|
||||
func mapNotFound(err error) error {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return consts.ErrRecordNotFound
|
||||
}
|
||||
return err
|
||||
}
|
||||
+48
-145
@@ -1,17 +1,12 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
package dao
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/message_gateway/errs"
|
||||
"Wavelet/plugins/domain/message_gateway/model"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/do"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
@@ -22,34 +17,9 @@ const (
|
||||
activePushEventCacheTTL = 24 * time.Hour
|
||||
)
|
||||
|
||||
var (
|
||||
cacheMu sync.RWMutex
|
||||
cacheSvc contracts.CacheService
|
||||
)
|
||||
|
||||
// SetCacheService sets the cache service singleton.
|
||||
func SetCacheService(s contracts.CacheService) {
|
||||
cacheMu.Lock()
|
||||
defer cacheMu.Unlock()
|
||||
cacheSvc = s
|
||||
}
|
||||
|
||||
// GetCache resolves the cache service for the current call.
|
||||
func GetCache(ctx context.Context) contracts.CacheService {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
}
|
||||
cacheMu.RLock()
|
||||
s := cacheSvc
|
||||
cacheMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
|
||||
// ListPushChannelsRecord returns all push channels ordered by creation time descending.
|
||||
func ListPushChannelsRecord(ctx context.Context) ([]model.PushChannel, error) {
|
||||
var channels []model.PushChannel
|
||||
func ListPushChannelsRecord(ctx context.Context) ([]entity.PushChannel, error) {
|
||||
var channels []entity.PushChannel
|
||||
if err := GetDB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -57,17 +27,17 @@ func ListPushChannelsRecord(ctx context.Context) ([]model.PushChannel, error) {
|
||||
}
|
||||
|
||||
// GetPushChannelByIDRecord loads a push channel by primary key.
|
||||
func GetPushChannelByIDRecord(ctx context.Context, id uint64) (model.PushChannel, error) {
|
||||
var channel model.PushChannel
|
||||
func GetPushChannelByIDRecord(ctx context.Context, id uint64) (entity.PushChannel, error) {
|
||||
var channel entity.PushChannel
|
||||
if err := GetDB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
|
||||
return model.PushChannel{}, mapNotFound(err)
|
||||
return entity.PushChannel{}, mapNotFound(err)
|
||||
}
|
||||
return channel, nil
|
||||
}
|
||||
|
||||
// GetPushChannelByNameRecord loads a push channel by its unique name.
|
||||
func GetPushChannelByNameRecord(ctx context.Context, name string) (*model.PushChannel, error) {
|
||||
var channel model.PushChannel
|
||||
func GetPushChannelByNameRecord(ctx context.Context, name string) (*entity.PushChannel, error) {
|
||||
var channel entity.PushChannel
|
||||
if err := GetDB(ctx).Where("name = ?", name).First(&channel).Error; err != nil {
|
||||
return nil, mapNotFound(err)
|
||||
}
|
||||
@@ -77,14 +47,14 @@ func GetPushChannelByNameRecord(ctx context.Context, name string) (*model.PushCh
|
||||
// CountPushChannelsByNameRecord returns how many channels share the given name.
|
||||
func CountPushChannelsByNameRecord(ctx context.Context, name string) (int64, error) {
|
||||
var count int64
|
||||
if err := GetDB(ctx).Model(&model.PushChannel{}).Where("name = ?", name).Count(&count).Error; err != nil {
|
||||
if err := GetDB(ctx).Model(&entity.PushChannel{}).Where("name = ?", name).Count(&count).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// CreatePushChannelRecord persists a new channel and invalidates cache.
|
||||
func CreatePushChannelRecord(ctx context.Context, channel *model.PushChannel) error {
|
||||
func CreatePushChannelRecord(ctx context.Context, channel *entity.PushChannel) error {
|
||||
if err := GetDB(ctx).Create(channel).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -93,7 +63,7 @@ func CreatePushChannelRecord(ctx context.Context, channel *model.PushChannel) er
|
||||
}
|
||||
|
||||
// SavePushChannelRecord updates a channel and invalidates cache.
|
||||
func SavePushChannelRecord(ctx context.Context, channel *model.PushChannel) error {
|
||||
func SavePushChannelRecord(ctx context.Context, channel *entity.PushChannel) error {
|
||||
if err := GetDB(ctx).Save(channel).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -102,7 +72,7 @@ func SavePushChannelRecord(ctx context.Context, channel *model.PushChannel) erro
|
||||
}
|
||||
|
||||
// DeletePushChannelRecord removes a channel and invalidates cache.
|
||||
func DeletePushChannelRecord(ctx context.Context, channel *model.PushChannel) error {
|
||||
func DeletePushChannelRecord(ctx context.Context, channel *entity.PushChannel) error {
|
||||
if err := GetDB(ctx).Delete(channel).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -111,14 +81,15 @@ func DeletePushChannelRecord(ctx context.Context, channel *model.PushChannel) er
|
||||
}
|
||||
|
||||
func getCachedOrQuery[T any](ctx context.Context, cacheKey string, ttl time.Duration, query func(db *gorm.DB, dest *T) error) (*T, error) {
|
||||
var val T
|
||||
if cache := GetCache(ctx); cache != nil {
|
||||
var val T
|
||||
if err := cache.Get(ctx, cacheKey, &val); err == nil {
|
||||
return &val, nil
|
||||
}
|
||||
}
|
||||
|
||||
db := GetDB(ctx)
|
||||
var val T
|
||||
if err := query(db, &val); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -131,8 +102,8 @@ func getCachedOrQuery[T any](ctx context.Context, cacheKey string, ttl time.Dura
|
||||
}
|
||||
|
||||
// GetActivePushChannelByName loads an enabled push channel, preferring the cache layer.
|
||||
func GetActivePushChannelByName(ctx context.Context, name string) (*model.PushChannel, error) {
|
||||
channel, err := getCachedOrQuery(ctx, "push:channel:active:"+name, activePushChannelCacheTTL, func(db *gorm.DB, dest *model.PushChannel) error {
|
||||
func GetActivePushChannelByName(ctx context.Context, name string) (*entity.PushChannel, error) {
|
||||
channel, err := getCachedOrQuery(ctx, "push:channel:active:"+name, activePushChannelCacheTTL, func(db *gorm.DB, dest *entity.PushChannel) error {
|
||||
return db.Where("name = ? AND enabled = ?", name, true).First(dest).Error
|
||||
})
|
||||
if err != nil {
|
||||
@@ -149,8 +120,8 @@ func DeleteActivePushChannelCache(ctx context.Context, name string) {
|
||||
}
|
||||
|
||||
// ListPushEventsRecord returns all push events ordered by creation time descending.
|
||||
func ListPushEventsRecord(ctx context.Context) ([]model.PushEvent, error) {
|
||||
var events []model.PushEvent
|
||||
func ListPushEventsRecord(ctx context.Context) ([]entity.PushEvent, error) {
|
||||
var events []entity.PushEvent
|
||||
if err := GetDB(ctx).Order("created_at DESC").Find(&events).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -158,19 +129,19 @@ func ListPushEventsRecord(ctx context.Context) ([]model.PushEvent, error) {
|
||||
}
|
||||
|
||||
// GetPushEventByIDRecord loads a push event by primary key.
|
||||
func GetPushEventByIDRecord(ctx context.Context, id uint64) (model.PushEvent, error) {
|
||||
var event model.PushEvent
|
||||
func GetPushEventByIDRecord(ctx context.Context, id uint64) (entity.PushEvent, error) {
|
||||
var event entity.PushEvent
|
||||
if err := GetDB(ctx).First(&event, id).Error; err != nil {
|
||||
return model.PushEvent{}, mapNotFound(err)
|
||||
return entity.PushEvent{}, mapNotFound(err)
|
||||
}
|
||||
return event, nil
|
||||
}
|
||||
|
||||
// GetPushEventByKeyRecord loads a push event by event key.
|
||||
func GetPushEventByKeyRecord(ctx context.Context, key string) (model.PushEvent, error) {
|
||||
var event model.PushEvent
|
||||
func GetPushEventByKeyRecord(ctx context.Context, key string) (entity.PushEvent, error) {
|
||||
var event entity.PushEvent
|
||||
if err := GetDB(ctx).Where("event_key = ?", key).First(&event).Error; err != nil {
|
||||
return model.PushEvent{}, mapNotFound(err)
|
||||
return entity.PushEvent{}, mapNotFound(err)
|
||||
}
|
||||
return event, nil
|
||||
}
|
||||
@@ -178,14 +149,14 @@ func GetPushEventByKeyRecord(ctx context.Context, key string) (model.PushEvent,
|
||||
// CountPushEventsByKeyRecord returns how many events use the given event key.
|
||||
func CountPushEventsByKeyRecord(ctx context.Context, key string) (int64, error) {
|
||||
var count int64
|
||||
if err := GetDB(ctx).Model(&model.PushEvent{}).Where("event_key = ?", key).Count(&count).Error; err != nil {
|
||||
if err := GetDB(ctx).Model(&entity.PushEvent{}).Where("event_key = ?", key).Count(&count).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// CreatePushEventRecord persists a new push event and invalidates cache.
|
||||
func CreatePushEventRecord(ctx context.Context, event *model.PushEvent) error {
|
||||
func CreatePushEventRecord(ctx context.Context, event *entity.PushEvent) error {
|
||||
if err := GetDB(ctx).Create(event).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -194,7 +165,7 @@ func CreatePushEventRecord(ctx context.Context, event *model.PushEvent) error {
|
||||
}
|
||||
|
||||
// SavePushEventRecord updates a push event and invalidates cache.
|
||||
func SavePushEventRecord(ctx context.Context, event *model.PushEvent) error {
|
||||
func SavePushEventRecord(ctx context.Context, event *entity.PushEvent) error {
|
||||
if err := GetDB(ctx).Save(event).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -203,7 +174,7 @@ func SavePushEventRecord(ctx context.Context, event *model.PushEvent) error {
|
||||
}
|
||||
|
||||
// UpdatePushEventEnabledRecord toggles the enabled flag for a push event.
|
||||
func UpdatePushEventEnabledRecord(ctx context.Context, event *model.PushEvent, enabled bool) error {
|
||||
func UpdatePushEventEnabledRecord(ctx context.Context, event *entity.PushEvent, enabled bool) error {
|
||||
event.Enabled = enabled
|
||||
if err := GetDB(ctx).Model(event).Update("enabled", enabled).Error; err != nil {
|
||||
return err
|
||||
@@ -213,7 +184,7 @@ func UpdatePushEventEnabledRecord(ctx context.Context, event *model.PushEvent, e
|
||||
}
|
||||
|
||||
// DeletePushEventRecord removes a push event and invalidates cache.
|
||||
func DeletePushEventRecord(ctx context.Context, event *model.PushEvent) error {
|
||||
func DeletePushEventRecord(ctx context.Context, event *entity.PushEvent) error {
|
||||
if err := GetDB(ctx).Delete(event).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -222,8 +193,8 @@ func DeletePushEventRecord(ctx context.Context, event *model.PushEvent) error {
|
||||
}
|
||||
|
||||
// ListActivePushEventsByTaskTypeRecord returns enabled events bound to a task type.
|
||||
func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string) ([]model.PushEvent, error) {
|
||||
var events []model.PushEvent
|
||||
func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string) ([]entity.PushEvent, error) {
|
||||
var events []entity.PushEvent
|
||||
if err := GetDB(ctx).Where("task_type = ? AND enabled = ?", taskType, true).Find(&events).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -231,8 +202,8 @@ func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string)
|
||||
}
|
||||
|
||||
// GetActivePushEventByKey loads an enabled push event, preferring the cache layer.
|
||||
func GetActivePushEventByKey(ctx context.Context, key string) (*model.PushEvent, error) {
|
||||
event, err := getCachedOrQuery(ctx, "push:event:active:"+key, activePushEventCacheTTL, func(db *gorm.DB, dest *model.PushEvent) error {
|
||||
func GetActivePushEventByKey(ctx context.Context, key string) (*entity.PushEvent, error) {
|
||||
event, err := getCachedOrQuery(ctx, "push:event:active:"+key, activePushEventCacheTTL, func(db *gorm.DB, dest *entity.PushEvent) error {
|
||||
return db.Where("event_key = ? AND enabled = ?", key, true).First(dest).Error
|
||||
})
|
||||
if err != nil {
|
||||
@@ -249,11 +220,14 @@ func DeleteActivePushEventCache(ctx context.Context, key string) {
|
||||
}
|
||||
|
||||
// ListPushHistoriesRecord returns paginated push history records.
|
||||
func ListPushHistoriesRecord(ctx context.Context, filter model.PushHistoryListFilter) (int64, []model.PushHistory, error) {
|
||||
query := GetDB(ctx).Model(&model.PushHistory{}).Order("created_at DESC")
|
||||
func ListPushHistoriesRecord(ctx context.Context, filter do.PushHistoryListFilter) (int64, []entity.PushHistory, error) {
|
||||
query := GetDB(ctx).Model(&entity.PushHistory{}).Order("created_at DESC")
|
||||
if filter.EventKey != "" {
|
||||
query = query.Where("event_key = ?", filter.EventKey)
|
||||
}
|
||||
if filter.Channel != "" {
|
||||
query = query.Where("channel = ?", filter.Channel)
|
||||
}
|
||||
if filter.Status != "" {
|
||||
query = query.Where("status = ?", filter.Status)
|
||||
}
|
||||
@@ -263,7 +237,7 @@ func ListPushHistoriesRecord(ctx context.Context, filter model.PushHistoryListFi
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
var results []model.PushHistory
|
||||
var results []entity.PushHistory
|
||||
offset := (filter.Page - 1) * filter.PageSize
|
||||
if err := query.Offset(offset).Limit(filter.PageSize).Find(&results).Error; err != nil {
|
||||
return 0, nil, err
|
||||
@@ -273,92 +247,21 @@ func ListPushHistoriesRecord(ctx context.Context, filter model.PushHistoryListFi
|
||||
}
|
||||
|
||||
// CreatePushHistoryRecord persists a push history audit record.
|
||||
func CreatePushHistoryRecord(ctx context.Context, history *model.PushHistory) error {
|
||||
func CreatePushHistoryRecord(ctx context.Context, history *entity.PushHistory) error {
|
||||
return GetDB(ctx).Create(history).Error
|
||||
}
|
||||
|
||||
// PushHistoryQuery returns a scoped query builder for push histories.
|
||||
func PushHistoryQuery(ctx context.Context) *gorm.DB {
|
||||
return GetDB(ctx).Model(&model.PushHistory{})
|
||||
return GetDB(ctx).Model(&entity.PushHistory{})
|
||||
}
|
||||
|
||||
// smtpConfigKeys are the system-config rows backing the built-in email channel.
|
||||
var smtpConfigKeys = []string{"smtp_host", "smtp_port", "smtp_username", "smtp_password"}
|
||||
|
||||
// LoadSMTPConfigRecord reads the SMTP settings in one query.
|
||||
//
|
||||
// A key that is simply absent leaves its field empty, which is how an unconfigured
|
||||
// mailer is represented. A read that fails is returned as an error, so callers
|
||||
// cannot mistake an unhealthy database for "no SMTP configured" and silently drop
|
||||
// the notification.
|
||||
func LoadSMTPConfigRecord(ctx context.Context) (model.SMTPConfig, error) {
|
||||
// DeletePushHistoriesBeforeRecord deletes push history records created before cutoff time.
|
||||
func DeletePushHistoriesBeforeRecord(ctx context.Context, cutoff time.Time) (int64, error) {
|
||||
db := GetDB(ctx)
|
||||
if db == nil {
|
||||
return model.SMTPConfig{}, errors.New("database not available")
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
var rows []struct {
|
||||
Key string
|
||||
Value string
|
||||
}
|
||||
if err := db.Table("w_system_configs").
|
||||
Select("key", "value").
|
||||
Where("key IN ?", smtpConfigKeys).
|
||||
Find(&rows).Error; err != nil {
|
||||
return model.SMTPConfig{}, fmt.Errorf("read smtp system configs: %w", err)
|
||||
}
|
||||
|
||||
var cfg model.SMTPConfig
|
||||
for _, row := range rows {
|
||||
switch row.Key {
|
||||
case "smtp_host":
|
||||
cfg.Host = row.Value
|
||||
case "smtp_port":
|
||||
cfg.Port = row.Value
|
||||
case "smtp_username":
|
||||
cfg.Username = row.Value
|
||||
case "smtp_password":
|
||||
cfg.Password = row.Value
|
||||
}
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
// userLookupColumns allow-lists the columns FindUserByFieldRecord may filter on.
|
||||
// The column name is concatenated into the WHERE clause, so anything not listed
|
||||
// here must never reach the database.
|
||||
var userLookupColumns = map[string]struct{}{
|
||||
"id": {},
|
||||
"username": {},
|
||||
}
|
||||
|
||||
// FindUserByFieldRecord is the user lookup fallback for when the UserService
|
||||
// contract is not wired yet. field must be one of userLookupColumns.
|
||||
func FindUserByFieldRecord(ctx context.Context, field string, value any) (*contracts.UserDTO, error) {
|
||||
if _, ok := userLookupColumns[field]; !ok {
|
||||
return nil, errs.ErrUnsupportedUserLookupField
|
||||
}
|
||||
db := GetDB(ctx)
|
||||
if db == nil {
|
||||
return nil, errs.ErrRecordNotFound
|
||||
}
|
||||
var user contracts.UserDTO
|
||||
if err := db.Table("w_users").Where(field+" = ?", value).First(&user).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
// FindFirstAdminUserRecord is the admin lookup fallback for when the UserService
|
||||
// contract is not wired yet.
|
||||
func FindFirstAdminUserRecord(ctx context.Context) (*contracts.UserDTO, error) {
|
||||
db := GetDB(ctx)
|
||||
if db == nil {
|
||||
return nil, errs.ErrRecordNotFound
|
||||
}
|
||||
var adminUser contracts.UserDTO
|
||||
if err := db.Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &adminUser, nil
|
||||
result := db.Where("created_at < ?", cutoff).Delete(&entity.PushHistory{})
|
||||
return result.RowsAffected, result.Error
|
||||
}
|
||||
@@ -0,0 +1,145 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package dao_test
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/msg_gateway/dao"
|
||||
"Wavelet/plugins/domain/msg_gateway/model/entity"
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type stubDBService struct{ db *gorm.DB }
|
||||
|
||||
func (s stubDBService) GORM() *gorm.DB { return s.db }
|
||||
func (s stubDBService) DB(_ context.Context) *gorm.DB { return s.db }
|
||||
func (s stubDBService) Named(_ string) *gorm.DB { return s.db }
|
||||
|
||||
func TestPushChannelDAO_CRUD(t *testing.T) {
|
||||
_ = idgen.Init(1)
|
||||
db, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
require.NoError(t, db.AutoMigrate(&entity.PushChannel{}, &entity.PushEvent{}, &entity.PushHistory{}))
|
||||
|
||||
dao.SetDBServiceForTest(stubDBService{db: db})
|
||||
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
ch := entity.PushChannel{
|
||||
Name: "test_webhook",
|
||||
Type: "custom",
|
||||
URL: "https://example.com/hook",
|
||||
Enabled: true,
|
||||
}
|
||||
require.NoError(t, dao.CreatePushChannelRecord(ctx, &ch))
|
||||
assert.NotZero(t, ch.ID)
|
||||
|
||||
loaded, err := dao.GetPushChannelByIDRecord(ctx, ch.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "test_webhook", loaded.Name)
|
||||
|
||||
active, err := dao.GetActivePushChannelByName(ctx, "test_webhook")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, ch.ID, active.ID)
|
||||
|
||||
channels, err := dao.ListPushChannelsRecord(ctx)
|
||||
require.NoError(t, err)
|
||||
assert.NotEmpty(t, channels)
|
||||
|
||||
require.NoError(t, dao.DeletePushChannelRecord(ctx, &ch))
|
||||
_, err = dao.GetPushChannelByIDRecord(ctx, ch.ID)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestPushEventDAO_CRUD(t *testing.T) {
|
||||
_ = idgen.Init(1)
|
||||
db, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
require.NoError(t, db.AutoMigrate(&entity.PushChannel{}, &entity.PushEvent{}, &entity.PushHistory{}))
|
||||
|
||||
dao.SetDBServiceForTest(stubDBService{db: db})
|
||||
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
ev := entity.PushEvent{
|
||||
EventKey: "test_event",
|
||||
Name: "测试事件",
|
||||
Channels: []string{"test_webhook"},
|
||||
Targets: []string{"admin"},
|
||||
Template: `{"title":"Hello"}`,
|
||||
Enabled: true,
|
||||
}
|
||||
require.NoError(t, dao.CreatePushEventRecord(ctx, &ev))
|
||||
assert.NotZero(t, ev.ID)
|
||||
|
||||
loaded, err := dao.GetPushEventByKeyRecord(ctx, "test_event")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "测试事件", loaded.Name)
|
||||
|
||||
require.NoError(t, dao.UpdatePushEventEnabledRecord(ctx, &ev, false))
|
||||
loadedDisabled, err := dao.GetPushEventByIDRecord(ctx, ev.ID)
|
||||
require.NoError(t, err)
|
||||
assert.False(t, loadedDisabled.Enabled)
|
||||
|
||||
require.NoError(t, dao.DeletePushEventRecord(ctx, &ev))
|
||||
_, err = dao.GetPushEventByIDRecord(ctx, ev.ID)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestPushHistoryDAO_Cleanup(t *testing.T) {
|
||||
_ = idgen.Init(1)
|
||||
db, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
require.NoError(t, db.AutoMigrate(&entity.PushHistory{}))
|
||||
|
||||
dao.SetDBServiceForTest(stubDBService{db: db})
|
||||
t.Cleanup(func() { dao.SetDBServiceForTest(nil) })
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
now := time.Now()
|
||||
oldTime := now.Add(-40 * 24 * time.Hour)
|
||||
recentTime := now.Add(-5 * 24 * time.Hour)
|
||||
|
||||
oldHistory := entity.PushHistory{
|
||||
EventKey: "login",
|
||||
Channel: "telegram",
|
||||
Target: "123",
|
||||
Title: "Old login",
|
||||
Content: "Old content",
|
||||
Level: "info",
|
||||
Status: "success",
|
||||
CreatedAt: oldTime,
|
||||
}
|
||||
recentHistory := entity.PushHistory{
|
||||
EventKey: "login",
|
||||
Channel: "telegram",
|
||||
Target: "123",
|
||||
Title: "Recent login",
|
||||
Content: "Recent content",
|
||||
Level: "info",
|
||||
Status: "success",
|
||||
CreatedAt: recentTime,
|
||||
}
|
||||
require.NoError(t, db.Create(&oldHistory).Error)
|
||||
require.NoError(t, db.Create(&recentHistory).Error)
|
||||
|
||||
cutoff := now.Add(-30 * 24 * time.Hour)
|
||||
deleted, err := dao.DeletePushHistoriesBeforeRecord(ctx, cutoff)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(1), deleted)
|
||||
|
||||
var count int64
|
||||
db.Model(&entity.PushHistory{}).Count(&count)
|
||||
assert.Equal(t, int64(1), count)
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user