merge(wavelet): sync upstream changes

This commit is contained in:
ryan
2026-09-03 09:29:28 +08:00
220 changed files with 12924 additions and 15487 deletions
+12 -1
View File
@@ -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",
+3 -1
View File
@@ -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
}
@@ -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 -8
View File
@@ -3,12 +3,7 @@
package auth
import "Wavelet/plugins/domain/auth/service"
// SessionConfig defines the session configuration declared by the auth plugin.
type SessionConfig struct {
SessionCookieName string `config:"session_cookie_name" env:"APP_SESSION_COOKIE_NAME" default:"wavelet_session"`
SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"`
SessionDomain string `config:"session_domain" env:"APP_SESSION_DOMAIN"`
SessionAge int `config:"session_age" env:"APP_SESSION_AGE" default:"86400"`
SessionHTTPOnly bool `config:"session_http_only" env:"APP_SESSION_HTTP_ONLY" default:"true"`
SessionSecure bool `config:"session_secure" env:"APP_SESSION_SECURE"`
}
type SessionConfig = service.SessionConfig
+52
View File
@@ -0,0 +1,52 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package consts defines constants, keys, and TTL values for the auth domain plugin.
package consts
import "time"
// CAP 默认参数
const (
DefaultCapChallengeCount = 1
DefaultCapChallengeSize = 32
DefaultCapChallengeDifficulty = 4
DefaultCapChallengeTTL = 10 * time.Minute
DefaultCapTokenTTL = 20 * time.Minute
RedeemTokenIDLength = 8 // 兑换 Token ID 字节长度
RedeemVerTokenLength = 15 // 兑换验证 Token 字节长度
TokenPartsCount = 2 // 兑换 Token 由两部分组成 (id:token)
ValuePartsCount = 2 // 存储值由 scope 和过期时间组成 (expNano|scope)
)
// CAP 动态配置键常量
const (
ConfigKeyCapLoginEnabled = "cap_login_enabled"
ConfigKeyCapChallengeCount = "cap_challenge_count"
ConfigKeyCapChallengeSize = "cap_challenge_size"
ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty"
ConfigKeyCapChallengeTTL = "cap_challenge_ttl"
// ConfigKeyCapTokenTTL 验证码 Token 过期时间键
// #nosec G101
ConfigKeyCapTokenTTL = "cap_token_ttl"
)
// HTTP 响应错误文案
const (
ErrCapTokenMissing = "验证码验证失败,缺少验证码凭证" //nolint:gosec // error message constant
ErrCapTokenInvalidOrExpired = "验证码校验失败或已过期,请重试" //nolint:gosec // error message constant
ErrCapNotConfigured = "captcha is not configured"
ErrChallengeGenerateFailed = "生成验证难题失败,请稍后再试"
ErrInvalidRequestParams = "无效的参数"
ErrSolutionVerifyFailed = "校验验证解答失败,请稍后再试"
)
// Redeem 结果码,属于 redeem 响应 JSON 的对外契约取值,禁止改写取值
const (
RedeemErrInvalidToken = "invalid_token"
RedeemErrNonceStoreFailed = "nonce_store_error"
RedeemErrAlreadyRedeemed = "already_redeemed"
RedeemErrSettingsLoad = "settings_load_error"
RedeemErrTokenStoreFailed = "token_store_error" //nolint:gosec // error code constant
)
@@ -1,11 +1,10 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// Package consts defines constants, keys, and TTL values for the auth domain plugin.
package consts
import (
"time"
)
import "time"
// Session and Context Keys
const (
@@ -14,7 +13,7 @@ const (
UserObjKey = "user_obj"
TokenAuthKey = "token_auth" // 标记当前请求是否通过 Access Token 鉴权
TokenAdminKey = "token_admin" // Access Token 本身是否具有管理员权限
SessionTokenKey = "oauth_session_token" //nolint:gosec // false positive: this is a session key, not hardcoded credentials
SessionTokenKey = "oauth_session_token" //nolint:gosec // false positive: session state key
PasswordHashKey = "password_hash"
SystemUsername = "system"
)
@@ -23,8 +22,8 @@ const (
const (
OAuthStateCacheKeyFormat = "oauth:state:%s"
OAuthStateCacheKeyExpiration = 10 * time.Minute
oauthStateLimitKeyFormat = "oauth:state:limit:%s"
oauthStateLimitMax = 10
OAuthStateLimitKeyFormat = "oauth:state:limit:%s"
OAuthStateLimitMax = 10
)
// OAuth Purpose Constants
@@ -37,3 +36,9 @@ const (
const (
AuthSourceTypeOIDC = "oidc"
)
// Cache TTLs
const (
TokenCacheTTL = 5 * time.Minute
UserCacheTTL = 5 * time.Minute
)
@@ -0,0 +1,53 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package consts defines constants, keys, and TTL values for the auth domain plugin.
package consts
// OAuth and Auth error messages
const (
ErrInvalidState = "非法登录请求"
ErrIDTokenVerifyFailed = "ID Token 验证失败" //nolint:gosec // error message constant
ErrIDTokenVerifyFailedFormat = "%s: %w"
ErrNonceMismatch = "nonce 不匹配,可能存在重放攻击"
ErrNoActiveAuthSource = "未配置可用认证源"
ErrServerAddressMissing = "服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试"
ErrAuthSourceRequired = "认证源不能为空"
ErrDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
ErrUsernameGenerateFailed = "无法生成可用用户名"
ErrUsernameFromSourceFailed = "无法从认证源获取用户名"
ErrAuthSourceDisabled = "认证源未启用"
ErrInvalidExternalAccountBindingID = "绑定记录 ID 无效"
ErrTokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // error message constant
ErrOAuthStateRateLimited = "请求授权过于频繁,请稍后重试"
ErrAuthSourceNameRequired = "认证源名称不能为空"
ErrAuthSourceNameInvalid = "认证源名称格式不正确"
ErrAuthSourceTypeUnsupported = "不支持的认证源类型"
ErrAuthSourceDiscoveryURLRequired = "Discovery URL 不能为空"
//nolint:gosec // error message constant
ErrAuthSourceClientCredentialsRequired = "启用认证源时必须配置 Client ID 和 Client Secret"
ErrAuthSourceIDRequired = "认证源 ID 不能为空"
ErrUserIDRequired = "用户 ID 不能为空"
ErrExternalAccountBindingIncomplete = "外部帐号绑定信息不完整"
ErrExternalAccountAlreadyBoundToAnother = "该外部帐号已被其他用户绑定"
ErrExternalAccountBindingIDRequired = "外部帐号绑定记录 ID 不能为空"
ErrInsufficientPermission = "权限不足"
ErrBannedAccount = "账号已被封禁"
ErrUnAuthorized = "未登录"
)
// Service 层与鉴权中间件内部错误文案
const (
ErrUserNotInContext = "auth: user not found in context"
ErrEmptyToken = "auth: empty token" //nolint:gosec // error message constant
ErrSystemUserTokenNotAllowed = "auth: system user token not allowed" //nolint:gosec // error message constant
ErrUnauthorizedInternal = "unauthorized"
ErrSystemUserLoginNotAllowed = "system user is not allowed to login"
)
// OAuth 回调会话校验错误文案
const (
ErrInvalidSessionContext = "invalid session context"
ErrSessionMismatchForOAuth = "session mismatch for oauth state"
ErrUserContextMismatch = "user context mismatch for oauth binding"
)
@@ -1,11 +1,13 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// Package controller provides HTTP handlers and middlewares for the auth plugin.
package controller
import (
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/plugins/domain/auth/model/dto"
"context"
"encoding/json"
@@ -17,7 +19,7 @@ func LogForAudit(ctx context.Context, user *contracts.UserDTO, c *gin.Context) {
if user == nil || c == nil {
return
}
auditLog := loginRequiredAuditLog{
auditLog := dto.LoginRequiredAuditLog{
UserID: user.ID,
Username: user.Username,
ClientIP: c.ClientIP(),
@@ -1,44 +1,59 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package 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()
}
+61
View File
@@ -0,0 +1,61 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package dao provides data access objects and caching for the auth domain plugin.
package dao
import (
"Wavelet/core/contracts"
"context"
"gorm.io/gorm"
)
// DAO aggregates all data access objects for the auth domain plugin.
type DAO struct {
dbSvc contracts.DBService
cacheSvc contracts.CacheService
limiterSvc contracts.LimiterService
}
// New creates a new DAO aggregate.
func New(dbSvc contracts.DBService, cacheSvc contracts.CacheService, limiterSvc contracts.LimiterService) *DAO {
return &DAO{
dbSvc: dbSvc,
cacheSvc: cacheSvc,
limiterSvc: limiterSvc,
}
}
// SetDBService updates the DBService reference.
func (d *DAO) SetDBService(db contracts.DBService) {
d.dbSvc = db
}
// SetCacheService updates the CacheService reference.
func (d *DAO) SetCacheService(cache contracts.CacheService) {
d.cacheSvc = cache
}
// SetLimiterService updates the LimiterService reference.
func (d *DAO) SetLimiterService(limiter contracts.LimiterService) {
d.limiterSvc = limiter
}
// DB returns the GORM DB instance associated with the request context.
func (d *DAO) DB(ctx context.Context) *gorm.DB {
if d.dbSvc != nil {
return d.dbSvc.DB(ctx)
}
return nil
}
// Cache returns the CacheService instance.
func (d *DAO) Cache() contracts.CacheService {
return d.cacheSvc
}
// Limiter returns the LimiterService instance.
func (d *DAO) Limiter() contracts.LimiterService {
return d.limiterSvc
}
@@ -0,0 +1,38 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package dao provides data access objects and caching for the auth domain plugin.
package dao
import (
"Wavelet/plugins/domain/auth/model/entity"
"context"
)
// FindExternalAccount 查询指定认证源的外部账号绑定
func (d *DAO) FindExternalAccount(ctx context.Context, authSourceID uint64, externalID string) (*entity.ExternalAccount, error) {
var account entity.ExternalAccount
if err := d.DB(ctx).Where("auth_source_id = ? AND external_id = ?", authSourceID, externalID).First(&account).Error; err != nil {
return nil, err
}
return &account, nil
}
// BindExternalAccount 绑定外部账号
func (d *DAO) BindExternalAccount(ctx context.Context, account *entity.ExternalAccount) error {
return d.DB(ctx).Create(account).Error
}
// ListExternalAccountsByUserID 获取用户绑定的所有外部账号
func (d *DAO) ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]entity.ExternalAccount, error) {
var accounts []entity.ExternalAccount
if err := d.DB(ctx).Where("user_id = ?", userID).Find(&accounts).Error; err != nil {
return nil, err
}
return accounts, nil
}
// UnbindExternalAccount 解绑外部账号
func (d *DAO) UnbindExternalAccount(ctx context.Context, id, userID uint64) error {
return d.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&entity.ExternalAccount{}).Error
}
@@ -0,0 +1,91 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package dao provides data access objects and caching for the auth domain plugin.
package dao
import (
"Wavelet/core/contracts"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/auth/model/do"
"context"
"time"
)
// GetAccessTokenByHash 按令牌哈希读取访问令牌记录(仅取鉴权所需字段)
func (d *DAO) GetAccessTokenByHash(ctx context.Context, tokenHash string) (*do.CachedToken, error) {
var row struct {
ID uint64
UserID uint64
IsAdmin bool
}
if err := d.DB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&row).Error; err != nil {
return nil, err
}
return &do.CachedToken{
ID: row.ID,
UserID: row.UserID,
IsAdmin: row.IsAdmin,
}, nil
}
// GetActiveUserByID 读取仍处于启用状态的用户
func (d *DAO) GetActiveUserByID(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
var user contracts.UserDTO
if err := d.DB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&user).Error; err != nil {
return nil, err
}
return &user, nil
}
// GetUserByID 按 ID 读取用户(不限制启用状态)
func (d *DAO) GetUserByID(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
var user contracts.UserDTO
if err := d.DB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil {
return nil, err
}
return &user, nil
}
// InsertUser 新建用户记录
func (d *DAO) InsertUser(ctx context.Context, user *contracts.UserDTO) error {
return d.DB(ctx).Table("w_users").Create(user).Error
}
// TouchUserLastLogin 刷新用户最后登录时间
func (d *DAO) TouchUserLastLogin(ctx context.Context, userID uint64, at time.Time) error {
return d.DB(ctx).Table("w_users").Where("id = ?", userID).Update("last_login_at", at).Error
}
// ListSimilarUsernames 查询与基础用户名相同或带 `-序号` 后缀的用户名(用于用户名去重)
func (d *DAO) ListSimilarUsernames(ctx context.Context, base string) ([]string, error) {
var existingUsernames []string
if err := d.DB(ctx).Table("w_users").
Where("username = ? OR username LIKE ? ESCAPE '\\'", base, util.EscapeLike(base)+"-%").
Pluck("username", &existingUsernames).Error; err != nil {
return nil, err
}
return existingUsernames, nil
}
// GetSystemConfigValue 读取系统配置项原始值
func (d *DAO) GetSystemConfigValue(ctx context.Context, key string) (string, error) {
var val string
if err := d.DB(ctx).Table("w_system_configs").Where("key = ?", key).Pluck("value", &val).Error; err != nil {
return "", err
}
return val, nil
}
// ListSystemConfigsByKeys 按键批量读取系统配置项
func (d *DAO) ListSystemConfigsByKeys(ctx context.Context, keys []string) ([]do.CapConfigRecord, error) {
var records []do.CapConfigRecord
db := d.DB(ctx)
if db == nil {
return nil, nil
}
if err := db.Table("w_system_configs").Where("key IN ?", keys).Find(&records).Error; err != nil {
return nil, err
}
return records, nil
}
-52
View File
@@ -1,52 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// OAuth and Auth error messages
const (
errInvalidState = "非法登录请求"
errIDTokenVerifyFailed = "ID Token 验证失败" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errIDTokenVerifyFailedFormat = "%s: %w"
errNonceMismatch = "nonce 不匹配,可能存在重放攻击"
errNoActiveAuthSource = "未配置可用认证源"
errServerAddressMissing = "服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试"
errAuthSourceRequired = "认证源不能为空"
errDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
errUsernameGenerateFailed = "无法生成可用用户名"
errUsernameFromSourceFailed = "无法从认证源获取用户名"
errAuthSourceDisabled = "认证源未启用"
errInvalidExternalAccountBindingID = "绑定记录 ID 无效"
ErrTokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errOAuthStateRateLimited = "请求授权过于频繁,请稍后重试"
errAuthSourceNameRequired = "认证源名称不能为空"
errAuthSourceNameInvalid = "认证源名称格式不正确"
errAuthSourceTypeUnsupported = "不支持的认证源类型"
errAuthSourceDiscoveryURLRequired = "Discovery URL 不能为空"
//nolint:gosec // error message, not hardcoded credentials
errAuthSourceClientCredentialsRequired = "启用认证源时必须配置 Client ID 和 Client Secret"
errAuthSourceIDRequired = "认证源 ID 不能为空"
errUserIDRequired = "用户 ID 不能为空"
errExternalAccountBindingIncomplete = "外部帐号绑定信息不完整"
errExternalAccountAlreadyBoundToAnother = "该外部帐号已被其他用户绑定"
errExternalAccountBindingIDRequired = "外部帐号绑定记录 ID 不能为空"
errInsufficientPermission = "权限不足"
errBannedAccount = "账号已被封禁"
errUnAuthorized = "未登录"
)
// Service 层与鉴权中间件内部错误文案(保持与重构前逐字一致)
const (
errUserNotInContext = "auth: user not found in context"
errEmptyToken = "auth: empty token" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errSystemUserTokenNotAllowed = "auth: system user token not allowed" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errUnauthorizedInternal = "unauthorized"
errSystemUserLoginNotAllowed = "system user is not allowed to login"
)
// OAuth 回调会话校验错误文案(保持与重构前逐字一致)
const (
errInvalidSessionContext = "invalid session context"
errSessionMismatchForOAuth = "session mismatch for oauth state"
errUserContextMismatch = "user context mismatch for oauth binding"
)
+338
View File
@@ -0,0 +1,338 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"Wavelet/plugins/domain/auth/consts"
"Wavelet/plugins/domain/auth/controller"
"Wavelet/plugins/domain/auth/dao"
"Wavelet/plugins/domain/auth/model/do"
"Wavelet/plugins/domain/auth/model/dto"
"Wavelet/plugins/domain/auth/model/entity"
"Wavelet/plugins/domain/auth/service"
"context"
"net/http"
"sync"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
)
// Exported Type Aliases for Backward Compatibility
//
//nolint:revive // backward compatibility type aliases with legacy naming
type (
AuthSource = entity.AuthSource
ExternalAccount = entity.ExternalAccount
CachedToken = do.CachedToken
CapRuntimeSettings = do.CapRuntimeSettings
AuthSourceView = dto.AuthSourceView
BasicUserInfo = dto.BasicUserInfo
OAuthAuthorizeResponse = dto.OAuthAuthorizeResponse
OAuthCallbackResult = dto.OAuthCallbackResult
CallbackRequest = dto.CallbackRequest
ChallengeResponse = dto.ChallengeResponse
RedeemResponse = dto.RedeemResponse
CaptchaManager = service.CaptchaManager
)
// Exported Constant Aliases for Backward Compatibility
const (
UserNameKey = consts.UserNameKey
UserIDKey = consts.UserIDKey
UserObjKey = consts.UserObjKey
TokenAuthKey = consts.TokenAuthKey
TokenAdminKey = consts.TokenAdminKey
SessionTokenKey = consts.SessionTokenKey
PasswordHashKey = consts.PasswordHashKey
SystemUsername = consts.SystemUsername
OAuthStateCacheKeyFormat = consts.OAuthStateCacheKeyFormat
OAuthStateCacheKeyExpiration = consts.OAuthStateCacheKeyExpiration
OAuthPurposeLogin = consts.OAuthPurposeLogin
OAuthPurposeBind = consts.OAuthPurposeBind
AuthSourceTypeOIDC = consts.AuthSourceTypeOIDC
ErrTokenAuthNotAllowed = consts.ErrTokenAuthNotAllowed
)
var (
defaultMu sync.RWMutex
defaultDAO = dao.New(nil, nil, nil)
defaultService = service.New(defaultDAO, SessionConfig{SessionCookieName: "wavelet_session", SessionAge: 86400, SessionHTTPOnly: true}, nil)
defaultCtrl = controller.New(defaultService)
)
func setDefaultRuntime(d *dao.DAO, s *service.Service, c *controller.Controller) {
defaultMu.Lock()
defer defaultMu.Unlock()
defaultDAO = d
defaultService = s
defaultCtrl = c
}
func getDefaultRuntime() (*dao.DAO, *service.Service, *controller.Controller) {
defaultMu.RLock()
defer defaultMu.RUnlock()
return defaultDAO, defaultService, defaultCtrl
}
// ParseUserID parses a string, int, or float64 user ID representation.
func ParseUserID(v any) uint64 {
return dto.ParseUserID(v)
}
// BuildBasicUserInfo converts UserDTO to BasicUserInfo.
func BuildBasicUserInfo(user *contracts.UserDTO, needChange bool) BasicUserInfo {
return dto.BuildBasicUserInfo(user, needChange)
}
// SetSessionConfig updates the active session configuration.
func SetSessionConfig(cfg SessionConfig) {
_, s, _ := getDefaultRuntime()
s.Session.SetConfig(cfg)
}
// GetSessionConfig returns the active session configuration.
func GetSessionConfig() SessionConfig {
_, s, _ := getDefaultRuntime()
return s.Session.Config()
}
// GetSessionOptions builds session cookie options based on config and maxAge.
func GetSessionOptions(maxAge int) sessions.Options {
_, s, _ := getDefaultRuntime()
return s.Session.GetSessionOptions(maxAge)
}
// StripCookieMaxAgeAndExpires removes max-age and expires from cookie header.
func StripCookieMaxAgeAndExpires(header http.Header, cookieName string) {
_, s, _ := getDefaultRuntime()
s.Session.StripCookieMaxAgeAndExpires(header, cookieName)
}
// GetUserIDFromSession extracts user ID from session.
func GetUserIDFromSession(s sessions.Session) uint64 {
return controller.GetUserIDFromSession(s)
}
// GetUserIDFromContext extracts user ID from Gin context.
func GetUserIDFromContext(c *gin.Context) uint64 {
return controller.GetUserIDFromContext(c)
}
// SetLoginSession sets the login session for the authenticated user.
func SetLoginSession(ctx context.Context, c *gin.Context, user *contracts.UserDTO, extras ...map[string]any) error {
_, s, _ := getDefaultRuntime()
session := sessions.Default(c)
isSessionCookie, err := s.Session.ApplyLoginSession(ctx, session, user, extras...)
if err != nil {
return err
}
if isSessionCookie {
s.Session.StripCookieMaxAgeAndExpires(c.Writer.Header(), s.Session.Config().SessionCookieName)
}
return nil
}
// RegisterWhitelist adds whitelist path patterns.
func RegisterWhitelist(patterns ...string) {
_, _, c := getDefaultRuntime()
c.RegisterWhitelist(patterns...)
}
// IsWhitelisted checks if the path matches the auth whitelist.
func IsWhitelisted(path string) bool {
_, _, c := getDefaultRuntime()
if wl := c.Whitelist(); wl != nil {
return wl.Match(path)
}
return false
}
// GetUserFromRequest extracts user from Request (Token or Session).
func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
d, _, _ := getDefaultRuntime()
return controller.GetUserFromRequest(c, d)
}
// LoginRequired returns authentication required middleware.
func LoginRequired() gin.HandlerFunc {
_, _, c := getDefaultRuntime()
return c.LoginRequired()
}
// AdminRequired returns admin authorization middleware.
func AdminRequired() gin.HandlerFunc {
_, _, c := getDefaultRuntime()
return c.AdminRequired()
}
// LoginAdminRequired alias for AdminRequired.
func LoginAdminRequired() gin.HandlerFunc {
return AdminRequired()
}
// DisallowTokenAuth returns middleware rejecting access token requests.
func DisallowTokenAuth() gin.HandlerFunc {
return controller.DisallowTokenAuth()
}
// GetCachedToken reads cached access token.
func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error) {
d, _, _ := getDefaultRuntime()
return d.GetCachedToken(ctx, tokenHash)
}
// SetCachedToken stores access token into cache.
func SetCachedToken(ctx context.Context, tokenHash string, token *CachedToken) {
d, _, _ := getDefaultRuntime()
d.SetCachedToken(ctx, tokenHash, token)
}
// InvalidateCachedToken invalidates access token cache.
func InvalidateCachedToken(ctx context.Context, tokenHash string) {
d, _, _ := getDefaultRuntime()
d.InvalidateCachedToken(ctx, tokenHash)
}
// GetCachedUser reads cached user.
func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
d, _, _ := getDefaultRuntime()
return d.GetCachedUser(ctx, userID)
}
// SetCachedUser stores user into cache.
func SetCachedUser(ctx context.Context, userID uint64, u *contracts.UserDTO) {
d, _, _ := getDefaultRuntime()
d.SetCachedUser(ctx, userID, u)
}
// InvalidateCachedUser invalidates user cache.
func InvalidateCachedUser(ctx context.Context, userID uint64) {
d, _, _ := getDefaultRuntime()
d.InvalidateCachedUser(ctx, userID)
}
// ResetAuthRAMCacheForTest clears RAM caches.
func ResetAuthRAMCacheForTest() {
dao.ResetRAMCacheForTest()
}
// StopAuthCacheListener compatibility stub.
func StopAuthCacheListener() {}
// SetCapSecret sets CAPTCHA secret.
func SetCapSecret(secret []byte) {
_, s, _ := getDefaultRuntime()
s.CapManager.SetSecret(secret)
}
// GetDefaultCapManager returns the singleton CAPTCHA manager.
func GetDefaultCapManager() *CaptchaManager {
_, s, _ := getDefaultRuntime()
return s.CapManager
}
// CurrentCapSettings returns current CAPTCHA runtime settings.
func CurrentCapSettings(ctx context.Context) (CapRuntimeSettings, error) {
_, s, _ := getDefaultRuntime()
return s.CapSettings.Current(ctx)
}
// CapProtectionEnabled checks if CAPTCHA is enabled.
func CapProtectionEnabled(ctx context.Context) bool {
_, s, _ := getDefaultRuntime()
return s.CapSettings.CapProtectionEnabled(ctx)
}
// InvalidateCapRuntimeSettings invalidates runtime CAPTCHA settings cache.
func InvalidateCapRuntimeSettings() {
_, s, _ := getDefaultRuntime()
s.CapSettings.Invalidate()
}
// ResetCapRuntimeSettingsForTest clears test CAPTCHA settings.
func ResetCapRuntimeSettingsForTest() {
InvalidateCapRuntimeSettings()
}
// InstallCapTestRuntimeSettings installs a test snapshot.
func InstallCapTestRuntimeSettings(settings CapRuntimeSettings) func() {
_, s, _ := getDefaultRuntime()
return s.CapSettings.InstallTestSnapshot(settings)
}
// VerifyCaptchaMiddleware returns captcha verification middleware.
func VerifyCaptchaMiddleware(mgr *service.CaptchaManager, scope string) gin.HandlerFunc {
_, s, _ := getDefaultRuntime()
return controller.VerifyCaptchaMiddleware(mgr, s.CapSettings, scope)
}
// Challenge HTTP handler.
func Challenge(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.Captcha.Challenge(c)
}
// Redeem HTTP handler.
func Redeem(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.Captcha.Redeem(c)
}
// GetLoginSources HTTP handler.
func GetLoginSources(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.OAuth.GetLoginSources(c)
}
// GetLoginURL HTTP handler.
func GetLoginURL(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.OAuth.GetLoginURL(c)
}
// Authorize HTTP handler.
func Authorize(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.OAuth.Authorize(c)
}
// Callback HTTP handler.
func Callback(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.OAuth.Callback(c)
}
// Logout HTTP handler.
func Logout(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.OAuth.Logout(c)
}
// UserInfo HTTP handler.
func UserInfo(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.UserInfo.UserInfo(c)
}
// ListExternalAccounts HTTP handler.
func ListExternalAccounts(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.OAuth.ListExternalAccounts(c)
}
// DeleteExternalAccount HTTP handler.
func DeleteExternalAccount(c *gin.Context) {
_, _, ctrl := getDefaultRuntime()
ctrl.OAuth.DeleteExternalAccount(c)
}
// InvalidateOIDCProviderCache invalidates OIDC provider cache entry.
func InvalidateOIDCProviderCache(issuer string) {
_, s, _ := getDefaultRuntime()
s.OIDCProviderCache.Invalidate(issuer)
}
-572
View File
@@ -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())
}
-195
View File
@@ -1,195 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/response"
"Wavelet/pkg/trace"
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"github.com/gin-gonic/gin"
)
// whitelist holds the no-auth route patterns. They are registered during Apply and
// matched on every request, so PathWhitelist parses them once up front.
var whitelist = extpoints.NewPathWhitelist()
// RegisterWhitelist registers route patterns that bypass mandatory authentication.
func RegisterWhitelist(patterns ...string) {
whitelist.Add(patterns...)
}
// IsWhitelisted checks if the specified path matches the auth whitelist.
func IsWhitelisted(path string) bool {
return whitelist.Match(path)
}
func hashToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
// currentUserIDFromRequestContext 是接入层向 Service 层暴露的登录态桥接。
//
// Session 读取必须依赖 *gin.Context,而 Service 层禁止 import gin,
// 因此该类型断言收敛在本(接入层)文件中。ok 为 false 表示 ctx 不是 *gin.Context。
func currentUserIDFromRequestContext(ctx context.Context) (uint64, bool) {
ginCtx, ok := ctx.(*gin.Context)
if !ok {
return 0, false
}
return GetUserIDFromContext(ginCtx), true
}
func getUserByToken(ctx context.Context, tokenStr string) (*contracts.UserDTO, *CachedToken, error) {
tokenHash := hashToken(tokenStr)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil || tokenRecord == nil {
tokenRecord, err = GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, nil, err
}
SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || user == nil || !user.IsActive {
user, err = GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, nil, err
}
SetCachedUser(ctx, tokenRecord.UserID, user)
}
return user, tokenRecord, nil
}
// GetUserFromRequest 从请求中获取当前用户(优先 Access Token,其次 Session)
func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
ctx := c.Request.Context()
var tokenStr string
tokenFromQuery := c.Query("token")
if tokenFromQuery != "" {
tokenStr = tokenFromQuery
} else {
authHeader := c.GetHeader("Authorization")
if len(authHeader) > 7 && authHeader[:7] == "Bearer " {
tokenStr = authHeader[7:]
}
}
// 优先使用 Access Token 鉴权
if tokenStr != "" {
if user, tokenRecord, err := getUserByToken(ctx, tokenStr); err == nil {
if user.Username == SystemUsername {
return nil, errors.New(errSystemUserLoginNotAllowed)
}
ginutil.SetToContext(c, contracts.AuthTokenAuthKey, true)
ginutil.SetToContext(c, contracts.AuthTokenAdminKey, tokenRecord.IsAdmin)
return user, nil
}
}
// 降级使用 Session 鉴权
userID := GetUserIDFromContext(c)
if userID <= 0 {
return nil, errors.New(errUnauthorizedInternal)
}
user, err := GetCachedUser(ctx, userID)
if err != nil || user == nil || !user.IsActive {
user, err = GetActiveUserByID(ctx, userID)
if err != nil {
return nil, err
}
SetCachedUser(ctx, userID, user)
}
ginutil.SetToContext(c, contracts.AuthTokenAuthKey, false)
ginutil.SetToContext(c, contracts.AuthTokenAdminKey, false)
if user.Username == "system" {
return nil, errors.New(errSystemUserLoginNotAllowed)
}
return user, nil
}
// LoginRequired 返回登录鉴权中间件,校验 Access Token 或 Session
func LoginRequired() gin.HandlerFunc {
return func(c *gin.Context) {
if IsWhitelisted(c.Request.URL.Path) {
c.Next()
return
}
_, span := trace.Start(c.Request.Context(), "LoginRequired")
defer span.End()
user, err := GetUserFromRequest(c)
if err != nil {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
LogForAudit(c.Request.Context(), user, c)
ginutil.SetToContext(c, contracts.AuthUserObjKey, user)
c.Next()
}
}
// AdminRequired 校验管理员权限(支持 Session 和 Token 鉴权)
func AdminRequired() gin.HandlerFunc {
return func(c *gin.Context) {
_, span := trace.Start(c.Request.Context(), "AdminRequired")
defer span.End()
user, err := GetUserFromRequest(c)
if err != nil {
response.AbortUnauthorized(c, errUnAuthorized)
return
}
isTokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey)
isTokenAdmin, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
// Logged-in but lacking admin permission is 403, not 401/404.
if isTokenAuth && !isTokenAdmin && !user.IsAdmin {
response.AbortForbidden(c, errInsufficientPermission)
return
}
if !isTokenAuth && !user.IsAdmin {
response.AbortForbidden(c, errInsufficientPermission)
return
}
LogForAudit(c.Request.Context(), user, c)
ginutil.SetToContext(c, contracts.AuthUserObjKey, user)
c.Next()
}
}
// LoginAdminRequired is an alias for AdminRequired.
func LoginAdminRequired() gin.HandlerFunc {
return AdminRequired()
}
// DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点
func DisallowTokenAuth() gin.HandlerFunc {
return func(c *gin.Context) {
if tokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
response.AbortForbidden(c, ErrTokenAuthNotAllowed)
return
}
c.Next()
}
}
@@ -0,0 +1,12 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package do provides domain data objects for the auth plugin.
package do
// CachedToken represents the minimal cached representation of an access token.
type CachedToken struct {
ID uint64 `json:"id"`
UserID uint64 `json:"user_id"`
IsAdmin bool `json:"is_admin"`
}
@@ -0,0 +1,75 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package do provides domain data objects for the auth plugin.
package do
import (
"Wavelet/plugins/domain/auth/consts"
"strconv"
"time"
)
// CapRuntimeSettings is the parsed CAPTCHA runtime configuration loaded from system_configs.
type CapRuntimeSettings struct {
LoginEnabled bool
ChallengeCount int
ChallengeSize int
ChallengeDifficulty int
ChallengeTTL time.Duration
TokenTTL time.Duration
}
// CapConfigRecord maps the columns selected from the system config table.
type CapConfigRecord struct {
Key string `gorm:"column:key"`
Value string `gorm:"column:value"`
}
// ParseCapRuntimeSettings parses system config key-value map into CapRuntimeSettings with fallback defaults.
func ParseCapRuntimeSettings(configs map[string]string) CapRuntimeSettings {
settings := CapRuntimeSettings{
ChallengeCount: consts.DefaultCapChallengeCount,
ChallengeSize: consts.DefaultCapChallengeSize,
ChallengeDifficulty: consts.DefaultCapChallengeDifficulty,
ChallengeTTL: consts.DefaultCapChallengeTTL,
TokenTTL: consts.DefaultCapTokenTTL,
}
if len(configs) == 0 {
return settings
}
if val, ok := configs[consts.ConfigKeyCapLoginEnabled]; ok {
if enabled, err := strconv.ParseBool(val); err == nil {
settings.LoginEnabled = enabled
}
}
if val, ok := configs[consts.ConfigKeyCapChallengeCount]; ok {
if count, err := strconv.Atoi(val); err == nil && count > 0 {
settings.ChallengeCount = count
}
}
if val, ok := configs[consts.ConfigKeyCapChallengeSize]; ok {
if size, err := strconv.Atoi(val); err == nil && size > 0 {
settings.ChallengeSize = size
}
}
if val, ok := configs[consts.ConfigKeyCapChallengeDifficulty]; ok {
if diff, err := strconv.Atoi(val); err == nil && diff > 0 {
settings.ChallengeDifficulty = diff
}
}
if val, ok := configs[consts.ConfigKeyCapChallengeTTL]; ok {
if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 {
settings.ChallengeTTL = time.Duration(ttlSeconds) * time.Second
}
}
if val, ok := configs[consts.ConfigKeyCapTokenTTL]; ok {
if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 {
settings.TokenTTL = time.Duration(ttlSeconds) * time.Second
}
}
return settings
}
@@ -0,0 +1,43 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package do_test
import (
"Wavelet/plugins/domain/auth/consts"
"Wavelet/plugins/domain/auth/model/do"
"testing"
"time"
"github.com/stretchr/testify/assert"
)
func TestParseCapRuntimeSettings(t *testing.T) {
t.Run("Default fallback on empty config", func(t *testing.T) {
settings := do.ParseCapRuntimeSettings(nil)
assert.False(t, settings.LoginEnabled)
assert.Equal(t, consts.DefaultCapChallengeCount, settings.ChallengeCount)
assert.Equal(t, consts.DefaultCapChallengeSize, settings.ChallengeSize)
assert.Equal(t, consts.DefaultCapChallengeDifficulty, settings.ChallengeDifficulty)
assert.Equal(t, consts.DefaultCapChallengeTTL, settings.ChallengeTTL)
assert.Equal(t, consts.DefaultCapTokenTTL, settings.TokenTTL)
})
t.Run("Parsed custom configs", func(t *testing.T) {
configs := map[string]string{
consts.ConfigKeyCapLoginEnabled: "true",
consts.ConfigKeyCapChallengeCount: "3",
consts.ConfigKeyCapChallengeSize: "64",
consts.ConfigKeyCapChallengeDifficulty: "5",
consts.ConfigKeyCapChallengeTTL: "300",
consts.ConfigKeyCapTokenTTL: "600",
}
settings := do.ParseCapRuntimeSettings(configs)
assert.True(t, settings.LoginEnabled)
assert.Equal(t, 3, settings.ChallengeCount)
assert.Equal(t, 64, settings.ChallengeSize)
assert.Equal(t, 5, settings.ChallengeDifficulty)
assert.Equal(t, 300*time.Second, settings.ChallengeTTL)
assert.Equal(t, 600*time.Second, settings.TokenTTL)
})
}
@@ -0,0 +1,33 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package do provides domain data objects for the auth plugin.
package do
import "encoding/json"
// OAuthStatePayload represents the cached state verification payload for OAuth flow.
type OAuthStatePayload struct {
SourceName string `json:"source_name"`
Purpose string `json:"purpose"`
UserID uint64 `json:"user_id,omitempty"`
SessionHash string `json:"session_hash"`
}
// Encode converts OAuthStatePayload to a JSON string.
func (p OAuthStatePayload) Encode() (string, error) {
data, err := json.Marshal(p)
if err != nil {
return "", err
}
return string(data), nil
}
// DecodeOAuthStatePayload parses a JSON string into OAuthStatePayload.
func DecodeOAuthStatePayload(value string) (OAuthStatePayload, error) {
var payload OAuthStatePayload
if err := json.Unmarshal([]byte(value), &payload); err != nil {
return OAuthStatePayload{}, err
}
return payload, nil
}
@@ -0,0 +1,32 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package do_test
import (
"Wavelet/plugins/domain/auth/model/do"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestOAuthStatePayload(t *testing.T) {
payload := do.OAuthStatePayload{
SourceName: "github",
Purpose: "login",
UserID: 12345,
SessionHash: "hash-abc-123",
}
encoded, err := payload.Encode()
require.NoError(t, err)
assert.NotEmpty(t, encoded)
decoded, err := do.DecodeOAuthStatePayload(encoded)
require.NoError(t, err)
assert.Equal(t, payload, decoded)
_, err = do.DecodeOAuthStatePayload("invalid-json")
assert.Error(t, err)
}
@@ -0,0 +1,45 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package dto provides data transfer objects and views for the auth plugin.
package dto
// AuthSourceView 登录源展示信息
type AuthSourceView struct {
ID uint64 `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
DisplayName string `json:"display_name"`
IsActive bool `json:"is_active"`
IconURL string `json:"icon_url"`
ClientSecretConfigured bool `json:"client_secret_configured"`
}
// OAuthAuthorizeResponse 授权 URL 响应
type OAuthAuthorizeResponse struct {
AuthorizeURL string `json:"authorize_url"`
}
// OAuthCallbackResult 回调处理结果
type OAuthCallbackResult struct {
Status string `json:"status"`
User *BasicUserInfo `json:"user,omitempty"`
}
// CallbackRequest OAuth 回调请求参数
type CallbackRequest struct {
State string `json:"state" binding:"required"`
Code string `json:"code" binding:"required"`
}
// ExternalAccountView 外部帐号绑定视图(脱敏展示用)
type ExternalAccountView struct {
ID uint64 `json:"id"`
AuthSourceID uint64 `json:"auth_source_id"`
AuthSourceName string `json:"auth_source_name"`
AuthSourceType string `json:"auth_source_type"`
AuthSourceLabel string `json:"auth_source_label"`
ExternalUsername string `json:"external_username"`
Email string `json:"email"`
CreatedAt string `json:"created_at"`
}
@@ -1,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"
}
-241
View File
@@ -1,241 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"encoding/json"
"errors"
"regexp"
"strconv"
"strings"
"time"
)
var authSourceNamePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_-]{0,79}$`)
// AuthSource 认证源实体
//
//nolint:revive // auth.AuthSource is standard domain entity name
type AuthSource struct {
ID uint64 `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"uniqueIndex;size:80;not null"`
Type string `json:"type" gorm:"size:20;not null"`
DisplayName string `json:"display_name" gorm:"size:100"`
IsActive bool `json:"is_active" gorm:"index;not null;default:false"`
ClientID string `json:"client_id" gorm:"size:255"`
ClientSecret string `json:"-" gorm:"size:1024"`
OpenIDDiscoveryURL string `json:"openid_discovery_url" gorm:"column:openid_discovery_url;size:1024"`
Scopes string `json:"scopes" gorm:"size:255"`
IconURL string `json:"icon_url" gorm:"size:1024"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
ClientSecretConfigured bool `json:"client_secret_configured" gorm:"-"`
}
// TableName 表名
func (AuthSource) TableName() string {
return "w_auth_sources"
}
// Normalize 对认证源字段进行标准化处理
func (source *AuthSource) Normalize() {
source.Type = strings.ToLower(strings.TrimSpace(source.Type))
source.Name = strings.TrimSpace(source.Name)
source.DisplayName = strings.TrimSpace(source.DisplayName)
source.ClientID = strings.TrimSpace(source.ClientID)
source.ClientSecret = strings.TrimSpace(source.ClientSecret)
source.OpenIDDiscoveryURL = strings.TrimSpace(source.OpenIDDiscoveryURL)
source.Scopes = strings.TrimSpace(source.Scopes)
source.IconURL = strings.TrimSpace(source.IconURL)
if source.DisplayName == "" {
source.DisplayName = source.Name
}
if source.Type == AuthSourceTypeOIDC && source.Scopes == "" {
source.Scopes = "openid profile email"
}
}
// Validate 校验认证源字段合法性
func (source *AuthSource) Validate() error {
source.Normalize()
if source.Name == "" {
return errors.New(errAuthSourceNameRequired)
}
if !authSourceNamePattern.MatchString(source.Name) {
return errors.New(errAuthSourceNameInvalid)
}
if source.Type != AuthSourceTypeOIDC {
return errors.New(errAuthSourceTypeUnsupported)
}
if source.OpenIDDiscoveryURL == "" {
//nolint:staticcheck // descriptive error constant
return errors.New(errAuthSourceDiscoveryURLRequired)
}
if source.IsActive && (source.ClientID == "" || source.ClientSecret == "") {
return errors.New(errAuthSourceClientCredentialsRequired)
}
return nil
}
// Sanitize 脱敏处理,将 ClientSecret 清空并设置 ClientSecretConfigured 标志
func (source *AuthSource) Sanitize() {
source.ClientSecretConfigured = source.ClientSecret != ""
source.ClientSecret = ""
}
// ExternalAccount 外部账号绑定实体
type ExternalAccount struct {
ID uint64 `json:"id" gorm:"primaryKey"`
AuthSourceID uint64 `json:"auth_source_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:1;index"`
UserID uint64 `json:"user_id" gorm:"index;not null"`
ExternalID string `json:"external_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:2;size:255;not null"`
ExternalUsername string `json:"external_username" gorm:"size:255"`
Email string `json:"email" gorm:"size:255"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// TableName 表名
func (ExternalAccount) TableName() string {
return "w_external_accounts"
}
// ExternalAccountView 外部帐号绑定视图(脱敏展示用)
type ExternalAccountView struct {
ID uint64 `json:"id"`
AuthSourceID uint64 `json:"auth_source_id"`
AuthSourceName string `json:"auth_source_name"`
AuthSourceType string `json:"auth_source_type"`
AuthSourceLabel string `json:"auth_source_label"`
ExternalUsername string `json:"external_username"`
Email string `json:"email"`
CreatedAt time.Time `json:"created_at"`
}
// AuthSourceView 登录源展示信息
//
//nolint:revive // auth.AuthSourceView is standard domain presentation struct
type AuthSourceView struct {
ID uint64 `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
DisplayName string `json:"display_name"`
IsActive bool `json:"is_active"`
IconURL string `json:"icon_url"`
ClientSecretConfigured bool `json:"client_secret_configured"`
}
// OAuthAuthorizeResponse 授权 URL 响应
type OAuthAuthorizeResponse struct {
AuthorizeURL string `json:"authorize_url"`
}
// OAuthCallbackResult 回调处理结果
type OAuthCallbackResult struct {
Status string `json:"status"`
User *BasicUserInfo `json:"user,omitempty"`
}
// CallbackRequest OAuth 回调请求参数
type CallbackRequest struct {
State string `json:"state" binding:"required"`
Code string `json:"code" binding:"required"`
}
// BasicUserInfo 用户基本信息结构体
type BasicUserInfo struct {
ID uint64 `json:"id,string"`
Username string `json:"username"`
Nickname string `json:"nickname"`
Email string `json:"email"`
AvatarURL string `json:"avatar_url"`
IsAdmin bool `json:"is_admin"`
NeedChangePassword bool `json:"need_change_password"`
Bio string `json:"bio"`
Phone string `json:"phone"`
Gender string `json:"gender"`
Website string `json:"website"`
Location string `json:"location"`
}
// BuildBasicUserInfo 将 UserDTO 转换为 BasicUserInfo
func BuildBasicUserInfo(user *contracts.UserDTO, needChange bool) BasicUserInfo {
if user == nil {
return BasicUserInfo{}
}
return BasicUserInfo{
ID: user.ID,
Username: user.Username,
Nickname: user.Nickname,
Email: user.Email,
AvatarURL: user.AvatarURL,
IsAdmin: user.IsAdmin,
NeedChangePassword: needChange || user.NeedChangePassword,
Bio: user.Bio,
Phone: user.Phone,
Gender: user.Gender,
Website: user.Website,
Location: user.Location,
}
}
type oauthStatePayload struct {
SourceName string `json:"source_name"`
Purpose string `json:"purpose"`
UserID uint64 `json:"user_id,omitempty"`
SessionHash string `json:"session_hash"`
}
func encodeOAuthStatePayload(payload oauthStatePayload) (string, error) {
data, err := json.Marshal(payload)
if err != nil {
return "", err
}
return string(data), nil
}
func decodeOAuthStatePayload(value string) (oauthStatePayload, error) {
var payload oauthStatePayload
if err := json.Unmarshal([]byte(value), &payload); err != nil {
return oauthStatePayload{}, err
}
return payload, nil
}
type loginRequiredAuditLog struct {
UserID uint64 `json:"user_id"`
Username string `json:"username"`
ClientIP string `json:"client_ip"`
Method string `json:"method"`
Path string `json:"path"`
RequestURI string `json:"request_uri"`
UserAgent string `json:"user_agent"`
Referer string `json:"referer"`
}
// ParseUserID parses a string or float64 user ID representation.
func ParseUserID(v any) uint64 {
switch val := v.(type) {
case uint64:
return val
case int64:
if val > 0 {
return uint64(val)
}
case int:
if val > 0 {
return uint64(val)
}
case float64:
if val > 0 {
return uint64(val)
}
case string:
if id, err := strconv.ParseUint(val, 10, 64); err == nil {
return id
}
}
return 0
}
@@ -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"])
}
}
+79 -37
View File
@@ -8,6 +8,9 @@ import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"Wavelet/plugins/domain/auth/controller"
"Wavelet/plugins/domain/auth/dao"
"Wavelet/plugins/domain/auth/service"
"context"
"embed"
"reflect"
@@ -83,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")
}
-211
View File
@@ -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
}
-262
View File
@@ -1,262 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"context"
"errors"
"sync"
)
type authServiceImpl struct{}
func newAuthService() contracts.AuthService {
return &authServiceImpl{}
}
func (s *authServiceImpl) RequireAuthMiddleware() any {
return LoginRequired()
}
func (s *authServiceImpl) RequireAdminMiddleware() any {
return AdminRequired()
}
// GetCurrentUser 从 context 中读取登录用户。
//
// 中间件通过 gin 的 c.Set(contracts.AuthUserObjKey, user) 写入登录态;
// *gin.Context 自身实现了 context.Context,且其 Value(key) 对 string 类型 key
// 等价于 c.Get(key)(未命中时再回落到 Request.Context().Value),
// 因此这里无需感知 gin 即可读取同一份登录态。
func (s *authServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
if v := ctx.Value(contracts.AuthUserObjKey); v != nil {
if u, ok := v.(*contracts.UserDTO); ok && u != nil {
return u, nil
}
}
return nil, errors.New(errUserNotInContext)
}
func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contracts.UserDTO, error) {
if token == "" {
return nil, errors.New(errEmptyToken)
}
tokenHash := hashToken(token)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil {
tokenRecord, err = GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, err
}
SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || user == nil || !user.IsActive {
user, err = GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, err
}
SetCachedUser(ctx, tokenRecord.UserID, user)
}
if user.Username == SystemUsername {
return nil, errors.New(errSystemUserTokenNotAllowed)
}
return user, nil
}
func (s *authServiceImpl) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) {
return "", nil
}
func (s *authServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64) error {
InvalidateCachedUser(ctx, userID)
return nil
}
// GetCurrentUserID 从请求登录态中读取用户 ID。
//
// Session 读取依赖 gin,属于接入层职责,因此这里通过接入层桥接函数
// currentUserIDFromRequestContext(见 middleware.go)取值,Service 层本身不感知 gin。
func (s *authServiceImpl) GetCurrentUserID(ctx context.Context) (uint64, error) {
userID, ok := currentUserIDFromRequestContext(ctx)
if !ok {
return 0, errors.New(errUserNotInContext)
}
return userID, nil
}
func (s *authServiceImpl) RevokeToken(ctx context.Context, tokenHash string) error {
InvalidateCachedToken(ctx, tokenHash)
return nil
}
func (s *authServiceImpl) DisallowTokenAuthMiddleware() any {
return DisallowTokenAuth()
}
func (s *authServiceImpl) InvalidateCachedUser(ctx context.Context, userID uint64) {
InvalidateCachedUser(ctx, userID)
}
func (s *authServiceImpl) InvalidateCachedToken(ctx context.Context, tokenHash string) {
InvalidateCachedToken(ctx, tokenHash)
}
func (s *authServiceImpl) ListAuthSources(ctx context.Context) ([]contracts.AuthSourceViewDTO, error) {
sources, err := ListAllAuthSources(ctx)
if err != nil {
return nil, err
}
views := make([]contracts.AuthSourceViewDTO, len(sources))
for i := range sources {
views[i] = contracts.AuthSourceViewDTO{
ID: sources[i].ID,
Name: sources[i].Name,
Type: sources[i].Type,
DisplayName: sources[i].DisplayName,
IsActive: sources[i].IsActive,
IconURL: sources[i].IconURL,
ClientSecretConfigured: sources[i].ClientSecret != "",
}
}
return views, nil
}
func (s *authServiceImpl) CreateAuthSource(ctx context.Context, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
model := AuthSource{
ID: source.ID,
Name: source.Name,
Type: source.Type,
DisplayName: source.DisplayName,
ClientID: source.ClientID,
ClientSecret: source.ClientSecret,
OpenIDDiscoveryURL: source.OpenIDDiscoveryURL,
Scopes: source.Scopes,
IconURL: source.IconURL,
IsActive: source.IsActive,
}
if err := model.Validate(); err != nil {
return nil, err
}
if err := CreateAuthSourceRecord(ctx, &model); err != nil {
return nil, err
}
model.Sanitize()
return toAuthSourceDTO(&model), nil
}
func (s *authServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
existing, err := GetAuthSourceByID(ctx, id)
if err != nil {
return nil, err
}
existing.DisplayName = source.DisplayName
existing.ClientID = source.ClientID
if source.ClientSecret != "" {
existing.ClientSecret = source.ClientSecret
}
existing.OpenIDDiscoveryURL = source.OpenIDDiscoveryURL
existing.Scopes = source.Scopes
existing.IconURL = source.IconURL
if err := existing.Validate(); err != nil {
return nil, err
}
if err := SaveAuthSourceRecord(ctx, existing); err != nil {
return nil, err
}
existing.Sanitize()
return toAuthSourceDTO(existing), nil
}
func (s *authServiceImpl) DeleteAuthSource(ctx context.Context, id uint64) error {
existing, err := GetAuthSourceByID(ctx, id)
if err != nil {
return err
}
return DeleteAuthSourceRecord(ctx, existing)
}
func (s *authServiceImpl) ToggleAuthSource(ctx context.Context, id uint64) (*contracts.AuthSourceDTO, error) {
existing, err := GetAuthSourceByID(ctx, id)
if err != nil {
return nil, err
}
existing.IsActive = !existing.IsActive
if err := SaveAuthSourceRecord(ctx, existing); err != nil {
return nil, err
}
existing.Sanitize()
return toAuthSourceDTO(existing), nil
}
func toAuthSourceDTO(s *AuthSource) *contracts.AuthSourceDTO {
if s == nil {
return nil
}
return &contracts.AuthSourceDTO{
ID: s.ID,
Name: s.Name,
Type: s.Type,
DisplayName: s.DisplayName,
ClientID: s.ClientID,
ClientSecret: s.ClientSecret,
OpenIDDiscoveryURL: s.OpenIDDiscoveryURL,
Scopes: s.Scopes,
IconURL: s.IconURL,
IsActive: s.IsActive,
CreatedAt: s.CreatedAt,
UpdatedAt: s.UpdatedAt,
}
}
type authRegistryImpl struct {
mu sync.RWMutex
providers map[string]contracts.OAuthProvider
}
func newAuthRegistry() contracts.AuthRegistry {
return &authRegistryImpl{
providers: make(map[string]contracts.OAuthProvider),
}
}
func (r *authRegistryImpl) RegisterOAuthProvider(name string, provider contracts.OAuthProvider) {
r.mu.Lock()
defer r.mu.Unlock()
r.providers[name] = provider
}
func (r *authRegistryImpl) GetOAuthProvider(name string) (contracts.OAuthProvider, bool) {
r.mu.RLock()
defer r.mu.RUnlock()
p, ok := r.providers[name]
return p, ok
}
func (r *authRegistryImpl) ListOAuthProviders() []string {
r.mu.RLock()
defer r.mu.RUnlock()
res := make([]string, 0, len(r.providers))
for name := range r.providers {
res = append(res, name)
}
return res
}
@@ -0,0 +1,49 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package service implements domain business services and orchestration for the auth plugin.
package service
import (
"Wavelet/core/contracts"
"sync"
)
// AuthRegistryImpl implements contracts.AuthRegistry.
type AuthRegistryImpl struct {
mu sync.RWMutex
providers map[string]contracts.OAuthProvider
}
// NewAuthRegistry creates a new AuthRegistryImpl.
func NewAuthRegistry() *AuthRegistryImpl {
return &AuthRegistryImpl{
providers: make(map[string]contracts.OAuthProvider),
}
}
// RegisterOAuthProvider registers an OAuthProvider by name.
func (r *AuthRegistryImpl) RegisterOAuthProvider(name string, provider contracts.OAuthProvider) {
r.mu.Lock()
defer r.mu.Unlock()
r.providers[name] = provider
}
// GetOAuthProvider retrieves an OAuthProvider by name.
func (r *AuthRegistryImpl) GetOAuthProvider(name string) (contracts.OAuthProvider, bool) {
r.mu.RLock()
defer r.mu.RUnlock()
p, ok := r.providers[name]
return p, ok
}
// ListOAuthProviders lists all registered provider names.
func (r *AuthRegistryImpl) ListOAuthProviders() []string {
r.mu.RLock()
defer r.mu.RUnlock()
res := make([]string, 0, len(r.providers))
for name := range r.providers {
res = append(res, name)
}
return res
}
@@ -0,0 +1,278 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package service implements domain business services and orchestration for the auth plugin.
package service
import (
"Wavelet/core/contracts"
"Wavelet/plugins/domain/auth/consts"
"Wavelet/plugins/domain/auth/dao"
"Wavelet/plugins/domain/auth/model/entity"
"context"
"crypto/sha256"
"encoding/hex"
"errors"
)
// HashToken computes SHA-256 hex digest of access token.
func HashToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
// UserIDExtractor extracts user ID from a request context.
type UserIDExtractor func(ctx context.Context) (uint64, bool)
// AuthServiceImpl implements contracts.AuthService.
type AuthServiceImpl struct {
dao *dao.DAO
requireAuthMiddleware any
requireAdminMiddleware any
disallowTokenMiddleware any
userIDExtractor UserIDExtractor
}
// NewAuthService creates a new AuthServiceImpl.
func NewAuthService(
d *dao.DAO,
requireAuth any,
requireAdmin any,
disallowToken any,
extractor UserIDExtractor,
) *AuthServiceImpl {
return &AuthServiceImpl{
dao: d,
requireAuthMiddleware: requireAuth,
requireAdminMiddleware: requireAdmin,
disallowTokenMiddleware: disallowToken,
userIDExtractor: extractor,
}
}
// SetMiddlewareHandlers wires middleware handlers into AuthService after controller initialization.
func (s *AuthServiceImpl) SetMiddlewareHandlers(requireAuth, requireAdmin, disallowToken any, extractor UserIDExtractor) {
s.requireAuthMiddleware = requireAuth
s.requireAdminMiddleware = requireAdmin
s.disallowTokenMiddleware = disallowToken
s.userIDExtractor = extractor
}
// RequireAuthMiddleware returns the authentication check middleware.
func (s *AuthServiceImpl) RequireAuthMiddleware() any {
return s.requireAuthMiddleware
}
// RequireAdminMiddleware returns the admin authorization middleware.
func (s *AuthServiceImpl) RequireAdminMiddleware() any {
return s.requireAdminMiddleware
}
// DisallowTokenAuthMiddleware returns the token rejection middleware.
func (s *AuthServiceImpl) DisallowTokenAuthMiddleware() any {
return s.disallowTokenMiddleware
}
// GetCurrentUser 从 context 中读取登录用户。
func (s *AuthServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
if v := ctx.Value(contracts.AuthUserObjKey); v != nil {
if u, ok := v.(*contracts.UserDTO); ok && u != nil {
return u, nil
}
}
return nil, errors.New(consts.ErrUserNotInContext)
}
// GetCurrentUserID 从请求登录态中读取用户 ID。
func (s *AuthServiceImpl) GetCurrentUserID(ctx context.Context) (uint64, error) {
if s.userIDExtractor != nil {
if userID, ok := s.userIDExtractor(ctx); ok {
return userID, nil
}
}
return 0, errors.New(consts.ErrUserNotInContext)
}
// VerifyToken 验证访问令牌并返回对应的用户。
func (s *AuthServiceImpl) VerifyToken(ctx context.Context, token string) (*contracts.UserDTO, error) {
if token == "" {
return nil, errors.New(consts.ErrEmptyToken)
}
tokenHash := HashToken(token)
tokenRecord, err := s.dao.GetCachedToken(ctx, tokenHash)
if err != nil {
tokenRecord, err = s.dao.GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, err
}
s.dao.SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := s.dao.GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || user == nil || !user.IsActive {
user, err = s.dao.GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, err
}
s.dao.SetCachedUser(ctx, tokenRecord.UserID, user)
}
if user.Username == consts.SystemUsername {
return nil, errors.New(consts.ErrSystemUserTokenNotAllowed)
}
return user, nil
}
// CreateSession establishes an authenticated session.
func (s *AuthServiceImpl) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) {
return "", nil
}
// RevokeUserSessions revokes active sessions and cached tokens for a user.
func (s *AuthServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64) error {
s.dao.InvalidateCachedUser(ctx, userID)
return nil
}
// RevokeToken invalidates a cached token by its hash.
func (s *AuthServiceImpl) RevokeToken(ctx context.Context, tokenHash string) error {
s.dao.InvalidateCachedToken(ctx, tokenHash)
return nil
}
// InvalidateCachedUser invalidates cached user profile.
func (s *AuthServiceImpl) InvalidateCachedUser(ctx context.Context, userID uint64) {
s.dao.InvalidateCachedUser(ctx, userID)
}
// InvalidateCachedToken invalidates cached access token.
func (s *AuthServiceImpl) InvalidateCachedToken(ctx context.Context, tokenHash string) {
s.dao.InvalidateCachedToken(ctx, tokenHash)
}
// ListAuthSources lists all configured authentication sources.
func (s *AuthServiceImpl) ListAuthSources(ctx context.Context) ([]contracts.AuthSourceViewDTO, error) {
sources, err := s.dao.ListAllAuthSources(ctx)
if err != nil {
return nil, err
}
views := make([]contracts.AuthSourceViewDTO, len(sources))
for i := range sources {
views[i] = contracts.AuthSourceViewDTO{
ID: sources[i].ID,
Name: sources[i].Name,
Type: sources[i].Type,
DisplayName: sources[i].DisplayName,
IsActive: sources[i].IsActive,
IconURL: sources[i].IconURL,
ClientSecretConfigured: sources[i].ClientSecret != "",
}
}
return views, nil
}
// CreateAuthSource creates a new authentication source.
func (s *AuthServiceImpl) CreateAuthSource(ctx context.Context, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
model := entity.AuthSource{
ID: source.ID,
Name: source.Name,
Type: source.Type,
DisplayName: source.DisplayName,
ClientID: source.ClientID,
ClientSecret: source.ClientSecret,
OpenIDDiscoveryURL: source.OpenIDDiscoveryURL,
Scopes: source.Scopes,
IconURL: source.IconURL,
IsActive: source.IsActive,
}
if err := model.Validate(); err != nil {
return nil, err
}
if err := s.dao.CreateAuthSource(ctx, &model); err != nil {
return nil, err
}
model.Sanitize()
return toAuthSourceDTO(&model), nil
}
// UpdateAuthSource updates an existing authentication source.
func (s *AuthServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
existing, err := s.dao.GetAuthSourceByID(ctx, id)
if err != nil {
return nil, err
}
existing.DisplayName = source.DisplayName
existing.ClientID = source.ClientID
if source.ClientSecret != "" {
existing.ClientSecret = source.ClientSecret
}
existing.OpenIDDiscoveryURL = source.OpenIDDiscoveryURL
existing.Scopes = source.Scopes
existing.IconURL = source.IconURL
if err := existing.Validate(); err != nil {
return nil, err
}
if err := s.dao.SaveAuthSource(ctx, existing); err != nil {
return nil, err
}
existing.Sanitize()
return toAuthSourceDTO(existing), nil
}
// DeleteAuthSource deletes an authentication source.
func (s *AuthServiceImpl) DeleteAuthSource(ctx context.Context, id uint64) error {
existing, err := s.dao.GetAuthSourceByID(ctx, id)
if err != nil {
return err
}
return s.dao.DeleteAuthSource(ctx, existing)
}
// ToggleAuthSource toggles active status of an authentication source.
func (s *AuthServiceImpl) ToggleAuthSource(ctx context.Context, id uint64) (*contracts.AuthSourceDTO, error) {
existing, err := s.dao.GetAuthSourceByID(ctx, id)
if err != nil {
return nil, err
}
existing.IsActive = !existing.IsActive
if err := s.dao.SaveAuthSource(ctx, existing); err != nil {
return nil, err
}
existing.Sanitize()
return toAuthSourceDTO(existing), nil
}
func toAuthSourceDTO(s *entity.AuthSource) *contracts.AuthSourceDTO {
if s == nil {
return nil
}
return &contracts.AuthSourceDTO{
ID: s.ID,
Name: s.Name,
Type: s.Type,
DisplayName: s.DisplayName,
ClientID: s.ClientID,
ClientSecret: s.ClientSecret,
OpenIDDiscoveryURL: s.OpenIDDiscoveryURL,
Scopes: s.Scopes,
IconURL: s.IconURL,
IsActive: s.IsActive,
CreatedAt: s.CreatedAt,
UpdatedAt: s.UpdatedAt,
}
}
@@ -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)
}
@@ -1,7 +1,8 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// Package service implements domain business services and orchestration for the auth plugin.
package service
import (
"context"
@@ -13,16 +14,18 @@ import (
"golang.org/x/sync/singleflight"
)
// oidcProviderCache 进程级 OIDC provider 缓存。
type oidcProviderCache struct {
// OIDCProviderCache 进程级 OIDC provider 缓存。
type OIDCProviderCache struct {
mu sync.RWMutex
entries map[string]*oidc.Provider // key: normalized issuer URL
sfGroup singleflight.Group
}
// globalOIDCProviderCache 是包级单例缓存,与进程同生命周期。
var globalOIDCProviderCache = &oidcProviderCache{
entries: make(map[string]*oidc.Provider),
// NewOIDCProviderCache creates a new OIDCProviderCache.
func NewOIDCProviderCache() *OIDCProviderCache {
return &OIDCProviderCache{
entries: make(map[string]*oidc.Provider),
}
}
// discoveryContext 从请求 ctx 提取 HTTP 客户端,并绑定到不可取消的 Background ctx。
@@ -34,8 +37,8 @@ func discoveryContext(ctx context.Context) context.Context {
return bg
}
// get 返回缓存的 provider;若无则通过 oidc.NewProvider 获取并写入缓存。
func (c *oidcProviderCache) get(ctx context.Context, issuer string) (*oidc.Provider, error) {
// Get 返回缓存的 provider;若无则通过 oidc.NewProvider 获取并写入缓存。
func (c *OIDCProviderCache) Get(ctx context.Context, issuer string) (*oidc.Provider, error) {
c.mu.RLock()
if p, ok := c.entries[issuer]; ok {
c.mu.RUnlock()
@@ -68,14 +71,9 @@ func (c *oidcProviderCache) get(ctx context.Context, issuer string) (*oidc.Provi
return v.(*oidc.Provider), nil //nolint:forcetypeassert
}
// invalidate 从缓存中移除指定 issuer 对应的 provider。
func (c *oidcProviderCache) invalidate(issuer string) {
// Invalidate 从缓存中移除指定 issuer 对应的 provider。
func (c *OIDCProviderCache) Invalidate(issuer string) {
c.mu.Lock()
delete(c.entries, issuer)
c.mu.Unlock()
}
// InvalidateOIDCProviderCache 从进程级缓存中清除指定 issuer 的 provider 条目。
func InvalidateOIDCProviderCache(issuer string) {
globalOIDCProviderCache.invalidate(issuer)
}
@@ -0,0 +1,51 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package service implements domain business services and orchestration for the auth plugin.
package service
import (
"Wavelet/plugins/domain/auth/dao"
"Wavelet/plugins/domain/auth/pow"
"time"
)
// Service aggregates all domain services for the auth plugin.
type Service struct {
DAO *dao.DAO
Session *SessionService
OAuth *OAuthService
OIDCProviderCache *OIDCProviderCache
CapSettings *CapSettingsManager
CapManager *CaptchaManager
AuthSvc *AuthServiceImpl
AuthRegistry *AuthRegistryImpl
}
// New creates a new Service container with all domain services wired up.
func New(d *dao.DAO, sessionCfg SessionConfig, capSecret []byte) *Service {
sessionSvc := NewSessionService(sessionCfg, d)
oidcCache := NewOIDCProviderCache()
oauthSvc := NewOAuthService(d, oidcCache, sessionSvc)
capSettings := NewCapSettingsManager(d)
var capStore pow.Store
if len(capSecret) > 0 {
capStore = pow.NewMemoryStore(1 * time.Minute)
}
capMgr := NewCaptchaManager(capSecret, capStore, capSettings)
authSvc := NewAuthService(d, nil, nil, nil, nil)
authRegistry := NewAuthRegistry()
return &Service{
DAO: d,
Session: sessionSvc,
OAuth: oauthSvc,
OIDCProviderCache: oidcCache,
CapSettings: capSettings,
CapManager: capMgr,
AuthSvc: authSvc,
AuthRegistry: authRegistry,
}
}
@@ -0,0 +1,177 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package service implements domain business services and orchestration for the auth plugin.
package service
import (
"Wavelet/core/contracts"
"Wavelet/plugins/domain/auth/consts"
"Wavelet/plugins/domain/auth/dao"
"context"
"crypto/sha256"
"encoding/hex"
"net/http"
"strconv"
"strings"
"sync"
"github.com/gin-contrib/sessions"
"github.com/google/uuid"
gsessions "github.com/gorilla/sessions"
)
// SessionConfig defines session settings.
type SessionConfig struct {
SessionCookieName string `config:"session_cookie_name" env:"APP_SESSION_COOKIE_NAME" default:"wavelet_session"`
SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"`
SessionDomain string `config:"session_domain" env:"APP_SESSION_DOMAIN"`
SessionAge int `config:"session_age" env:"APP_SESSION_AGE" default:"86400"`
SessionHTTPOnly bool `config:"session_http_only" env:"APP_SESSION_HTTP_ONLY" default:"true"`
SessionSecure bool `config:"session_secure" env:"APP_SESSION_SECURE"`
}
// SessionService manages HTTP session operations, cookies, and tokens.
type SessionService struct {
mu sync.RWMutex
config SessionConfig
dao *dao.DAO
}
// NewSessionService creates a new SessionService.
func NewSessionService(cfg SessionConfig, d *dao.DAO) *SessionService {
return &SessionService{
config: cfg,
dao: d,
}
}
// SetConfig updates the active session configuration.
func (s *SessionService) SetConfig(cfg SessionConfig) {
s.mu.Lock()
defer s.mu.Unlock()
s.config = cfg
}
// Config returns the current session configuration.
func (s *SessionService) Config() SessionConfig {
s.mu.RLock()
defer s.mu.RUnlock()
return s.config
}
// GetSessionOptions 根据配置构建 Session 选项
func (s *SessionService) GetSessionOptions(maxAge int) sessions.Options {
cfg := s.Config()
return sessions.Options{
Path: "/",
Domain: cfg.SessionDomain,
MaxAge: maxAge,
HttpOnly: cfg.SessionHTTPOnly,
Secure: cfg.SessionSecure,
SameSite: http.SameSiteLaxMode,
}
}
// StripCookieMaxAgeAndExpires 从 Set-Cookie 响应头中移除 Max-Age 和 Expires,从而使其成为浏览器会话 Cookie
func (s *SessionService) StripCookieMaxAgeAndExpires(header http.Header, cookieName string) {
headers := header["Set-Cookie"]
if len(headers) == 0 {
return
}
newHeaders := make([]string, 0, len(headers))
for _, h := range headers {
if strings.HasPrefix(h, cookieName+"=") {
parts := strings.Split(h, ";")
newParts := make([]string, 0, len(parts))
for _, p := range parts {
trimmed := strings.TrimSpace(p)
lower := strings.ToLower(trimmed)
if strings.HasPrefix(lower, "max-age=") || strings.HasPrefix(lower, "expires=") {
continue
}
newParts = append(newParts, p)
}
newHeaders = append(newHeaders, strings.Join(newParts, ";"))
} else {
newHeaders = append(newHeaders, h)
}
}
header["Set-Cookie"] = newHeaders
}
// EnsureSessionToken returns or generates the session unique token.
func (s *SessionService) EnsureSessionToken(session sessions.Session) (string, bool) {
token, ok := session.Get(consts.SessionTokenKey).(string)
if !ok || token == "" {
token = uuid.NewString()
session.Set(consts.SessionTokenKey, token)
return token, true
}
return token, false
}
// HashSessionToken hashes the session token using SHA-256.
func (s *SessionService) HashSessionToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
// RotateSessionID forces session ID rotation to prevent session fixation attacks.
func (s *SessionService) RotateSessionID(session sessions.Session) {
if inner, ok := session.(interface{ Session() *gsessions.Session }); ok {
if sess := inner.Session(); sess != nil {
sess.ID = ""
}
}
}
// CalculateSessionMaxAge dynamically calculates max age and whether it's a browser-session cookie.
func (s *SessionService) CalculateSessionMaxAge(ctx context.Context) (int, bool) {
cfg := s.Config()
maxAge := cfg.SessionAge
isSessionCookie := false
if s.dao != nil {
val, err := s.dao.GetSystemConfigValue(ctx, "login_session_ttl_hours")
if err == nil && val != "" {
if ttlHours, err := strconv.Atoi(val); err == nil {
switch {
case ttlHours == -1:
// 永不过期,设置为 10 年
maxAge = 10 * 365 * 24 * 3600
case ttlHours > 0:
maxAge = ttlHours * 3600
case ttlHours == 0:
isSessionCookie = true
}
}
}
}
return maxAge, isSessionCookie
}
// ApplyLoginSession writes the authenticated user into a freshly rotated session.
func (s *SessionService) ApplyLoginSession(ctx context.Context, session sessions.Session, user *contracts.UserDTO, extras ...map[string]any) (bool, error) {
session.Clear()
s.RotateSessionID(session)
session.Set(consts.UserIDKey, strconv.FormatUint(user.ID, 10))
session.Set(consts.UserNameKey, user.Username)
if len(extras) > 0 {
for key, value := range extras[0] {
session.Set(key, value)
}
}
maxAge, isSessionCookie := s.CalculateSessionMaxAge(ctx)
session.Options(s.GetSessionOptions(maxAge))
if err := session.Save(); err != nil {
return false, err
}
return isSessionCookie, nil
}
-169
View File
@@ -1,169 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/core/contracts"
"context"
"crypto/sha256"
"encoding/hex"
"net/http"
"strconv"
"strings"
"sync"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
gsessions "github.com/gorilla/sessions"
)
var (
sessConfigMu sync.RWMutex
sessConfig = SessionConfig{
SessionCookieName: "wavelet_session",
SessionAge: 86400,
SessionHTTPOnly: true,
}
)
// SetSessionConfig updates the active session configuration.
func SetSessionConfig(cfg SessionConfig) {
sessConfigMu.Lock()
defer sessConfigMu.Unlock()
sessConfig = cfg
}
// GetSessionConfig returns the active session configuration.
func GetSessionConfig() SessionConfig {
sessConfigMu.RLock()
defer sessConfigMu.RUnlock()
return sessConfig
}
// GetSessionOptions 根据配置构建 Session 选项
func GetSessionOptions(maxAge int) sessions.Options {
cfg := GetSessionConfig()
return sessions.Options{
Path: "/",
Domain: cfg.SessionDomain,
MaxAge: maxAge,
HttpOnly: cfg.SessionHTTPOnly,
Secure: cfg.SessionSecure,
SameSite: http.SameSiteLaxMode,
}
}
// StripCookieMaxAgeAndExpires 从 Set-Cookie 响应头中移除 Max-Age 和 Expires,从而使其成为浏览器会话 Cookie
func StripCookieMaxAgeAndExpires(header http.Header, cookieName string) {
headers := header["Set-Cookie"]
if len(headers) == 0 {
return
}
newHeaders := make([]string, 0, len(headers))
for _, h := range headers {
if strings.HasPrefix(h, cookieName+"=") {
parts := strings.Split(h, ";")
newParts := make([]string, 0, len(parts))
for _, p := range parts {
trimmed := strings.TrimSpace(p)
lower := strings.ToLower(trimmed)
if strings.HasPrefix(lower, "max-age=") || strings.HasPrefix(lower, "expires=") {
continue
}
newParts = append(newParts, p)
}
newHeaders = append(newHeaders, strings.Join(newParts, ";"))
} else {
newHeaders = append(newHeaders, h)
}
}
header["Set-Cookie"] = newHeaders
}
// GetUserIDFromSession 从 Session 中提取用户 ID
func GetUserIDFromSession(s sessions.Session) uint64 {
val := s.Get(UserIDKey)
return ParseUserID(val)
}
// GetUserIDFromContext 从 Gin Context 的 Session 中提取用户 ID
func GetUserIDFromContext(c *gin.Context) (uid uint64) {
defer func() {
_ = recover()
}()
session := sessions.Default(c)
return GetUserIDFromSession(session)
}
func ensureSessionToken(s sessions.Session) (string, bool) {
token, ok := s.Get(SessionTokenKey).(string)
if !ok || token == "" {
token = uuid.NewString()
s.Set(SessionTokenKey, token)
return token, true
}
return token, false
}
func hashSessionToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
func rotateSessionID(s sessions.Session) {
if inner, ok := s.(interface{ Session() *gsessions.Session }); ok {
if sess := inner.Session(); sess != nil {
sess.ID = ""
}
}
}
// SetLoginSession writes the authenticated user into a freshly rotated session.
func SetLoginSession(ctx context.Context, c *gin.Context, user *contracts.UserDTO, extras ...map[string]any) error {
session := sessions.Default(c)
session.Clear()
rotateSessionID(session)
session.Set(UserIDKey, strconv.FormatUint(user.ID, 10))
session.Set(UserNameKey, user.Username)
if len(extras) > 0 {
for key, value := range extras[0] {
session.Set(key, value)
}
}
// 根据系统配置动态设置 Session 过期时间
cfg := GetSessionConfig()
maxAge := cfg.SessionAge
isSessionCookie := false
val, err := GetSystemConfigValue(ctx, "login_session_ttl_hours")
if err == nil && val != "" {
if ttlHours, err := strconv.Atoi(val); err == nil {
switch {
case ttlHours == -1:
// 永不过期,设置为 10 年
maxAge = 10 * 365 * 24 * 3600
case ttlHours > 0:
maxAge = ttlHours * 3600
case ttlHours == 0:
isSessionCookie = true
}
}
}
session.Options(GetSessionOptions(maxAge))
if err := session.Save(); err != nil {
return err
}
if isSessionCookie {
StripCookieMaxAgeAndExpires(c.Writer.Header(), cfg.SessionCookieName)
}
return nil
}
@@ -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"},
}
-24
View File
@@ -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
)
-38
View File
@@ -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()
}
}
-111
View File
@@ -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)
}
}
}
-56
View File
@@ -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() {}
-182
View File
@@ -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
}
+147 -20
View File
@@ -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, "&", "&amp;")
s = strings.ReplaceAll(s, "<", "&lt;")
s = strings.ReplaceAll(s, ">", "&gt;")
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
}
@@ -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,
@@ -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")
}
@@ -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
}
@@ -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
)
@@ -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"
)
@@ -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))
}
@@ -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
@@ -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
@@ -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"
@@ -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
}
@@ -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
}
@@ -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