diff --git a/internal/apps/admin/auth_source/routers.go b/internal/apps/admin/auth_source/routers.go deleted file mode 100644 index b4bbf430..00000000 --- a/internal/apps/admin/auth_source/routers.go +++ /dev/null @@ -1,250 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package auth_source 提供认证源管理功能 -package auth_source - -import ( - "errors" - "fmt" - "net/http" - "strings" - - "github.com/Rain-kl/Wavelet/internal/apps/admin" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -// AuthSourceRequest 创建或更新认证源的请求参数 -type AuthSourceRequest struct { - Name string `json:"name"` - Type string `json:"type"` - DisplayName string `json:"display_name"` - IsActive bool `json:"is_active"` - ClientID string `json:"client_id"` - ClientSecret string `json:"client_secret"` - OpenIDDiscoveryURL string `json:"openid_discovery_url"` - Scopes string `json:"scopes"` - IconURL string `json:"icon_url"` -} - -// ToggleAuthSourceRequest 切换认证源启用状态的请求参数 -type ToggleAuthSourceRequest struct { - IsActive bool `json:"is_active"` -} - -// ListAuthSources 获取认证源列表 -// @Summary 获取认证源列表 -// @Description 返回所有已配置的 OAuth/OIDC 认证源列表,包括已启用和未启用的,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=[]model.AuthSource} "认证源列表" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/auth-sources [get] -func ListAuthSources(c *gin.Context) { - sources, err := repository.GetAuthSources(c.Request.Context()) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(sources)) -} - -// CreateAuthSource 创建认证源 -// @Summary 创建认证源 -// @Description 创建一个新的 OAuth/OIDC 认证源配置,认证源名称必须唯一且符合命名规范,需要管理员权限 -// @Tags admin -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param request body auth_source.AuthSourceRequest true "创建认证源参数" -// @Success 200 {object} response.Any{data=model.AuthSource} "创建成功,返回认证源信息" -// @Failure 400 {object} response.Any "参数错误或验证失败" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Router /api/v1/admin/auth-sources [post] -func CreateAuthSource(c *gin.Context) { - var req AuthSourceRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - source := model.AuthSource{ - Name: req.Name, - Type: req.Type, - DisplayName: req.DisplayName, - IsActive: req.IsActive, - ClientID: req.ClientID, - ClientSecret: req.ClientSecret, - OpenIDDiscoveryURL: req.OpenIDDiscoveryURL, - Scopes: req.Scopes, - IconURL: req.IconURL, - } - if err := repository.CreateAuthSource(c.Request.Context(), &source); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - _ = repository.InvalidateAuthSourceCache(c.Request.Context()) - source.Sanitize() - c.JSON(http.StatusOK, response.OK(source)) -} - -// UpdateAuthSource 更新认证源 -// @Summary 更新认证源 -// @Description 更新指定 ID 的认证源配置。若 client_secret 字段为空,则保留原有密钥不变,需要管理员权限 -// @Tags admin -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param id path uint64 true "认证源 ID 或名称" -// @Param request body auth_source.AuthSourceRequest true "更新认证源参数" -// @Success 200 {object} response.Any{data=model.AuthSource} "更新成功,返回更新后的认证源信息" -// @Failure 400 {object} response.Any "参数错误或验证失败" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/auth-sources/{id} [put] -func UpdateAuthSource(c *gin.Context) { - id, err := parseSourceID(c) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - var req AuthSourceRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - // 记录更新前的 Discovery URL,以便更新成功后清除旧缓存条目。 - existing, _ := repository.GetAuthSourceByID(c.Request.Context(), id) - - source := model.AuthSource{ - ID: id, - Name: req.Name, - Type: req.Type, - DisplayName: req.DisplayName, - IsActive: req.IsActive, - ClientID: req.ClientID, - ClientSecret: req.ClientSecret, - OpenIDDiscoveryURL: req.OpenIDDiscoveryURL, - Scopes: req.Scopes, - IconURL: req.IconURL, - } - keepSecret := source.ClientSecret == "" - if err := repository.UpdateAuthSource(c.Request.Context(), &source, keepSecret); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - // Discovery URL 可能已变更,清除旧、新 issuer 的 provider 缓存, - // 确保下次登录时重新拉取最新 OIDC 元数据。 - if existing != nil { - oauth.InvalidateOIDCProviderCache(normalizeIssuer(existing.OpenIDDiscoveryURL)) - } - oauth.InvalidateOIDCProviderCache(normalizeIssuer(req.OpenIDDiscoveryURL)) - _ = repository.InvalidateAuthSourceCache(c.Request.Context()) - - updated, err := repository.GetAuthSourceByID(c.Request.Context(), id) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - updated.Sanitize() - c.JSON(http.StatusOK, response.OK(updated)) -} - -// ToggleAuthSource 切换认证源启用状态 -// @Summary 切换认证源启用状态 -// @Description 启用或禁用指定认证源。尝试启用时将验证 Client ID 和 Client Secret 是否已配置,需要管理员权限 -// @Tags admin -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param id path uint64 true "认证源 ID 或名称" -// @Param request body auth_source.ToggleAuthSourceRequest true "启用状态" -// @Success 200 {object} response.Any{data=string} "切换成功" -// @Failure 400 {object} response.Any "验证失败或认证源不存在" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Router /api/v1/admin/auth-sources/{id}/toggle [put] -func ToggleAuthSource(c *gin.Context) { - id, err := parseSourceID(c) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - var req ToggleAuthSourceRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - if err := repository.ToggleAuthSource(c.Request.Context(), id, req.IsActive); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - _ = repository.InvalidateAuthSourceCache(c.Request.Context()) - c.JSON(http.StatusOK, response.OKNil()) -} - -// DeleteAuthSource 删除认证源 -// @Summary 删除认证源 -// @Description 删除指定认证源及其关联的所有外部帐号绑定记录,警告:删除后相关用户将无法通过该源登录,需要管理员权限 -// @Tags admin -// @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 "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Router /api/v1/admin/auth-sources/{id} [delete] -func DeleteAuthSource(c *gin.Context) { - id, err := parseSourceID(c) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - if err := repository.DeleteAuthSource(c.Request.Context(), id); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - _ = repository.InvalidateAuthSourceCache(c.Request.Context()) - c.JSON(http.StatusOK, response.OKNil()) -} - -func parseSourceID(c *gin.Context) (uint64, error) { - raw := c.Param("id") - if raw == "" { - return 0, errors.New(admin.InvalidAuthSourceID) - } - source, err := repository.GetAuthSourceByName(c.Request.Context(), raw) - if err == nil { - return source.ID, nil - } - var id uint64 - if _, scanErr := fmt.Sscanf(raw, "%d", &id); scanErr != nil || id == 0 { - return 0, errors.New(admin.InvalidAuthSourceID) - } - return id, nil -} - -// normalizeIssuer 将 Discovery URL 规范化为 issuer 基础 URL, -// 与 oauth.buildOAuthConfig 中的规范化逻辑保持一致。 -func normalizeIssuer(discoveryURL string) string { - issuer := strings.TrimSuffix(strings.TrimSpace(discoveryURL), "/") - issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration") - issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server") - return issuer -} diff --git a/internal/apps/admin/auth_source/routers_test.go b/internal/apps/admin/auth_source/routers_test.go deleted file mode 100644 index 3d05bf49..00000000 --- a/internal/apps/admin/auth_source/routers_test.go +++ /dev/null @@ -1,333 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package auth_source - -import ( - "bytes" - "encoding/json" - "net/http" - "net/http/httptest" - "testing" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -func setupTestRouter(authUser *model.User) *gin.Engine { - r := testhelper.NewTestGinEngine() - adminGroup := r.Group("/api/v1/admin") - - // Mock authentication middleware - adminGroup.Use(func(c *gin.Context) { - if authUser != nil { - oauth.SetToContext(c, oauth.UserObjKey, authUser) - } - c.Next() - }) - - adminGroup.GET("/auth-sources", ListAuthSources) - adminGroup.POST("/auth-sources", CreateAuthSource) - adminGroup.PUT("/auth-sources/:id", UpdateAuthSource) - adminGroup.PUT("/auth-sources/:id/toggle", ToggleAuthSource) - adminGroup.DELETE("/auth-sources/:id", DeleteAuthSource) - return r -} - -func TestListAuthSources(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - // Seed source - source := model.AuthSource{ - ID: 1, - Name: "google", - Type: "oidc", - DisplayName: "Google Auth", - IsActive: true, - ClientID: "client_id_123", - ClientSecret: "client_secret_456", - OpenIDDiscoveryURL: "https://accounts.google.com", - } - dbConn.Create(&source) - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - req, _ := http.NewRequest("GET", "/api/v1/admin/auth-sources", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Errorf("expected 200 OK, got %d", w.Code) - } - - var resp response.Any - _ = json.Unmarshal(w.Body.Bytes(), &resp) - - dataBytes, _ := json.Marshal(resp.Data) - var sources []model.AuthSource - _ = json.Unmarshal(dataBytes, &sources) - - if len(sources) != 1 { - t.Errorf("expected 1 auth source, got %d", len(sources)) - } - if sources[0].Name != "google" { - t.Errorf("expected name 'google', got '%s'", sources[0].Name) - } - // Verify sanitize removed the secret - if sources[0].ClientSecret != "" { - t.Error("client secret should be sanitized") - } - if !sources[0].ClientSecretConfigured { - t.Error("client secret configured flag should be true") - } -} - -func TestCreateAuthSource(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("create successfully", func(t *testing.T) { - reqPayload := AuthSourceRequest{ - Name: "github", - Type: "oidc", - DisplayName: "GitHub OIDC", - IsActive: true, - ClientID: "client_id_gh", - ClientSecret: "client_secret_gh", - OpenIDDiscoveryURL: "https://github.com", - } - body, _ := json.Marshal(reqPayload) - req, _ := http.NewRequest("POST", "/api/v1/admin/auth-sources", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - // Verify database - var src model.AuthSource - dbConn.Where("name = ?", "github").First(&src) - if src.ClientID != "client_id_gh" { - t.Errorf("expected client_id_gh, got '%s'", src.ClientID) - } - }) - - t.Run("create invalid validation failure", func(t *testing.T) { - reqPayload := AuthSourceRequest{ - Name: "invalid name!", - Type: "oidc", - DisplayName: "Invalid", - IsActive: true, - ClientID: "client_id_val", - ClientSecret: "client_secret_val", - OpenIDDiscoveryURL: "https://discovery.url", - } - body, _ := json.Marshal(reqPayload) - req, _ := http.NewRequest("POST", "/api/v1/admin/auth-sources", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusBadRequest { - t.Errorf("expected 400 Bad Request, got %d", w.Code) - } - }) -} - -func TestUpdateAuthSource(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - // Seed source - source := model.AuthSource{ - ID: 1, - Name: "microsoft", - Type: "oidc", - DisplayName: "Microsoft", - IsActive: true, - ClientID: "old_client_id", - ClientSecret: "old_secret", - OpenIDDiscoveryURL: "https://login.microsoftonline.com", - } - dbConn.Create(&source) - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("update keep client secret", func(t *testing.T) { - reqPayload := AuthSourceRequest{ - Name: "microsoft", - Type: "oidc", - DisplayName: "Microsoft Updated", - IsActive: true, - ClientID: "new_client_id", - ClientSecret: "", // empty implies keeping existing secret - OpenIDDiscoveryURL: "https://login.microsoftonline.com", - } - body, _ := json.Marshal(reqPayload) - req, _ := http.NewRequest("PUT", "/api/v1/admin/auth-sources/1", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - var src model.AuthSource - dbConn.First(&src, 1) - if src.DisplayName != "Microsoft Updated" { - t.Errorf("expected display name update, got '%s'", src.DisplayName) - } - if src.ClientSecret != "old_secret" { - t.Errorf("expected old secret to be preserved, got '%s'", src.ClientSecret) - } - }) - - t.Run("update new client secret", func(t *testing.T) { - reqPayload := AuthSourceRequest{ - Name: "microsoft", - Type: "oidc", - DisplayName: "Microsoft Updated Again", - IsActive: true, - ClientID: "new_client_id", - ClientSecret: "brand_new_secret", - OpenIDDiscoveryURL: "https://login.microsoftonline.com", - } - body, _ := json.Marshal(reqPayload) - req, _ := http.NewRequest("PUT", "/api/v1/admin/auth-sources/microsoft", bytes.NewBuffer(body)) // Using Name instead of ID - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - var src model.AuthSource - dbConn.First(&src, 1) - if src.ClientSecret != "brand_new_secret" { - t.Errorf("expected secret update, got '%s'", src.ClientSecret) - } - }) -} - -func TestToggleAuthSource(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - source := model.AuthSource{ - ID: 1, - Name: "test_source", - Type: "oidc", - DisplayName: "Test Source", - IsActive: false, - ClientID: "", - ClientSecret: "", - OpenIDDiscoveryURL: "https://test.discovery.url", - } - dbConn.Create(&source) - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("cannot activate without credentials", func(t *testing.T) { - payload := ToggleAuthSourceRequest{IsActive: true} - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("PUT", "/api/v1/admin/auth-sources/1/toggle", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusBadRequest { - t.Errorf("expected 400 Bad Request when activating without client_id/secret, got %d", w.Code) - } - }) - - t.Run("toggle success after setting credentials", func(t *testing.T) { - // Set credentials first - dbConn.Model(&model.AuthSource{}).Where("id = ?", 1).Updates(map[string]interface{}{ - "client_id": "id", - "client_secret": "secret", - }) - - payload := ToggleAuthSourceRequest{IsActive: true} - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("PUT", "/api/v1/admin/auth-sources/1/toggle", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - var src model.AuthSource - dbConn.First(&src, 1) - if !src.IsActive { - t.Error("auth source should be activated") - } - }) -} - -func TestDeleteAuthSource(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - source := model.AuthSource{ - ID: 1, - Name: "delete_me", - Type: "oidc", - DisplayName: "Delete Me", - IsActive: true, - ClientID: "id", - ClientSecret: "secret", - OpenIDDiscoveryURL: "https://delete.me", - } - dbConn.Create(&source) - - externalAccount := model.ExternalAccount{ - ID: 10, - AuthSourceID: 1, - UserID: 50, - ExternalID: "ext_50", - } - dbConn.Create(&externalAccount) - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - req, _ := http.NewRequest("DELETE", "/api/v1/admin/auth-sources/1", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Errorf("expected 200 OK, got %d", w.Code) - } - - // Verify AuthSource is deleted - var srcCount int64 - dbConn.Model(&model.AuthSource{}).Where("id = ?", 1).Count(&srcCount) - if srcCount != 0 { - t.Error("AuthSource should be deleted from the database") - } - - // Verify ExternalAccount bindings are also deleted - var extCount int64 - dbConn.Model(&model.ExternalAccount{}).Where("auth_source_id = ?", 1).Count(&extCount) - if extCount != 0 { - t.Error("related ExternalAccount bindings should be deleted") - } -} diff --git a/internal/apps/admin/cache/logics.go b/internal/apps/admin/cache/logics.go deleted file mode 100644 index 76f8446a..00000000 --- a/internal/apps/admin/cache/logics.go +++ /dev/null @@ -1,14 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package cache - -import ( - "context" - - "github.com/Rain-kl/Wavelet/internal/repository" -) - -func saveOrUpdateConfig(ctx context.Context, key, value string) error { - return repository.SaveOrUpdateSystemConfig(ctx, key, value) -} diff --git a/internal/apps/admin/cache/routers.go b/internal/apps/admin/cache/routers.go deleted file mode 100644 index b9b1c799..00000000 --- a/internal/apps/admin/cache/routers.go +++ /dev/null @@ -1,100 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package cache provides HTTP handlers for managing disk cache. -package cache - -import ( - "net/http" - "strconv" - - "github.com/Rain-kl/Wavelet/internal/infra/diskcache" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -type updateCacheConfigRequest struct { - MaxSizeMB int64 `json:"max_size_mb" binding:"required,min=1"` - TTLMinutes int64 `json:"ttl_minutes" binding:"required,min=0"` - LRUEnabled bool `json:"lru_enabled"` -} - -// GetCacheStatus 获取磁盘缓存状态与当前统计数据 -// @Summary 获取缓存状态 -// @Description 获取当前系统磁盘缓存的使用情况(已占用字节、Key 数量等)与策略配置 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=diskcache.Status} "获取成功" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/cache/status [get] -func GetCacheStatus(c *gin.Context) { - status := diskcache.GetGlobalCache().Status() - c.JSON(http.StatusOK, response.OK(status)) -} - -// UpdateCacheConfig 更新磁盘缓存策略配置 -// @Summary 更新缓存配置 -// @Description 更改磁盘缓存最大容量限制、文件生存时间(TTL)以及是否启用 LRU 淘汰淘汰算法,并进行热更新 -// @Tags admin -// @Accept json -// @Produce json -// @Param request body cache.updateCacheConfigRequest true "缓存配置请求体" -// @Security SessionCookie -// @Success 200 {object} response.Any "更新成功" -// @Failure 400 {object} response.Any "参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "服务内部错误" -// @Router /api/v1/admin/cache/config [post] -func UpdateCacheConfig(c *gin.Context) { - var req updateCacheConfigRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - ctx := c.Request.Context() - - if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil { - response.AbortInternal(c, err.Error()) - return - } - - if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil { - response.AbortInternal(c, err.Error()) - return - } - - if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil { - response.AbortInternal(c, err.Error()) - return - } - - diskcache.GetGlobalCache().ReloadConfig(ctx) - - c.JSON(http.StatusOK, response.OKNil()) -} - -// ClearCache 一键清空所有磁盘缓存数据 -// @Summary 清空缓存 -// @Description 清除系统磁盘缓存目录中的所有临时文件,并重置缓存容量和 Key 追踪数据 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any "清理成功" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "服务内部错误" -// @Router /api/v1/admin/cache/clear [post] -func ClearCache(c *gin.Context) { - if err := diskcache.GetGlobalCache().Clear(); err != nil { - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OKNil()) -} diff --git a/internal/apps/admin/db_manage/routers.go b/internal/apps/admin/db_manage/routers.go deleted file mode 100644 index 972c32db..00000000 --- a/internal/apps/admin/db_manage/routers.go +++ /dev/null @@ -1,498 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package db_manage provides router handlers for managing database tables, -// overview information, and executing custom SQL queries. -package db_manage - -import ( - "database/sql" - "fmt" - "math" - "net/http" - "os" - "strings" - "time" - - "github.com/Rain-kl/Wavelet/internal/infra/config" - "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/gin-gonic/gin" - "gorm.io/gorm" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -const ( - binaryKB = 0 - binaryMB = 1 - binaryGB = 2 - valueThreshold = 10 - maxStringLength = 200 -) - -// DBOverviewResponse 数据库运行概览响应结构体 -type DBOverviewResponse struct { - Type string `json:"type"` - Version string `json:"version"` - Name string `json:"name"` - Size string `json:"size"` - TableCount int64 `json:"table_count"` - Connections int64 `json:"connections"` -} - -// GetTableDataRequest 分页拉取表数据请求结构体 -type GetTableDataRequest struct { - Table string `form:"table" binding:"required"` - Page int `form:"page,default=1"` - PageSize int `form:"pageSize,default=10"` -} - -// TableDataResponse 动态数据表响应结构体 -type TableDataResponse struct { - Columns []string `json:"columns"` - Total int64 `json:"total"` - Results []map[string]interface{} `json:"results"` -} - -// ExecuteSQLRequest 执行自定义 SQL 请求结构体 -type ExecuteSQLRequest struct { - SQL string `json:"sql" binding:"required"` -} - -// ExecuteSQLResponse 执行自定义 SQL 响应结构体 -type ExecuteSQLResponse struct { - Type string `json:"type"` // "select" 或 "exec" - Columns []string `json:"columns,omitempty"` - Results []map[string]interface{} `json:"results,omitempty"` - AffectedRows int64 `json:"affected_rows"` - ExecutionTimeMs int64 `json:"execution_time_ms"` -} - -// formatBytes 格式化字节大小为可读字符串 -func formatBytes(bytes uint64) string { - const unit = 1024 - if bytes < unit { - return fmt.Sprintf("%d B", bytes) - } - div, exp := int64(unit), 0 - for n := bytes / unit; n >= unit; n /= unit { - div *= unit - exp++ - } - value := float64(bytes) / float64(div) - var suffix string - switch exp { - case binaryKB: - suffix = "KiB" - case binaryMB: - suffix = "MiB" - case binaryGB: - suffix = "GiB" - default: - suffix = "TiB" - } - - if value == math.Trunc(value) { - if value >= valueThreshold { - return fmt.Sprintf("%.0f %s", value, suffix) - } - return fmt.Sprintf("%.1f %s", value, suffix) - } - return fmt.Sprintf("%.1f %s", value, suffix) -} - -// getSQLiteOverview 获取 SQLite 数据库概览信息 -func getSQLiteOverview(gormDB *gorm.DB) (DBOverviewResponse, error) { - name := config.Config.Database.SQLitePath - if name == "" { - name = "./data/wavelet.db" - } - - var version string - var ver string - if err := gormDB.Raw("SELECT sqlite_version()").Scan(&ver).Error; err == nil { - version = "SQLite " + ver - } else { - version = "SQLite" - } - - var sizeStr string - if fi, err := os.Stat(name); err == nil { - size := fi.Size() - if size < 0 { - size = 0 - } - sizeStr = formatBytes(uint64(size)) - } else { - sizeStr = "0 B" - } - - var tableCount int64 - if err := gormDB.Raw("SELECT count(*) FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%'").Scan(&tableCount).Error; err != nil { - tableCount = 0 - } - - var connCount int64 - if sqlDB, err := gormDB.DB(); err == nil { - connCount = int64(sqlDB.Stats().OpenConnections) - } else { - connCount = 1 - } - - return DBOverviewResponse{ - Type: "sqlite", - Version: version, - Name: name, - Size: sizeStr, - TableCount: tableCount, - Connections: connCount, - }, nil -} - -// getPostgresOverview 获取 PostgreSQL 数据库概览信息 -func getPostgresOverview(gormDB *gorm.DB) (DBOverviewResponse, error) { - name := config.Config.Database.Database - - var version string - var ver string - if err := gormDB.Raw("SELECT version()").Scan(&ver).Error; err == nil { - version = ver - } else { - version = "PostgreSQL" - } - - var sizeStr string - var sizeBytes sql.NullInt64 - if err := gormDB.Raw("SELECT pg_database_size(current_database())").Scan(&sizeBytes).Error; err == nil && sizeBytes.Valid { - size := sizeBytes.Int64 - if size < 0 { - size = 0 - } - sizeStr = formatBytes(uint64(size)) - } else { - sizeStr = "0 B" - } - - var tableCount int64 - if err := gormDB.Raw("SELECT count(*) FROM information_schema.tables WHERE table_schema = current_schema()").Scan(&tableCount).Error; err != nil { - tableCount = 0 - } - - var connCount int64 - var pgc sql.NullInt64 - if err := gormDB.Raw("SELECT count(*) FROM pg_stat_activity WHERE datname = current_database()").Scan(&pgc).Error; err == nil && pgc.Valid { - connCount = pgc.Int64 - } else { - if sqlDB, err := gormDB.DB(); err == nil { - connCount = int64(sqlDB.Stats().OpenConnections) - } else { - connCount = 1 - } - } - - return DBOverviewResponse{ - Type: "postgres", - Version: version, - Name: name, - Size: sizeStr, - TableCount: tableCount, - Connections: connCount, - }, nil -} - -// GetDBOverview 获取数据库运行概览 -// @Summary 获取数据库运行概览 -// @Description 获取数据库类型、版本、名称、文件大小、表数量及当前连接数,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=db_manage.DBOverviewResponse} "获取成功" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/db-manage/overview [get] -func GetDBOverview(c *gin.Context) { - gormDB := db.DB(c.Request.Context()) - if gormDB == nil { - response.AbortInternal(c, "数据库未初始化") - return - } - - var overview DBOverviewResponse - var err error - - if !config.Config.Database.Enabled { - overview, err = getSQLiteOverview(gormDB) - } else { - overview, err = getPostgresOverview(gormDB) - } - - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - - c.JSON(http.StatusOK, response.OK(overview)) -} - -// ListDBTables 获取数据库所有表名 -// @Summary 获取数据库所有表名 -// @Description 返回当前数据库的所有用户自定义表名称列表,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=[]string} "获取成功" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/db-manage/tables [get] -func ListDBTables(c *gin.Context) { - gormDB := db.DB(c.Request.Context()) - if gormDB == nil { - response.AbortInternal(c, "数据库未初始化") - return - } - - var tables []string - var err error - - if !config.Config.Database.Enabled { - err = gormDB.Raw("SELECT name FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%' ORDER BY name").Scan(&tables).Error - } else { - err = gormDB.Raw("SELECT table_name FROM information_schema.tables WHERE table_schema = current_schema() ORDER BY table_name").Scan(&tables).Error - } - - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - - c.JSON(http.StatusOK, response.OK(tables)) -} - -// GetDBTableData 获取数据表 data -func GetDBTableData(c *gin.Context) { - var req GetTableDataRequest - if err := c.ShouldBindQuery(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - gormDB := db.DB(c.Request.Context()) - if gormDB == nil { - response.AbortInternal(c, "数据库未初始化") - return - } - - // 安全转义表名并拼接 - quotedTable := `"` + strings.ReplaceAll(req.Table, `"`, `""`) + `"` - - var total int64 - if err := gormDB.Raw("SELECT count(*) FROM " + quotedTable).Scan(&total).Error; err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - offset := (req.Page - 1) * req.PageSize - if offset < 0 { - offset = 0 - } - limit := req.PageSize - if limit <= 0 { - limit = 10 - } - - rows, err := gormDB.Raw("SELECT * FROM "+quotedTable+" LIMIT ? OFFSET ?", limit, offset).Rows() - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - defer func() { - _ = rows.Close() - }() - - cols, err := rows.Columns() - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - - results, err := scanTableRows(rows, cols) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - - c.JSON(http.StatusOK, response.OK(TableDataResponse{ - Columns: cols, - Total: total, - Results: results, - })) -} - -// scanTableRows 扫描并提取数据表行数据,做截断处理 -func scanTableRows(rows *sql.Rows, cols []string) ([]map[string]interface{}, error) { - results := make([]map[string]interface{}, 0) - for rows.Next() { - columns := make([]interface{}, len(cols)) - columnPointers := make([]interface{}, len(cols)) - for i := range columns { - columnPointers[i] = &columns[i] - } - - if err := rows.Scan(columnPointers...); err != nil { - return nil, err - } - - rowMap := make(map[string]interface{}) - for i, colName := range cols { - val := columns[i] - if b, ok := val.([]byte); ok { - strVal := string(b) - runes := []rune(strVal) - if len(runes) > maxStringLength { - strVal = string(runes[:maxStringLength]) + "..." - } - rowMap[colName] = strVal - } else if str, ok := val.(string); ok { - runes := []rune(str) - if len(runes) > maxStringLength { - str = string(runes[:maxStringLength]) + "..." - } - rowMap[colName] = str - } else { - rowMap[colName] = val - } - } - results = append(results, rowMap) - } - return results, nil -} - -// executeSQLQuery 执行并解析查询类 SQL 语句 -func executeSQLQuery(gormDB *gorm.DB, sqlStr string, startTime time.Time) (ExecuteSQLResponse, error) { - rows, err := gormDB.Raw(sqlStr).Rows() - if err != nil { - return ExecuteSQLResponse{}, err - } - defer func() { - _ = rows.Close() - }() - - cols, err := rows.Columns() - if err != nil { - return ExecuteSQLResponse{}, err - } - - results := make([]map[string]interface{}, 0) - for rows.Next() { - columns := make([]interface{}, len(cols)) - columnPointers := make([]interface{}, len(cols)) - for i := range columns { - columnPointers[i] = &columns[i] - } - - if err := rows.Scan(columnPointers...); err != nil { - return ExecuteSQLResponse{}, err - } - - rowMap := make(map[string]interface{}) - for i, colName := range cols { - val := columns[i] - if b, ok := val.([]byte); ok { - rowMap[colName] = string(b) - } else { - rowMap[colName] = val - } - } - results = append(results, rowMap) - } - - executionTime := time.Since(startTime).Milliseconds() - return ExecuteSQLResponse{ - Type: "select", - Columns: cols, - Results: results, - AffectedRows: int64(len(results)), - ExecutionTimeMs: executionTime, - }, nil -} - -// executeSQLMutation 执行修改/更新类 SQL 语句 -func executeSQLMutation(gormDB *gorm.DB, sqlStr string, startTime time.Time) (ExecuteSQLResponse, error) { - tx := gormDB.Exec(sqlStr) - if tx.Error != nil { - return ExecuteSQLResponse{}, tx.Error - } - - executionTime := time.Since(startTime).Milliseconds() - return ExecuteSQLResponse{ - Type: "exec", - AffectedRows: tx.RowsAffected, - ExecutionTimeMs: executionTime, - }, nil -} - -// ExecuteSQL 执行 SQL 查询 -// @Summary 执行 SQL 查询 -// @Description 在当前数据库中执行任意自定义 SQL,如果是查询语句将返回格式化后的列与数据集,否则返回受影响行数,需要管理员权限 -// @Tags admin -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param request body db_manage.ExecuteSQLRequest true "SQL 请求参数" -// @Success 200 {object} response.Any{data=db_manage.ExecuteSQLResponse} "执行完毕" -// @Failure 400 {object} response.Any "SQL 语句错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/db-manage/query [post] -func ExecuteSQL(c *gin.Context) { - var req ExecuteSQLRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - gormDB := db.DB(c.Request.Context()) - if gormDB == nil { - response.AbortInternal(c, "数据库未初始化") - return - } - - trimmedSQL := strings.TrimSpace(req.SQL) - if trimmedSQL == "" { - response.AbortBadRequest(c, "SQL 语句不能为空") - return - } - - startTime := time.Now() - - // 识别是否是查询语句(SELECT, SHOW, EXPLAIN 等) - isQuery := false - lowerSQL := strings.ToLower(trimmedSQL) - queryKeywords := []string{"select", "show", "explain", "describe", "pragma"} - for _, kw := range queryKeywords { - if strings.HasPrefix(lowerSQL, kw) { - isQuery = true - break - } - } - - var resp ExecuteSQLResponse - var err error - - if isQuery { - resp, err = executeSQLQuery(gormDB, trimmedSQL, startTime) - } else { - resp, err = executeSQLMutation(gormDB, trimmedSQL, startTime) - } - - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - c.JSON(http.StatusOK, response.OK(resp)) -} diff --git a/internal/apps/admin/errs.go b/internal/apps/admin/errs.go deleted file mode 100644 index c869ed34..00000000 --- a/internal/apps/admin/errs.go +++ /dev/null @@ -1,15 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package admin 提供管理后台功能 -package admin - -// 管理后台错误消息常量 -const ( - AdminRequired = "未经授权访问" - TokenAdminRequired = "该访问令牌没有管理员权限,无法访问管理端点" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - InvalidAuthSourceID = "认证源 ID 无效" - InvalidCursorParam = "无效的 cursor 参数" - InvalidTaskExecutionID = "无效的任务执行记录 ID" -) diff --git a/internal/apps/admin/logs/routers.go b/internal/apps/admin/logs/routers.go deleted file mode 100644 index 958777fe..00000000 --- a/internal/apps/admin/logs/routers.go +++ /dev/null @@ -1,428 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package logs 提供日志查询与分析功能 -package logs - -import ( - "context" - "encoding/json" - "fmt" - "net/http" - "strconv" - "time" - - "github.com/Rain-kl/Wavelet/internal/apps/admin" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/repository/logstore" - "github.com/Rain-kl/Wavelet/pkg/logger" - "github.com/Rain-kl/Wavelet/pkg/util" - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -const ( - defaultLimit = 200 - maxLimit = 500 - maxPageSize = 100 - hoursInDay = 24 - analyticsDays = 7 - topActiveLimit = 10 -) - -// logsResponse 历史日志查询响应 -type logsResponse struct { - Lines []logger.LogEntry `json:"lines"` - HasMore bool `json:"has_more"` - NextCursor int `json:"next_cursor"` // 用于加载更早日志的 cursor -} - -// GetLogs 获取历史日志 -// @Summary 获取系统日志 -// @Description 分页获取系统历史日志,cursor=0 获取最新日志,cursor>0 获取更早日志 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Param cursor query int false "日志游标,0=获取最新" default(0) -// @Param limit query int false "每页条数" default(200) -// @Success 200 {object} response.Any{data=logs.logsResponse} "日志列表" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Router /api/v1/admin/logs [get] -func GetLogs(c *gin.Context) { - cursorStr := c.DefaultQuery("cursor", "0") - limitStr := c.DefaultQuery("limit", "200") - - var cursor, limit int - if _, err := parsePositiveInt(cursorStr, &cursor); err != nil { - response.AbortWithError(c, http.StatusBadRequest, admin.InvalidCursorParam) - return - } - if _, err := parsePositiveInt(limitStr, &limit); err != nil || limit <= 0 { - limit = defaultLimit - } - if limit > maxLimit { - limit = maxLimit - } - - entries, hasMore := logger.GlobalRingBuffer.Query(cursor, limit) - - resp := logsResponse{ - Lines: entries, - HasMore: hasMore, - } - if len(entries) > 0 { - resp.NextCursor = entries[0].Index - } - - c.JSON(http.StatusOK, response.OK(resp)) -} - -// wsMessage WebSocket 消息格式 -type wsMessage struct { - Type string `json:"type"` // "log" | "error" - Data json.RawMessage `json:"data"` -} - -// HandleLogWebSocket WebSocket 端点,实时推送系统日志 -// @Summary 系统日志实时推送 -// @Description 通过 WebSocket 实时推送系统日志,需要管理员权限 -// @Tags admin -// @Router /api/v1/admin/logs/ws [get] -func HandleLogWebSocket(c *gin.Context) { - upgrader := getUpgrader() - - conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) - if err != nil { - return - } - defer func() { _ = conn.Close() }() - - // 订阅 ring buffer - ch := logger.GlobalRingBuffer.Subscribe() - defer logger.GlobalRingBuffer.Unsubscribe(ch) - - // 在独立 goroutine 中读取客户端消息(保持连接活跃 + 检测断开) - done := make(chan struct{}) - util.Go(func() { - defer close(done) - for { - _, _, err := conn.ReadMessage() - if err != nil { - return - } - } - }) - - // 主循环:推送日志 - for { - select { - case <-done: - return - case entry, ok := <-ch: - if !ok { - return - } - data, _ := json.Marshal(entry) - msg := wsMessage{Type: "log", Data: data} - payload, _ := json.Marshal(msg) - if err := conn.WriteMessage(1, payload); err != nil { - return - } - } - } -} - -// accessLogItem 访问日志单条数据 -type accessLogItem struct { - ID uint64 `json:"id,string"` - UserID uint64 `json:"user_id,string"` - Username string `json:"username"` - Nickname string `json:"nickname"` - Path string `json:"path"` - Method string `json:"method"` - IP string `json:"ip"` - UserAgent string `json:"user_agent"` - Headers string `json:"headers"` - Status int32 `json:"status"` - Latency int64 `json:"latency"` - CreatedAt string `json:"created_at"` -} - -// accessLogsResponse 访问日志查询响应 -type accessLogsResponse struct { - Total uint64 `json:"total"` - List []accessLogItem `json:"list"` -} - -func buildAccessLogFilter(ctx context.Context, c *gin.Context) (logstore.AccessLogFilter, error) { - filter := logstore.AccessLogFilter{} - - username := c.Query("username") - if username != "" { - userIDs, err := repository.ListUserIDsByUsernameContains(ctx, username) - if err != nil { - return filter, fmt.Errorf("查询用户信息失败: %w", err) - } - filter.UserIDs = userIDs - } - - if path := c.Query("path"); path != "" { - filter.Path = path - } - - if startTime := c.Query("start_time"); startTime != "" { - if t, err := parseAccessLogTime(startTime); err == nil { - filter.StartTime = &t - } - } - - if endTime := c.Query("end_time"); endTime != "" { - if t, err := parseAccessLogTime(endTime); err == nil { - filter.EndTime = &t - } - } - - return filter, nil -} - -func parseAccessLogTime(value string) (time.Time, error) { - if t, err := time.Parse(time.RFC3339, value); err == nil { - return t, nil - } - return time.Parse("2006-01-02 15:04:05", value) -} - -func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) { - if len(list) == 0 { - return - } - - userIDs := make([]uint64, 0, len(list)) - seen := make(map[uint64]struct{}, len(list)) - for _, item := range list { - if _, ok := seen[item.UserID]; ok { - continue - } - seen[item.UserID] = struct{}{} - userIDs = append(userIDs, item.UserID) - } - - userMap := make(map[uint64]struct{ Username, Nickname string }) - if users, err := repository.ListUsersByIDs(ctx, userIDs); err == nil { - for _, u := range users { - userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname} - } - } - for i := range list { - if info, ok := userMap[list[i].UserID]; ok { - list[i].Username = info.Username - list[i].Nickname = info.Nickname - } - } -} - -// GetAccessLogs 获取 ClickHouse 异步采集的访问日志 -// @Summary 获取用户访问日志 -// @Description 分页并按照用户、接口路径、时间范围等维度检索用户访问日志列表(需要管理员权限) -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Param page query int false "页码" default(1) -// @Param page_size query int false "每页条数" default(20) -// @Param username query string false "用户名模糊搜索" -// @Param path query string false "接口路径模糊搜索" -// @Param start_time query string false "起始时间(RFC3339 或 YYYY-MM-DD HH:MM:SS)" -// @Param end_time query string false "结束时间(RFC3339 或 YYYY-MM-DD HH:MM:SS)" -// @Success 200 {object} response.Any{data=logs.accessLogsResponse} "访问日志列表" -// @Failure 400 {object} response.Any "参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/logs/access [get] -func GetAccessLogs(c *gin.Context) { - ctx := c.Request.Context() - store, err := logstore.Active(ctx) - if err != nil { - response.AbortInternal(c, "日志存储初始化失败") - return - } - - page, _ := strconv.Atoi(c.DefaultQuery("page", "1")) - if page < 1 { - page = 1 - } - pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20")) - if pageSize < 1 { - pageSize = 20 - } - if pageSize > maxPageSize { - pageSize = maxPageSize - } - - filter, err := buildAccessLogFilter(ctx, c) - if err != nil { - response.AbortWithError(c, http.StatusInternalServerError, err.Error()) - return - } - if filter.UserIDs != nil && len(filter.UserIDs) == 0 { - c.JSON(http.StatusOK, response.OK(accessLogsResponse{Total: 0, List: []accessLogItem{}})) - return - } - - logs, total, err := store.UserAccessLogs.List(ctx, filter, page, pageSize) - if err != nil { - response.AbortWithError(c, http.StatusInternalServerError, err.Error()) - return - } - if total == 0 { - c.JSON(http.StatusOK, response.OK(accessLogsResponse{Total: 0, List: []accessLogItem{}})) - return - } - - list := make([]accessLogItem, len(logs)) - for i, logItem := range logs { - list[i] = accessLogItem{ - ID: logItem.ID, - UserID: logItem.UserID, - Path: logItem.Path, - Method: logItem.Method, - IP: logItem.IP, - UserAgent: logItem.UserAgent, - Headers: logItem.Headers, - Status: logItem.Status, - Latency: logItem.Latency, - CreatedAt: logItem.CreatedAt.Format(time.RFC3339), - } - } - enrichAccessLogsWithUsers(ctx, list) - - c.JSON(http.StatusOK, response.OK(accessLogsResponse{ - Total: total, - List: list, - })) -} - -// trendItem 趋势图数据点 -type trendItem struct { - Date string `json:"date"` - Count uint64 `json:"count"` -} - -// browserItem 浏览器占比排行 -type browserItem struct { - Browser string `json:"browser"` - Count uint64 `json:"count"` -} - -// topUserItem 活跃用户数据 -type topUserItem struct { - UserID uint64 `json:"user_id,string"` - Username string `json:"username"` - Nickname string `json:"nickname"` - Count uint64 `json:"count"` -} - -// logsAnalyticsResponse 访问日志数据分析结果 -type logsAnalyticsResponse struct { - Trend []trendItem `json:"trend"` - Browsers []browserItem `json:"browsers"` - TopUsers []topUserItem `json:"top_users"` -} - -// GetLogsAnalytics 获取 ClickHouse 访问日志图表聚合指标 -// @Summary 获取访问日志分析数据 -// @Description 聚合统计最近 7 天的每日访问趋势、浏览器分布以及前 10 名最活跃用户排行(需要管理员权限) -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=logs.logsAnalyticsResponse} "分析统计数据" -// @Failure 500 {object} response.Any "内部错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Router /api/v1/admin/logs/analytics [get] -func GetLogsAnalytics(c *gin.Context) { - ctx := c.Request.Context() - store, err := logstore.Active(ctx) - if err != nil { - response.AbortInternal(c, "日志存储初始化失败") - return - } - - startTime := time.Now().AddDate(0, 0, -(analyticsDays - 1)).Truncate(hoursInDay * time.Hour) - - trendPoints, err := store.UserAccessLogs.GetDailyTrend(ctx, analyticsDays) - if err != nil { - response.AbortWithError(c, http.StatusInternalServerError, "查询访问趋势失败: "+err.Error()) - return - } - trendList := make([]trendItem, len(trendPoints)) - for i, point := range trendPoints { - trendList[i] = trendItem{ - Date: point.Date, - Count: point.Count, - } - } - - browserPoints, err := store.UserAccessLogs.GetBrowserDistribution(ctx, startTime) - if err != nil { - response.AbortWithError(c, http.StatusInternalServerError, "查询浏览器分布失败: "+err.Error()) - return - } - browserList := make([]browserItem, len(browserPoints)) - for i, point := range browserPoints { - browserList[i] = browserItem{ - Browser: point.Browser, - Count: point.Count, - } - } - - topUserPoints, err := store.UserAccessLogs.GetTopActiveUsers(ctx, startTime, topActiveLimit) - if err != nil { - response.AbortWithError(c, http.StatusInternalServerError, "查询活跃用户失败: "+err.Error()) - return - } - - topUsers := make([]topUserItem, len(topUserPoints)) - userIDs := make([]uint64, len(topUserPoints)) - for i, point := range topUserPoints { - topUsers[i] = topUserItem{ - UserID: point.UserID, - Count: point.Count, - } - userIDs[i] = point.UserID - } - - if len(userIDs) > 0 { - userProfileMap := make(map[uint64]struct { - Username string - Nickname string - }) - users, errProfile := repository.ListUsersByIDs(ctx, userIDs) - if errProfile == nil { - for _, u := range users { - userProfileMap[u.ID] = struct { - Username string - Nickname string - }{ - Username: u.Username, - Nickname: u.Nickname, - } - } - } - for i := range topUsers { - if profile, ok := userProfileMap[topUsers[i].UserID]; ok { - topUsers[i].Username = profile.Username - topUsers[i].Nickname = profile.Nickname - } - } - } - - c.JSON(http.StatusOK, response.OK(logsAnalyticsResponse{ - Trend: trendList, - Browsers: browserList, - TopUsers: topUsers, - })) -} diff --git a/internal/apps/admin/logs/tasks.go b/internal/apps/admin/logs/tasks.go deleted file mode 100644 index ed52b0cd..00000000 --- a/internal/apps/admin/logs/tasks.go +++ /dev/null @@ -1,219 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package logs - -import ( - "context" - "encoding/json" - "errors" - "fmt" - - "github.com/Rain-kl/Wavelet/internal/infra/config" - "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/repository/logstore" - "github.com/Rain-kl/Wavelet/pkg/logger" - "github.com/Rain-kl/Wavelet/plugins/domain/risk_control" -) - -const ( - // LogDBSwitchTask 切换日志数据库任务标识。 - LogDBSwitchTask = "logs:db_switch" - // TaskTypeLogDBSwitch 管理端任务类型。 - TaskTypeLogDBSwitch = "logs_db_switch" - - copyBatchSize = 1000 - targetPostgres = "postgres" - targetSQLite = "sqlite" - targetClickHouse = "clickhouse" -) - -// LogDBSwitchMeta 描述切换日志数据库任务。 -var LogDBSwitchMeta = task.TaskMeta{ - Type: TaskTypeLogDBSwitch, - AsynqTask: LogDBSwitchTask, - Name: "切换日志数据库", - Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)", - SupportsTime: false, - MaxRetry: task.DefaultMaxRetry, - Queue: task.QueueDefault, - Retryable: true, - Params: []task.TaskParam{ - {Name: "target", Label: "目标日志库", Type: "string", Required: true, - Placeholder: "postgres|sqlite|clickhouse", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse"}, - }, -} - -type logDBSwitchPayload struct { - Target string `json:"target"` -} - -// LogDBSwitchHandler 切换日志数据库任务处理器。 -type LogDBSwitchHandler struct{} - -// ValidatePayload 校验并规范化参数。 -func (h *LogDBSwitchHandler) ValidatePayload(payload []byte) ([]byte, error) { - var p logDBSwitchPayload - if err := json.Unmarshal(payload, &p); err != nil { - return nil, fmt.Errorf("参数解析失败: %w", err) - } - p.Target = normalizeTarget(p.Target) - if !validTarget(p.Target) { - return nil, fmt.Errorf("目标日志库不合法: %s", p.Target) - } - out, err := json.Marshal(p) - if err != nil { - return nil, err - } - return out, nil -} - -func normalizeTarget(v string) string { - switch v { - case targetPostgres, "postgresql": - return targetPostgres - case targetSQLite, "sqlite3": - return targetSQLite - case targetClickHouse, "ch": - return targetClickHouse - } - return v -} - -func validTarget(v string) bool { - return v == targetPostgres || v == targetSQLite || v == targetClickHouse -} - -// Execute 执行迁移。 -func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) { - var p logDBSwitchPayload - if err := json.Unmarshal(payload, &p); err != nil { - return nil, fmt.Errorf("参数解析失败: %w", err) - } - p.Target = normalizeTarget(p.Target) - if err := validateSwitch(ctx, p.Target); err != nil { - return nil, err - } - - source, err := currentLogDatabase(ctx) - if err != nil { - task.AppendLog(ctx, "读取日志主库失败: %v", err) - return nil, err - } - task.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target) - - if err := setMigrationFlag(ctx, "migrating"); err != nil { - return nil, err - } - defer func() { - if err := setMigrationFlag(ctx, ""); err != nil { - logger.ErrorF(ctx, "清除日志迁移冻结标记失败: %v", err) - } - }() - - if err := risk_control.Drain(ctx); err != nil { - return nil, fmt.Errorf("排空日志写入队列失败: %w", err) - } - - src, err := logstore.Active(ctx) - if err != nil { - return nil, err - } - dst, err := logstore.BuildForMigration(ctx, p.Target) - if err != nil { - return nil, err - } - - if _, err := dst.UserAccessLogs.DeleteAll(ctx); err != nil { - return nil, fmt.Errorf("清空目标用户访问日志失败: %w", err) - } - from, to, err := src.UserAccessLogs.MigrationRange(ctx) - if err != nil { - return nil, fmt.Errorf("读取源库时间范围失败: %w", err) - } - if !from.IsZero() && !to.IsZero() { - if err := dst.UserAccessLogs.EnsurePartitions(ctx, from, to); err != nil { - return nil, fmt.Errorf("预建目标分区失败: %w", err) - } - } - - if err := copyUserAccessLogs(ctx, src, dst); err != nil { - return nil, err - } - if err := flipLogDatabase(ctx, p.Target); err != nil { - return nil, err - } - logstore.InvalidateCache() - task.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target) - return &task.TaskResult{Message: fmt.Sprintf("日志数据库已从 %s 切换为 %s", source, p.Target)}, nil -} - -func validateSwitch(ctx context.Context, target string) error { - source, err := currentLogDatabase(ctx) - if err != nil { - return err - } - if source == target { - return errors.New("目标日志库与当前日志库相同,无需迁移") - } - switch target { - case targetClickHouse: - if !config.Config.ClickHouse.Enabled { - return errors.New("ClickHouse 未启用,无法迁移到 ClickHouse") - } - case targetPostgres: - if !config.Config.Database.Enabled { - return errors.New("PostgreSQL 未启用(当前主库为 SQLite),无法迁移到 PostgreSQL") - } - case targetSQLite: - if config.Config.Database.Enabled { - return errors.New("当前主库为 PostgreSQL,日志库不能设置为 SQLite") - } - } - return nil -} - -func currentLogDatabase(ctx context.Context) (string, error) { - cfg, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyLogDatabase) - if err != nil { - return "", fmt.Errorf("读取日志主库失败: %w", err) - } - if cfg.Value == "" { - return "", errors.New("日志主库配置为空") - } - return cfg.Value, nil -} - -func setMigrationFlag(ctx context.Context, v string) error { - return repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyLogDBMigration, v) -} - -func flipLogDatabase(ctx context.Context, target string) error { - return repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyLogDatabase, target) -} - -func copyUserAccessLogs(ctx context.Context, src, dst *logstore.Store) error { - var afterID uint64 - var copied int - for { - rows, err := src.UserAccessLogs.ListForMigration(ctx, afterID, copyBatchSize) - if err != nil { - return fmt.Errorf("读取源用户访问日志失败: %w", err) - } - if len(rows) == 0 { - break - } - if err := dst.UserAccessLogs.BatchInsert(ctx, rows); err != nil { - return fmt.Errorf("写入目标用户访问日志失败: %w", err) - } - afterID = rows[len(rows)-1].ID - copied += len(rows) - task.AppendLog(ctx, "已复制用户访问日志 %d 条", copied) - if len(rows) < copyBatchSize { - break - } - } - return nil -} diff --git a/internal/apps/admin/logs/utils.go b/internal/apps/admin/logs/utils.go deleted file mode 100644 index 3c93edcc..00000000 --- a/internal/apps/admin/logs/utils.go +++ /dev/null @@ -1,63 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package logs - -import ( - "net/http" - "net/url" - "strconv" - "strings" - - "github.com/gorilla/websocket" - - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" -) - -// getUpgrader 返回 WebSocket 升级器并执行 Origin 安全检查以防止 CSWSH 攻击 -func getUpgrader() *websocket.Upgrader { - return &websocket.Upgrader{ - CheckOrigin: func(r *http.Request) bool { - origin := r.Header.Get("Origin") - if origin == "" { - return true - } - - // 1. 同源检查 (Same-origin check) - u, err := url.Parse(origin) - if err == nil && strings.EqualFold(u.Host, r.Host) { - return true - } - - // 2. 检查配置的允许跨域 Origin (Check allowed origins in system config) - ctx := r.Context() - if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress); err == nil && sc.Value != "" { - originToCheck := strings.TrimRight(strings.TrimSpace(origin), "/") - allowedOrigins := strings.Split(sc.Value, ",") - for _, allowed := range allowedOrigins { - allowed = strings.TrimRight(strings.TrimSpace(allowed), "/") - if allowed != "" && strings.EqualFold(allowed, originToCheck) { - return true - } - } - } - return false - }, - } -} - -// parsePositiveInt 解析非负整数字符串 -func parsePositiveInt(s string, result *int) (bool, error) { - if s == "" { - *result = 0 - return true, nil - } - n, err := strconv.Atoi(s) - if err != nil || n < 0 { - return false, err - } - *result = n - return true, nil -} diff --git a/internal/apps/admin/logs/utils_test.go b/internal/apps/admin/logs/utils_test.go deleted file mode 100644 index 3768e2eb..00000000 --- a/internal/apps/admin/logs/utils_test.go +++ /dev/null @@ -1,118 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package logs - -import ( - "context" - "net/http" - "testing" - - "gorm.io/driver/sqlite" - "gorm.io/gorm" - - "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/model" -) - -func setupTestDB(t *testing.T) *gorm.DB { - dbConn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) - if err != nil { - t.Fatalf("failed to open sqlite in memory: %v", err) - } - err = dbConn.AutoMigrate(&model.SystemConfig{}) - if err != nil { - t.Fatalf("failed to migrate schema: %v", err) - } - db.SetDB(dbConn) - return dbConn -} - -func TestWebSocketCheckOrigin(t *testing.T) { - dbConn := setupTestDB(t) - - // Clean up global DB after test - defer db.SetDB(nil) - - // Seed ConfigKeyServerAddress with allowed frontend origin - allowedOrigin := "http://localhost:3000" - if err := dbConn.Create(&model.SystemConfig{ - Key: model.ConfigKeyServerAddress, - Value: allowedOrigin, - }).Error; err != nil { - t.Fatalf("failed to seed server address config: %v", err) - } - - upgrader := getUpgrader() - if upgrader.CheckOrigin == nil { - t.Fatal("expected CheckOrigin to be defined") - } - - tests := []struct { - name string - origin string - host string - wantOK bool - }{ - { - name: "empty origin (non-browser clients)", - origin: "", - host: "localhost:8000", - wantOK: true, - }, - { - name: "same-origin request", - origin: "http://localhost:8000", - host: "localhost:8000", - wantOK: true, - }, - { - name: "same-origin request case insensitive", - origin: "HTTP://LOCALHOST:8000", - host: "localhost:8000", - wantOK: true, - }, - { - name: "configured allowed origin request", - origin: "http://localhost:3000", - host: "localhost:8000", - wantOK: true, - }, - { - name: "configured allowed origin request with trailing slash", - origin: "http://localhost:3000/", - host: "localhost:8000", - wantOK: true, - }, - { - name: "unauthorized third-party origin", - origin: "http://evil.com", - host: "localhost:8000", - wantOK: false, - }, - { - name: "invalid origin format", - origin: "::not-a-valid-url", - host: "localhost:8000", - wantOK: false, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - req, err := http.NewRequestWithContext(context.Background(), "GET", "/api/v1/admin/logs/ws", nil) - if err != nil { - t.Fatalf("failed to create request: %v", err) - } - req.Host = tt.host - if tt.origin != "" { - req.Header.Set("Origin", tt.origin) - } - - got := upgrader.CheckOrigin(req) - if got != tt.wantOK { - t.Errorf("CheckOrigin() = %v, want %v", got, tt.wantOK) - } - }) - } -} diff --git a/internal/apps/admin/message_gateway/errs.go b/internal/apps/admin/message_gateway/errs.go deleted file mode 100644 index 8ca41cc0..00000000 --- a/internal/apps/admin/message_gateway/errs.go +++ /dev/null @@ -1,14 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package message_gateway - -const ( - errNameRequired = "name is required" - 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" - errChannelProbeFailed = "channel probe failed" - maskedSecret = "********" -) diff --git a/internal/apps/admin/message_gateway/handlers.go b/internal/apps/admin/message_gateway/handlers.go deleted file mode 100644 index 402dae19..00000000 --- a/internal/apps/admin/message_gateway/handlers.go +++ /dev/null @@ -1,157 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package message_gateway - -import ( - "net/http" - "strconv" - - "github.com/Rain-kl/Wavelet/internal/shared/response" - "github.com/gin-gonic/gin" -) - -// ListChannelDefinitions returns form schemas for supported channel types. -// @Summary List message gateway channel definitions -// @Description Returns form field definitions for Telegram and QQ channels -// @Tags admin-message-gateway -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=[]Definition} -// @Router /api/v1/admin/message-gateway/channels/definitions [get] -func ListChannelDefinitions(c *gin.Context) { - c.JSON(http.StatusOK, response.OK(channelDefinitions())) -} - -// ListChannels lists configured messaging channels with secrets masked. -// @Summary List message gateway channels -// @Description Returns all messaging channels; secrets are masked -// @Tags admin-message-gateway -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=[]ChannelDTO} -// @Router /api/v1/admin/message-gateway/channels [get] -func ListChannels(c *gin.Context) { - rows, err := listChannels(c.Request.Context()) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(rows)) -} - -// CreateChannel creates a messaging channel. -// @Summary Create message gateway channel -// @Description Creates a Telegram or QQ channel with encrypted credentials -// @Tags admin-message-gateway -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param request body CreateChannelRequest true "create body" -// @Success 200 {object} response.Any{data=ChannelDTO} -// @Failure 400 {object} response.Any -// @Router /api/v1/admin/message-gateway/channels [post] -func CreateChannel(c *gin.Context) { - var req CreateChannelRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - dto, err := createChannel(c.Request.Context(), req) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(dto)) -} - -// UpdateChannel patches a messaging channel. Empty secrets keep the previous values. -// @Summary Update message gateway channel -// @Description Updates a channel; empty secrets keep the current ciphertext -// @Tags admin-message-gateway -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param id path int true "channel id" -// @Param request body UpdateChannelRequest true "update body" -// @Success 200 {object} response.Any{data=ChannelDTO} -// @Failure 400 {object} response.Any -// @Failure 404 {object} response.Any -// @Router /api/v1/admin/message-gateway/channels/{id} [patch] -func UpdateChannel(c *gin.Context) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, "invalid channel id") - return - } - var req UpdateChannelRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - dto, err := updateChannel(c.Request.Context(), id, req) - if err != nil { - if err.Error() == errChannelNotFound { - response.AbortNotFound(c, err.Error()) - return - } - response.AbortBadRequest(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(dto)) -} - -// DeleteChannel removes a channel and its bindings/pairing codes. -// @Summary Delete message gateway channel -// @Description Deletes a channel and cascaded bindings and pairing codes -// @Tags admin-message-gateway -// @Produce json -// @Security SessionCookie -// @Param id path int true "channel id" -// @Success 200 {object} response.Any -// @Failure 404 {object} response.Any -// @Router /api/v1/admin/message-gateway/channels/{id} [delete] -func DeleteChannel(c *gin.Context) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, "invalid channel id") - return - } - if err := deleteChannel(c.Request.Context(), id); err != nil { - if err.Error() == errChannelNotFound { - response.AbortNotFound(c, err.Error()) - return - } - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OKNil()) -} - -// TestChannel probes stored credentials (Telegram getMe or QQ token). -// @Summary Test message gateway channel -// @Description Probes stored credentials without returning secrets -// @Tags admin-message-gateway -// @Produce json -// @Security SessionCookie -// @Param id path int true "channel id" -// @Success 200 {object} response.Any -// @Failure 400 {object} response.Any -// @Failure 404 {object} response.Any -// @Router /api/v1/admin/message-gateway/channels/{id}/test [post] -func TestChannel(c *gin.Context) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, "invalid channel id") - return - } - if err := probeChannel(c.Request.Context(), id); err != nil { - if err.Error() == errChannelNotFound { - response.AbortNotFound(c, err.Error()) - return - } - response.AbortBadRequest(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OKNil()) -} diff --git a/internal/apps/admin/message_gateway/logics.go b/internal/apps/admin/message_gateway/logics.go deleted file mode 100644 index 45509c0b..00000000 --- a/internal/apps/admin/message_gateway/logics.go +++ /dev/null @@ -1,347 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package message_gateway - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "io" - "net/http" - "strings" - "time" - - appgw "github.com/Rain-kl/Wavelet/internal/apps/message_gateway" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/tencent-connect/botgo/token" - "gorm.io/gorm" -) - -const defaultTelegramAPI = "https://api.telegram.org" - -// 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"` - Name string `json:"name"` - Fields []Field `json:"fields"` -} - -// CreateChannelRequest is the admin create body. -type CreateChannelRequest struct { - Name string `json:"name"` - Type string `json:"type"` - Enabled *bool `json:"enabled"` - BotToken string `json:"bot_token"` - AppID string `json:"app_id"` - AppSecret string `json:"app_secret"` - BaseURL string `json:"base_url"` - PortalHost string `json:"portal_host"` - Sandbox string `json:"sandbox"` -} - -// UpdateChannelRequest is the admin patch body. -type UpdateChannelRequest struct { - Name *string `json:"name"` - Enabled *bool `json:"enabled"` - BotToken string `json:"bot_token"` - AppID string `json:"app_id"` - AppSecret string `json:"app_secret"` - BaseURL *string `json:"base_url"` - PortalHost *string `json:"portal_host"` - Sandbox *string `json:"sandbox"` -} - -// ChannelDTO is a list/detail view with secrets masked. -type ChannelDTO struct { - ID uint64 `json:"id,string"` - Name string `json:"name"` - Type string `json:"type"` - OwnerScope string `json:"owner_scope"` - Enabled bool `json:"enabled"` - BotToken string `json:"bot_token,omitempty"` - AppID string `json:"app_id,omitempty"` - AppSecret string `json:"app_secret,omitempty"` - BaseURL string `json:"base_url,omitempty"` - PortalHost string `json:"portal_host,omitempty"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at"` -} - -func channelDefinitions() []Definition { - return []Definition{ - { - Type: model.MessageChannelTypeTelegram, - Name: "Telegram", - Fields: []Field{ - {Key: "bot_token", Type: "password", Required: true}, - {Key: "base_url", Type: "text"}, - }, - }, - { - Type: model.MessageChannelTypeQQ, - Name: "QQ", - Fields: []Field{ - {Key: "app_id", Required: true}, - {Key: "app_secret", Type: "password", Required: true}, - {Key: "portal_host", Type: "text"}, - }, - }, - } -} - -func createChannel(ctx context.Context, req CreateChannelRequest) (ChannelDTO, error) { - name := strings.TrimSpace(req.Name) - if name == "" { - return ChannelDTO{}, errors.New(errNameRequired) - } - typ := strings.TrimSpace(req.Type) - creds, extra, err := credentialsFromCreate(req) - if err != nil { - return ChannelDTO{}, err - } - cipher, err := appgw.EncryptCredentials(creds) - if err != nil { - return ChannelDTO{}, err - } - enabled := true - if req.Enabled != nil { - enabled = *req.Enabled - } - row := &model.MessageChannel{ - Name: name, - Type: typ, - OwnerScope: model.MessageOwnerScopeSystem, - Enabled: enabled, - Credentials: cipher, - Extra: appgw.EncodeExtra(extra), - } - if err := repository.CreateMessageChannel(ctx, row); err != nil { - return ChannelDTO{}, err - } - return toDTO(row, creds, extra), nil -} - -func updateChannel(ctx context.Context, id uint64, req UpdateChannelRequest) (ChannelDTO, error) { - row, err := repository.GetMessageChannel(ctx, id) - if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return ChannelDTO{}, errors.New(errChannelNotFound) - } - return ChannelDTO{}, err - } - creds, err := appgw.DecryptCredentials(row.Credentials) - if err != nil { - creds = map[string]string{} - } - extra := appgw.ParseExtra(row.Extra) - if req.Name != nil { - name := strings.TrimSpace(*req.Name) - if name == "" { - return ChannelDTO{}, errors.New(errNameRequired) - } - row.Name = name - } - if req.Enabled != nil { - row.Enabled = *req.Enabled - } - if token := strings.TrimSpace(req.BotToken); token != "" { - creds["bot_token"] = token - } - if appID := strings.TrimSpace(req.AppID); appID != "" { - creds["app_id"] = appID - } - if secret := strings.TrimSpace(req.AppSecret); secret != "" { - creds["app_secret"] = secret - } - if req.BaseURL != nil { - extra["base_url"] = strings.TrimSpace(*req.BaseURL) - } - if req.PortalHost != nil { - extra["portal_host"] = strings.TrimSpace(*req.PortalHost) - } - if req.Sandbox != nil { - extra["sandbox"] = strings.TrimSpace(*req.Sandbox) - } - if err := validateCredentials(row.Type, creds); err != nil { - return ChannelDTO{}, err - } - cipher, err := appgw.EncryptCredentials(creds) - if err != nil { - return ChannelDTO{}, err - } - row.Credentials = cipher - row.Extra = appgw.EncodeExtra(extra) - if err := repository.UpdateMessageChannel(ctx, row); err != nil { - return ChannelDTO{}, err - } - return toDTO(row, creds, extra), nil -} - -func listChannels(ctx context.Context) ([]ChannelDTO, error) { - rows, err := repository.ListMessageChannels(ctx) - if err != nil { - return nil, err - } - out := make([]ChannelDTO, 0, len(rows)) - for i := range rows { - creds, err := appgw.DecryptCredentials(rows[i].Credentials) - if err != nil { - creds = map[string]string{} - } - out = append(out, toDTO(&rows[i], creds, appgw.ParseExtra(rows[i].Extra))) - } - return out, nil -} - -func deleteChannel(ctx context.Context, id uint64) error { - if _, err := repository.GetMessageChannel(ctx, id); err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return errors.New(errChannelNotFound) - } - return err - } - return repository.DeleteMessageChannel(ctx, id) -} - -func probeChannel(ctx context.Context, id uint64) error { - row, err := repository.GetMessageChannel(ctx, id) - if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return errors.New(errChannelNotFound) - } - return err - } - creds, err := appgw.DecryptCredentials(row.Credentials) - if err != nil { - return err - } - extra := appgw.ParseExtra(row.Extra) - if err := probeCredentials(ctx, row.Type, creds, extra); err != nil { - return fmt.Errorf("%s: %w", errChannelProbeFailed, err) - } - return nil -} - -func credentialsFromCreate(req CreateChannelRequest) (map[string]string, map[string]string, error) { - typ := strings.TrimSpace(req.Type) - creds := map[string]string{} - extra := map[string]string{} - switch typ { - case model.MessageChannelTypeTelegram: - creds["bot_token"] = strings.TrimSpace(req.BotToken) - if base := strings.TrimSpace(req.BaseURL); base != "" { - extra["base_url"] = base - } - case model.MessageChannelTypeQQ: - creds["app_id"] = strings.TrimSpace(req.AppID) - creds["app_secret"] = strings.TrimSpace(req.AppSecret) - if host := strings.TrimSpace(req.PortalHost); host != "" { - extra["portal_host"] = host - } else { - extra["portal_host"] = "q.qq.com" - } - if sandbox := strings.TrimSpace(req.Sandbox); sandbox != "" { - extra["sandbox"] = sandbox - } - default: - return nil, nil, errors.New(errTypeInvalid) - } - if err := validateCredentials(typ, creds); err != nil { - return nil, nil, err - } - return creds, extra, nil -} - -func validateCredentials(typ string, creds map[string]string) error { - switch typ { - case model.MessageChannelTypeTelegram: - if strings.TrimSpace(creds["bot_token"]) == "" { - return errors.New(errTelegramTokenRequired) - } - case model.MessageChannelTypeQQ: - if strings.TrimSpace(creds["app_id"]) == "" || strings.TrimSpace(creds["app_secret"]) == "" { - return errors.New(errQQCredentialsRequired) - } - default: - return errors.New(errTypeInvalid) - } - return nil -} - -func toDTO(row *model.MessageChannel, creds, extra map[string]string) ChannelDTO { - dto := ChannelDTO{ - ID: row.ID, - Name: row.Name, - Type: row.Type, - OwnerScope: row.OwnerScope, - Enabled: row.Enabled, - CreatedAt: row.CreatedAt, - UpdatedAt: row.UpdatedAt, - } - if strings.TrimSpace(creds["bot_token"]) != "" { - dto.BotToken = maskedSecret - } - if id := strings.TrimSpace(creds["app_id"]); id != "" { - dto.AppID = id - } - if strings.TrimSpace(creds["app_secret"]) != "" { - dto.AppSecret = maskedSecret - } - dto.BaseURL = extra["base_url"] - dto.PortalHost = extra["portal_host"] - return dto -} - -func probeCredentials(ctx context.Context, typ string, creds, extra map[string]string) error { - switch typ { - case model.MessageChannelTypeTelegram: - base := strings.TrimSpace(extra["base_url"]) - if base == "" { - base = defaultTelegramAPI - } - url := strings.TrimRight(base, "/") + "/bot" + creds["bot_token"] + "/getMe" - req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) - if err != nil { - return err - } - resp, err := http.DefaultClient.Do(req) - if err != nil { - return err - } - defer func() { _ = resp.Body.Close() }() - const probeBodyLimit = 4096 - body, _ := io.ReadAll(io.LimitReader(resp.Body, probeBodyLimit)) - if resp.StatusCode != http.StatusOK { - return fmt.Errorf("telegram getMe status %d", resp.StatusCode) - } - var parsed struct { - OK bool `json:"ok"` - } - if err := json.Unmarshal(body, &parsed); err != nil { - return err - } - if !parsed.OK { - return errors.New("telegram getMe returned ok=false") - } - return nil - case model.MessageChannelTypeQQ: - src := token.NewQQBotTokenSource(&token.QQBotCredentials{ - AppID: creds["app_id"], - AppSecret: creds["app_secret"], - }) - _, err := src.Token() - return err - default: - return errors.New(errTypeInvalid) - } -} diff --git a/internal/apps/admin/message_gateway/logics_test.go b/internal/apps/admin/message_gateway/logics_test.go deleted file mode 100644 index 334718b9..00000000 --- a/internal/apps/admin/message_gateway/logics_test.go +++ /dev/null @@ -1,85 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package message_gateway - -import ( - "context" - "strings" - "testing" - - appgw "github.com/Rain-kl/Wavelet/internal/apps/message_gateway" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/testhelper" -) - -func TestCreateChannel_TelegramRequiresToken(t *testing.T) { - _, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - _, err := createChannel(context.Background(), CreateChannelRequest{Name: "tg", Type: "telegram"}) - if err == nil { - t.Fatal("createChannel() error = nil, want token required") - } -} - -func TestCreateChannel_StoresCiphertextNotPlaintext(t *testing.T) { - _, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - const token = "secret-token-xyz" - dto, err := createChannel(context.Background(), CreateChannelRequest{ - Name: "tg", - Type: "telegram", - BotToken: token, - }) - if err != nil { - t.Fatalf("createChannel() error = %v", err) - } - if strings.Contains(dto.BotToken, token) { - t.Fatalf("createChannel() dto leaked plaintext token %q", dto.BotToken) - } - row, err := repository.GetMessageChannel(context.Background(), dto.ID) - if err != nil { - t.Fatalf("GetMessageChannel() error = %v", err) - } - if strings.Contains(row.Credentials, token) { - t.Fatalf("stored credentials contain plaintext token") - } - creds, err := appgw.DecryptCredentials(row.Credentials) - if err != nil { - t.Fatalf("DecryptCredentials() error = %v", err) - } - if creds["bot_token"] != token { - t.Fatalf("DecryptCredentials() bot_token = %q, want %q", creds["bot_token"], token) - } -} - -func TestUpdateChannel_EmptySecretKeepsPrevious(t *testing.T) { - _, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - created, err := createChannel(context.Background(), CreateChannelRequest{ - Name: "tg", - Type: "telegram", - BotToken: "old-token", - }) - if err != nil { - t.Fatalf("createChannel() error = %v", err) - } - name := "renamed" - if _, err := updateChannel(context.Background(), created.ID, UpdateChannelRequest{Name: &name}); err != nil { - t.Fatalf("updateChannel() error = %v", err) - } - row, err := repository.GetMessageChannel(context.Background(), created.ID) - if err != nil { - t.Fatalf("GetMessageChannel() error = %v", err) - } - if row.Name != "renamed" { - t.Fatalf("updateChannel() name = %q, want renamed", row.Name) - } - creds, err := appgw.DecryptCredentials(row.Credentials) - if err != nil { - t.Fatalf("DecryptCredentials() error = %v", err) - } - if creds["bot_token"] != "old-token" { - t.Fatalf("updateChannel() bot_token = %q, want old-token", creds["bot_token"]) - } -} diff --git a/internal/apps/admin/message_gateway/routers.go b/internal/apps/admin/message_gateway/routers.go deleted file mode 100644 index 6005f5e6..00000000 --- a/internal/apps/admin/message_gateway/routers.go +++ /dev/null @@ -1,20 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package message_gateway provides admin HTTP APIs for messaging channels. -package message_gateway - -import "github.com/gin-gonic/gin" - -// RegisterRoutes mounts admin message-gateway APIs under /admin. -func RegisterRoutes(adminRouter *gin.RouterGroup) { - g := adminRouter.Group("/message-gateway") - { - g.GET("/channels/definitions", ListChannelDefinitions) - g.GET("/channels", ListChannels) - g.POST("/channels", CreateChannel) - g.PATCH("/channels/:id", UpdateChannel) - g.DELETE("/channels/:id", DeleteChannel) - g.POST("/channels/:id/test", TestChannel) - } -} diff --git a/internal/apps/admin/middlewares.go b/internal/apps/admin/middlewares.go deleted file mode 100644 index c40215cb..00000000 --- a/internal/apps/admin/middlewares.go +++ /dev/null @@ -1,46 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package admin - -import ( - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/shared/response" - "github.com/Rain-kl/Wavelet/pkg/logger" - otel_trace "github.com/Rain-kl/Wavelet/pkg/trace" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/gin-gonic/gin" -) - -// LoginAdminRequired 返回管理员权限校验中间件 -func LoginAdminRequired() gin.HandlerFunc { - return func(c *gin.Context) { - // init trace - ctx, span := otel_trace.Start(c.Request.Context(), "LoginAdminRequired") - defer span.End() - - user, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) - - // 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限 - if tokenAuth, _ := oauth.GetFromContext[bool](c, oauth.TokenAuthKey); tokenAuth { - tokenAdmin, _ := oauth.GetFromContext[bool](c, oauth.TokenAdminKey) - if !tokenAdmin { - response.AbortNotFound(c, TokenAdminRequired) - return - } - } - - if !user.IsAdmin { - response.AbortNotFound(c, AdminRequired) - return - } - - // log - logger.InfoF(ctx, "[LoginAdminRequired] %d %s", user.ID, user.Username) - - // next - c.Next() - } -} diff --git a/internal/apps/admin/push/channels.go b/internal/apps/admin/push/channels.go deleted file mode 100644 index 0b205734..00000000 --- a/internal/apps/admin/push/channels.go +++ /dev/null @@ -1,278 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package push - -import ( - "encoding/json" - "errors" - "net/http" - "strconv" - "strings" - - pkgpush "github.com/Rain-kl/Wavelet/pkg/push" - "github.com/gin-gonic/gin" - "gorm.io/gorm" - - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -// ListChannelDefinitions 获取各种消息通道的表单配置定义列表 -// @Summary 获取所有消息通道配置字段定义 -// @Description 返回系统支持的所有消息通道类型(如飞书、邮件、自定义、Telegram)的动态表单定义,需要管理员权限 -// @Tags admin-push -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=[]Definition} "通道配置定义列表" -// @Router /api/v1/admin/push/channels/definitions [get] -func ListChannelDefinitions(c *gin.Context) { - c.JSON(http.StatusOK, response.OK(ListDefinitions())) -} - -// ListChannels 获取消息通道列表 -// @Summary 获取所有消息通道 -// @Description 返回系统配置的所有消息通道列表,需要管理员权限 -// @Tags admin-push -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=[]model.PushChannel} "消息通道列表" -// @Router /api/v1/admin/push/channels [get] -func ListChannels(c *gin.Context) { - channels, err := listPushChannels(c.Request.Context()) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(channels)) -} - -// CreateChannelRequest 创建通道参数 -type CreateChannelRequest struct { - Name string `json:"name" binding:"required"` - Description string `json:"description"` - Type string `json:"type" binding:"required"` - Token string `json:"token"` - URL string `json:"url"` - Other string `json:"other"` - Enabled bool `json:"enabled"` -} - -// CreateChannel 创建消息通道 -// @Summary 创建消息通道 -// @Description 新建一个消息通道配置,需要管理员权限 -// @Tags admin-push -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param request body CreateChannelRequest true "创建参数" -// @Success 200 {object} response.Any{data=model.PushChannel} "创建成功" -// @Router /api/v1/admin/push/channels [post] -func CreateChannel(c *gin.Context) { - var req CreateChannelRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - channel, err := createPushChannel(c.Request.Context(), req) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(channel)) -} - -// UpdateChannelRequest 修改通道参数 -type UpdateChannelRequest struct { - Description string `json:"description"` - Type string `json:"type" binding:"required"` - Token string `json:"token"` - URL string `json:"url"` - Other string `json:"other"` - Enabled bool `json:"enabled"` -} - -// UpdateChannel 更新消息通道 -// @Summary 更新消息通道 -// @Description 修改消息通道配置,需要管理员权限 -// @Tags admin-push -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param id path uint64 true "通道ID" -// @Param request body UpdateChannelRequest true "更新参数" -// @Success 200 {object} response.Any{data=model.PushChannel} "更新成功" -// @Router /api/v1/admin/push/channels/{id} [put] -func UpdateChannel(c *gin.Context) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, "invalid channel id") - return - } - - var req UpdateChannelRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - channel, err := updatePushChannel(c.Request.Context(), id, req) - if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - response.AbortNotFound(c, "channel not found") - return - } - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(channel)) -} - -// DeleteChannel 删除消息通道 -// @Summary 删除消息通道 -// @Description 根据ID删除消息通道,需要管理员权限 -// @Tags admin-push -// @Produce json -// @Security SessionCookie -// @Param id path uint64 true "通道ID" -// @Success 200 {object} response.Any "删除成功" -// @Router /api/v1/admin/push/channels/{id} [delete] -func DeleteChannel(c *gin.Context) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, "invalid channel id") - return - } - - if err := deletePushChannel(c.Request.Context(), id); err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - response.AbortNotFound(c, "channel not found") - return - } - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OKNil()) -} - -// TestChannelRequest 测试通道连通性参数 -type TestChannelRequest struct { - Name string `json:"name"` - Type string `json:"type"` - Token string `json:"token"` - URL string `json:"url"` - Other string `json:"other"` - Target string `json:"target"` -} - -// TestChannel 测试通道连通性 -// @Summary 测试通道连通性 -// @Description 触发一次临时的或现有的通道连通性推送测试,需要管理员权限 -// @Tags admin-push -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param request body TestChannelRequest true "测试参数" -// @Success 200 {object} response.Any "测试触发成功" -// @Router /api/v1/admin/push/channels/test [post] -func TestChannel(c *gin.Context) { - var req TestChannelRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - ctx := c.Request.Context() - url, token, other, channelType, err := loadChannelForTest(ctx, req) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - if channelType == channelEmail { - url, token, other = resolveSMTPConfig(ctx, url, token, other) - } - - tempChannel := model.PushChannel{ - Name: "test_temp", - URL: url, - Token: token, - Other: other, - Type: channelType, - Enabled: true, - } - if err := tempChannel.Validate(); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - url = tempChannel.URL - - var config pkgpush.Config - var renderedJSON string - switch channelType { - case channelLark: - config = pkgpush.Config{Channel: channelLark, URL: url, Secret: token} - renderedJSON = other - case channelEmail: - config = pkgpush.Config{Channel: channelEmail, URL: url, Key: token, Secret: other} - case channelTelegram: - config = pkgpush.Config{Channel: channelTelegram, URL: url, Secret: token, Key: other} - default: - config = pkgpush.Config{Channel: channelCustom, URL: url} - customPushReq := CustomPushRequest{ - Title: "通道测试通知", - Content: "这是一条来自系统的消息通道连通性测试消息。", - Description: "系统通道测试", - URL: "https://example.com", - To: req.Target, - } - renderedJSON = renderCustomPayload(other, customPushReq) - } - - payload := SendPayload{ - EventKey: "test_channel", - Config: config, - Target: req.Target, - Body: NotificationMessage{ - Title: "通道测试通知", - Content: "这是一条来自系统的消息通道连通性测试消息。", - Level: defaultLevelInfo, - }, - Template: renderedJSON, - } - if err := enqueuePushTask(ctx, payload); err != nil { - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OKNil()) -} - -// CustomPushRequest 外部公开推送请求参数 -type CustomPushRequest struct { - Title string `json:"title" form:"title"` - Description string `json:"description" form:"description"` - Content string `json:"content" form:"content"` - URL string `json:"url" form:"url"` - To string `json:"to" form:"to"` - Token string `json:"token" form:"token"` -} - -func escapeJSONString(s string) string { - b, _ := json.Marshal(s) - const minJSONLen = 2 - if len(b) >= minJSONLen { - return string(b[1 : len(b)-1]) - } - return s -} - -func renderCustomPayload(template string, req CustomPushRequest) string { - result := template - result = strings.ReplaceAll(result, "$title", escapeJSONString(req.Title)) - result = strings.ReplaceAll(result, "$description", escapeJSONString(req.Description)) - result = strings.ReplaceAll(result, "$content", escapeJSONString(req.Content)) - result = strings.ReplaceAll(result, "$url", escapeJSONString(req.URL)) - result = strings.ReplaceAll(result, "$to", escapeJSONString(req.To)) - return result -} diff --git a/internal/apps/admin/push/channels_definition.go b/internal/apps/admin/push/channels_definition.go deleted file mode 100644 index 0d810569..00000000 --- a/internal/apps/admin/push/channels_definition.go +++ /dev/null @@ -1,183 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package push - -import "sync" - -const ( - // KeyURL represents the URL field key - KeyURL = "url" - // KeyToken represents the Token field key - KeyToken = "token" - // KeyOther represents the Other field key - KeyOther = "other" - - // TypeText represents standard text input type - TypeText = "text" - // TypePassword represents password input type - TypePassword = "password" - // TypeTextarea represents textarea input type - TypeTextarea = "textarea" -) - -// Field represents a form field configuration for a channel. -type Field struct { - Key string `json:"key"` // unique key for the field (e.g. url, token, other) - Label string `json:"label"` // human readable label (e.g. "Webhook 地址") - Type string `json:"type"` // input type: "text" | "password" | "textarea" - Required bool `json:"required"` // whether this field is required - Placeholder string `json:"placeholder"` // input placeholder - Description string `json:"description"` // field explanation/help text -} - -// Definition represents the metadata and form schema for a notification channel. -type Definition struct { - Type string `json:"type"` // channel type (e.g., custom, lark, email) - Name string `json:"name"` // display name - Description string `json:"description"` // short description - Fields []Field `json:"fields"` // form fields -} - -var ( - defMu sync.RWMutex - definitions = make(map[string]Definition) -) - -// RegisterChannelDefinition registers a channel definition. -func RegisterChannelDefinition(def Definition) { - defMu.Lock() - defer defMu.Unlock() - definitions[def.Type] = def -} - -// ListDefinitions returns all registered channel definitions. -func ListDefinitions() []Definition { - defMu.RLock() - defer defMu.RUnlock() - - // We want a stable order: custom, lark, telegram, email - order := []string{channelCustom, channelLark, channelTelegram, channelEmail} - res := make([]Definition, 0, len(definitions)) - for _, t := range order { - if d, ok := definitions[t]; ok { - res = append(res, d) - } - } - // Add any others - for t, d := range definitions { - found := false - for _, o := range order { - if o == t { - found = true - break - } - } - if !found { - res = append(res, d) - } - } - return res -} - -func init() { - // Register custom webhook channel - RegisterChannelDefinition(Definition{ - Type: channelCustom, - Name: "自定义消息通道", - Description: "使用自定义 HTTP POST 请求向外部 Webhook 发送数据。", - Fields: []Field{ - { - Key: KeyURL, - Label: "请求地址", - Type: TypeText, - Required: true, - Placeholder: "在此填写完整的请求地址,必须使用 HTTPS 协议", - Description: "接口请求的完整 HTTPS URL,例如 https://api.example.com/webhook", - }, - { - Key: KeyOther, - Label: "请求体 (JSON)", - Type: TypeTextarea, - Required: true, - Placeholder: "在此输入请求体,支持模板变量,必须为合法的 JSON 格式", - Description: "可使用的变量:$title, $description, $content, $url, $to。例如 {\"text\": \"$content\"}", - }, - }, - }) - - // Register Lark robot channel - RegisterChannelDefinition(Definition{ - Type: channelLark, - Name: "飞书群机器人", - Description: "配置飞书群自定义机器人的 Webhook 接口投递。", - Fields: []Field{ - { - Key: KeyURL, - Label: "Webhook 地址", - Type: TypeText, - Required: true, - Placeholder: "https://open.feishu.cn/open-apis/bot/v2/hook/YOUR_TOKEN", - Description: "从飞书群机器人设置中复制 of Webhook URL", - // Note: using 'of' was in feishu.go, let's keep original wording or fix it - }, - { - Key: KeyToken, - Label: "签名校验密钥 (Secret) (可选)", - Type: TypeText, - Required: false, - Placeholder: "可选,若机器人启用了安全设置中的签名校验,请在此输入", - Description: "飞书群机器人安全设置中的签名校验 Key", - }, - { - Key: KeyOther, - Label: "自定义卡片 JSON 模版 (可选)", - Type: TypeTextarea, - Required: false, - Placeholder: "可选,留空则默认使用系统内置的精美互动卡片", - Description: "若填写,必须是合法的飞书卡片 JSON 格式", - }, - }, - }) - - // Register Telegram channel - RegisterChannelDefinition(Definition{ - Type: channelTelegram, - Name: "Telegram 机器人", - Description: "配置 Telegram 机器人推送消息。", - Fields: []Field{ - { - Key: KeyURL, - Label: "API 基础地址 (可选)", - Type: TypeText, - Required: false, - Placeholder: "https://api.telegram.org", - Description: "接口请求的 HTTPS 基础地址,留空默认为 https://api.telegram.org", - }, - { - Key: KeyToken, - Label: "机器人 Token (Bot Token)", - Type: TypePassword, - Required: true, - Placeholder: "在此输入 Telegram 机器人的 Bot Token", - Description: "通过 BotFather 申请到的机器人 Access Token", - }, - { - Key: KeyOther, - Label: "默认会话 ID (Chat ID) (可选)", - Type: TypeText, - Required: false, - Placeholder: "例如 -100123456789 或 @channel_name", - Description: "默认的消息接收 Chat ID。如果通知事件中未配置 targets,将推送到此 ID", - }, - }, - }) - - // Register Email channel - RegisterChannelDefinition(Definition{ - Type: channelEmail, - Name: "邮件推送通道", - Description: "邮件推送通道直接使用系统全局 SMTP 设置进行发送,无需在此填写服务器配置。", - Fields: []Field{}, - }) -} diff --git a/internal/apps/admin/push/constants.go b/internal/apps/admin/push/constants.go deleted file mode 100644 index 0b0e8860..00000000 --- a/internal/apps/admin/push/constants.go +++ /dev/null @@ -1,15 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package push - -const ( - channelCustom = "custom" - channelEmail = "email" - channelLark = "lark" - channelTelegram = "telegram" - defaultLevelInfo = "INFO" - keyTitle = "title" - keyContent = "content" - keyLevel = "level" -) diff --git a/internal/apps/admin/push/custom_events/admin_login.go b/internal/apps/admin/push/custom_events/admin_login.go deleted file mode 100644 index 5fb25226..00000000 --- a/internal/apps/admin/push/custom_events/admin_login.go +++ /dev/null @@ -1,38 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package custom_events defines custom push notification events. -package custom_events - -import ( - "context" - "time" - - "github.com/Rain-kl/Wavelet/internal/apps/admin/push" - "github.com/Rain-kl/Wavelet/internal/listener" -) - -// AdminLogin is the metadata definition for the admin login event. -var AdminLogin = push.EventMetadata{ - Key: "admin_login", - Name: "管理员登录", - DefaultTemplate: push.NotificationMessage{ - Title: "管理员登录提醒", - Content: "管理员 {{user.username}} 于 {{time}} 从 IP {{ip}} 登录系统。", - Level: "INFO", - }, - Description: "当管理员成功登录系统时触发此通知", -} - -func handleAdminLogin(ctx context.Context, event listener.AdminLoggedIn) { - if event.User == nil { - return - } - - body := map[string]any{ - "user": event.User, - "ip": event.IP, - "time": time.Now().Format("2006-01-02 15:04:05"), - } - push.DefaultTrigger.Trigger(ctx, AdminLogin, body) -} diff --git a/internal/apps/admin/push/custom_events/admin_login_test.go b/internal/apps/admin/push/custom_events/admin_login_test.go deleted file mode 100644 index 45d6d9c2..00000000 --- a/internal/apps/admin/push/custom_events/admin_login_test.go +++ /dev/null @@ -1,176 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package custom_events - -import ( - "context" - "encoding/json" - "sync" - "testing" - "time" - - "github.com/Rain-kl/Wavelet/internal/apps/admin/push" - "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/internal/listener" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/hibiken/asynq" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - "gorm.io/gorm" -) - -var registerOnce sync.Once - -func ensureRegistered() { - registerOnce.Do(Register) -} - -func setupAdminLoginIntegrationTest(t *testing.T) (*gorm.DB, func()) { - t.Helper() - - dbConn, mr, cleanup := testhelper.SetupTestEnvironment(t) - - err := dbConn.AutoMigrate( - &model.PushEvent{}, - &model.PushHistory{}, - &model.PushChannel{}, - ) - require.NoError(t, err) - - sysUser := &model.User{ - ID: 999, - Username: "system", - Nickname: "系统", - Password: "*", - IsActive: true, - } - require.NoError(t, dbConn.Create(sysUser).Error) - - task.AsynqClient = asynq.NewClient(asynq.RedisClientOpt{Addr: mr.Addr()}) - task.RegisterHandler(push.SendNotificationTask, &push.PushHandler{}) - task.RegisterTaskMeta(push.SendNotificationMeta) - - ensureRegistered() - - require.NoError(t, push.SyncEvents(context.Background())) - - return dbConn, func() { - cleanup() - if task.AsynqClient != nil { - task.AsynqClient.Close() - task.AsynqClient = nil - } - } -} - -func seedMockPushChannel(t *testing.T, dbConn *gorm.DB) *model.PushChannel { - t.Helper() - - channel := &model.PushChannel{ - Name: "mock_channel", - Type: "custom", - URL: "https://webhook.site/admin-login", - Other: `{"text": "$content"}`, - Enabled: true, - } - require.NoError(t, dbConn.Create(channel).Error) - return channel -} - -func enableAdminLoginEvent(t *testing.T, dbConn *gorm.DB, channelName string, targets []string) { - t.Helper() - - var event model.PushEvent - require.NoError(t, dbConn.Where("event_key = ?", AdminLogin.Key).First(&event).Error) - - event.Enabled = true - event.Channels = []string{channelName} - event.Targets = targets - require.NoError(t, repository.SavePushEvent(context.Background(), &event)) -} - -func waitForAsyncTrigger(t *testing.T) { - t.Helper() - time.Sleep(100 * time.Millisecond) -} - -func countPushTasks(t *testing.T, dbConn *gorm.DB) int64 { - t.Helper() - - var count int64 - require.NoError(t, dbConn.Model(&model.TaskExecution{}). - Where("task_type = ?", push.SendNotificationTask). - Count(&count).Error) - return count -} - -func TestAdminLoginPushIntegration(t *testing.T) { - dbConn, cleanup := setupAdminLoginIntegrationTest(t) - defer cleanup() - - channel := seedMockPushChannel(t, dbConn) - defer dbConn.Delete(channel) - - enableAdminLoginEvent(t, dbConn, channel.Name, []string{"ops_team"}) - - adminUser := &model.User{ - ID: 1001, - Username: "super_admin", - IsAdmin: true, - IsActive: true, - } - require.NoError(t, dbConn.Create(adminUser).Error) - - t.Run("admin login emits push task with user and ip", func(t *testing.T) { - dbConn.Where("task_type = ?", push.SendNotificationTask).Delete(&model.TaskExecution{}) - - listener.EmitAdminLoggedIn(context.Background(), adminUser, "203.0.113.42") - waitForAsyncTrigger(t) - - var execution model.TaskExecution - require.NoError(t, dbConn.Where("task_type = ?", push.SendNotificationTask).First(&execution).Error) - - var payload push.SendPayload - require.NoError(t, json.Unmarshal([]byte(execution.Payload), &payload)) - - assert.Equal(t, AdminLogin.Key, payload.EventKey) - assert.Equal(t, "ops_team", payload.Target) - assert.Equal(t, "管理员登录提醒", payload.Body.Title) - assert.Contains(t, payload.Body.Content, "super_admin") - assert.Contains(t, payload.Body.Content, "203.0.113.42") - }) - - t.Run("non-admin login does not trigger push", func(t *testing.T) { - dbConn.Where("task_type = ?", push.SendNotificationTask).Delete(&model.TaskExecution{}) - - nonAdmin := &model.User{ - ID: 2002, - Username: "regular_user", - IsAdmin: false, - IsActive: true, - } - require.NoError(t, dbConn.Create(nonAdmin).Error) - - listener.EmitAdminLoggedIn(context.Background(), nonAdmin, "198.51.100.1") - waitForAsyncTrigger(t) - - assert.Equal(t, int64(0), countPushTasks(t, dbConn)) - }) - - t.Run("disabled admin login event does not enqueue push", func(t *testing.T) { - dbConn.Where("task_type = ?", push.SendNotificationTask).Delete(&model.TaskExecution{}) - - var event model.PushEvent - require.NoError(t, dbConn.Where("event_key = ?", AdminLogin.Key).First(&event).Error) - event.Enabled = false - require.NoError(t, repository.SavePushEvent(context.Background(), &event)) - - listener.EmitAdminLoggedIn(context.Background(), adminUser, "10.0.0.1") - waitForAsyncTrigger(t) - - assert.Equal(t, int64(0), countPushTasks(t, dbConn)) - }) -} diff --git a/internal/apps/admin/push/custom_events/register.go b/internal/apps/admin/push/custom_events/register.go deleted file mode 100644 index 08ec205c..00000000 --- a/internal/apps/admin/push/custom_events/register.go +++ /dev/null @@ -1,17 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package custom_events - -import ( - "github.com/Rain-kl/Wavelet/internal/apps/admin/push" - "github.com/Rain-kl/Wavelet/internal/listener" -) - -// Register wires push notification handlers for domain events and registers -// built-in event metadata. Must be called once during application bootstrap -// before push.SyncEvents. -func Register() { - push.RegisterBuiltInEvent(AdminLogin) - listener.OnAdminLoggedIn(handleAdminLogin) -} diff --git a/internal/apps/admin/push/events.go b/internal/apps/admin/push/events.go deleted file mode 100644 index 16d02209..00000000 --- a/internal/apps/admin/push/events.go +++ /dev/null @@ -1,344 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package push defines push notification HTTP routes, background tasks, and events. -package push - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "strings" - - "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/pkg/logger" - pkgpush "github.com/Rain-kl/Wavelet/pkg/push" - "github.com/Rain-kl/Wavelet/pkg/util" - "gorm.io/gorm" -) - -// NotificationMessage represents the structured notification message payload. -type NotificationMessage struct { - Title string `json:"title"` - Content string `json:"content"` - Level string `json:"level"` - Ext map[string]any `json:"ext,omitempty"` -} - -// Flatten converts the structured NotificationMessage back to a flat map (original json structure). -func (m NotificationMessage) Flatten() map[string]any { - res := map[string]any{ - keyTitle: m.Title, - keyContent: m.Content, - keyLevel: m.Level, - } - for k, v := range m.Ext { - res[k] = v - } - return res -} - -// EventMetadata represents the metadata of a push notification event. -type EventMetadata struct { - Key string `json:"key"` - Name string `json:"name"` - DefaultTemplate NotificationMessage `json:"default_template"` - Description string `json:"description"` -} - -// SendPayload 异步投递推送载荷 (供 task/Worker 使用) -type SendPayload struct { - EventKey string `json:"event_key"` - Config pkgpush.Config `json:"config"` - Target string `json:"target"` - Body NotificationMessage `json:"body"` - Template string `json:"template"` -} - -// BuiltInEvents lists all built-in events defined in custom_events. -var BuiltInEvents []EventMetadata - -// RegisterBuiltInEvent registers a built-in event definition. -func RegisterBuiltInEvent(meta EventMetadata) { - BuiltInEvents = append(BuiltInEvents, meta) -} - -// EventTrigger represents the unified event trigger class. -type EventTrigger struct{} - -// DefaultTrigger is the singleton instance of EventTrigger. -var DefaultTrigger = &EventTrigger{} - -// Trigger receives event metadata and processes the event notification dispatch asynchronously. -// -//nolint:contextcheck -func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map[string]any) { - asyncCtx := context.WithoutCancel(ctx) - util.Go(func() { - if body == nil { - body = make(map[string]any) - } - if _, hasUser := body["user"]; !hasUser || body["user"] == nil { - body["user"] = getSystemUser(asyncCtx) - } - - eventPtr, err := repository.GetActivePushEventByKey(asyncCtx, meta.Key) - if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return - } - logger.ErrorF(asyncCtx, "push_event_trigger: failed to get active event %s: %v", meta.Key, err) - return - } - event := *eventPtr - if len(event.Channels) == 0 { - return - } - - flatBody := getFlatBody(body) - msg, _ := t.buildMessage(&event, meta, flatBody, body) - t.enqueuePushTasks(asyncCtx, meta, &event, msg, flatBody) - }) -} - -func (t *EventTrigger) buildMessage(event *model.PushEvent, meta EventMetadata, flatBody map[string]any, body map[string]any) (NotificationMessage, string) { - var msg NotificationMessage - renderedTemplate := "" - - templateSource := event.Template - if templateSource != "" { - var err error - msg, renderedTemplate, err = t.parseCustomTemplate(event, templateSource, flatBody) - if err != nil { - msg.Title = event.Name - msg.Content = renderedTemplate - msg.Level = defaultLevelInfo - } - } else { - msg = t.parseDefaultTemplate(meta, flatBody) - } - - if msg.Ext == nil { - msg.Ext = make(map[string]any) - } - for k, v := range body { - if k == keyTitle || k == keyContent || k == keyLevel { - continue - } - if _, exists := msg.Ext[k]; !exists { - msg.Ext[k] = v - } - } - - return msg, renderedTemplate -} - -func (t *EventTrigger) parseCustomTemplate(event *model.PushEvent, templateSource string, flatBody map[string]any) (NotificationMessage, string, error) { - var msg NotificationMessage - renderedTemplate := pkgpush.ParseTemplate(templateSource, flatBody) - - var tMap map[string]any - if err := json.Unmarshal([]byte(renderedTemplate), &tMap); err != nil { - return msg, renderedTemplate, err - } - - if title, ok := tMap[keyTitle].(string); ok && title != "" { - msg.Title = title - } else { - msg.Title = event.Name - } - delete(tMap, keyTitle) - - if content, ok := tMap[keyContent].(string); ok && content != "" { - msg.Content = content - } else { - msg.Content = renderedTemplate - } - delete(tMap, keyContent) - - if level, ok := tMap[keyLevel].(string); ok && level != "" { - msg.Level = level - } else { - msg.Level = defaultLevelInfo - } - delete(tMap, keyLevel) - - msg.Ext = tMap - return msg, renderedTemplate, nil -} - -func (t *EventTrigger) parseDefaultTemplate(meta EventMetadata, flatBody map[string]any) NotificationMessage { - var msg NotificationMessage - msg.Title = pkgpush.ParseTemplate(meta.DefaultTemplate.Title, flatBody) - msg.Content = pkgpush.ParseTemplate(meta.DefaultTemplate.Content, flatBody) - msg.Level = pkgpush.ParseTemplate(meta.DefaultTemplate.Level, flatBody) - - if meta.DefaultTemplate.Ext != nil { - msg.Ext = make(map[string]any) - for k, v := range meta.DefaultTemplate.Ext { - if strVal, ok := v.(string); ok { - msg.Ext[k] = pkgpush.ParseTemplate(strVal, flatBody) - } else { - msg.Ext[k] = v - } - } - } - return msg -} - -func (t *EventTrigger) enqueuePushTasks(ctx context.Context, meta EventMetadata, event *model.PushEvent, msg NotificationMessage, flatBody map[string]any) { - for _, channelName := range event.Channels { - customChannel, err := repository.GetActivePushChannelByName(ctx, channelName) - if err == nil { - t.enqueueCustomPushChannelTasks(ctx, meta, event, customChannel, msg, flatBody) - continue - } - logger.WarnF(ctx, "push_event_trigger: channel %q not found in DB or disabled: %v", channelName, err) - } -} - -func (t *EventTrigger) enqueueCustomPushChannelTasks(ctx context.Context, meta EventMetadata, event *model.PushEvent, channel *model.PushChannel, msg NotificationMessage, flatBody map[string]any) { - if len(event.Targets) == 0 { - t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, "", msg) - return - } - - for _, target := range event.Targets { - resolvedTarget := resolveTarget(ctx, target, flatBody, channel.Name) - t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, resolvedTarget, msg) - } -} - -func (t *EventTrigger) enqueueSingleCustomPushChannelTask(ctx context.Context, meta EventMetadata, channel *model.PushChannel, target string, msg NotificationMessage) { - var config pkgpush.Config - var renderedTemplate string - - switch channel.Type { - case channelLark: - config = pkgpush.Config{Channel: channelLark, URL: channel.URL, Secret: channel.Token} - renderedTemplate = channel.Other - case channelEmail: - url, token, other := resolveSMTPConfig(ctx, channel.URL, channel.Token, channel.Other) - config = pkgpush.Config{Channel: channelEmail, URL: url, Key: token, Secret: other} - case channelTelegram: - config = pkgpush.Config{Channel: channelTelegram, URL: channel.URL, Secret: channel.Token, Key: channel.Other} - default: - config = pkgpush.Config{Channel: channelCustom, URL: channel.URL} - customPushReq := CustomPushRequest{ - Title: msg.Title, - Content: msg.Content, - Description: meta.Description, - To: target, - } - if urlVal, ok := msg.Ext["url"].(string); ok { - customPushReq.URL = urlVal - } - renderedTemplate = renderCustomPayload(channel.Other, customPushReq) - } - - payload := SendPayload{ - EventKey: meta.Key, - Config: config, - Target: target, - Body: msg, - Template: renderedTemplate, - } - if err := enqueuePushTask(ctx, payload); err != nil { - logger.ErrorF(ctx, "push_event_trigger: enqueuePushTask failed for %s channel %s -> %s: %v", channel.Type, channel.Name, target, err) - } -} - -func enqueuePushTask(ctx context.Context, payload SendPayload) error { - payloadBytes, err := json.Marshal(payload) - if err != nil { - return err - } - _, err = task.DispatchTask(ctx, "send_notification", payloadBytes, "system") - return err -} - -func getFlatBody(body map[string]any) map[string]any { - jsonBytes, err := json.Marshal(body) - if err != nil { - return body - } - var jsonMap map[string]any - if err := json.Unmarshal(jsonBytes, &jsonMap); err != nil { - return body - } - - flatResult := make(map[string]any) - flattenMap("", jsonMap, flatResult) - return flatResult -} - -func flattenMap(prefix string, m map[string]any, result map[string]any) { - for k, v := range m { - key := k - if prefix != "" { - key = prefix + "." + k - } - if nestedMap, ok := v.(map[string]any); ok { - flattenMap(key, nestedMap, result) - } else { - result[key] = v - } - } -} - -func resolveTarget(ctx context.Context, target string, flatBody map[string]any, channel string) string { - target = strings.TrimSpace(target) - if target == "" { - return "" - } - - resolved := resolveDynamicKeyword(target, flatBody) - if strings.Contains(resolved, "@") { - return resolved - } - if val, matched := resolveSystemTarget(ctx, resolved, channel); matched { - return val - } - - user, found := resolveTargetUser(ctx, resolved, channel) - if !found { - return resolved - } - if channel == channelEmail && user.Email != "" { - return user.Email - } - if channel != channelEmail && user.Username != "" { - return user.Username - } - return resolved -} - -func resolveDynamicKeyword(target string, flatBody map[string]any) string { - switch target { - case "user.id", "id": - if val, ok := flatBody["user.id"]; ok { - return fmt.Sprintf("%v", val) - } - if val, ok := flatBody["id"]; ok { - return fmt.Sprintf("%v", val) - } - case "user.username", "username": - if val, ok := flatBody["user.username"]; ok { - return fmt.Sprintf("%v", val) - } - if val, ok := flatBody["username"]; ok { - return fmt.Sprintf("%v", val) - } - case "user.email", channelEmail: - if val, ok := flatBody["user.email"]; ok { - return fmt.Sprintf("%v", val) - } - if val, ok := flatBody["email"]; ok { - return fmt.Sprintf("%v", val) - } - } - return target -} diff --git a/internal/apps/admin/push/logics.go b/internal/apps/admin/push/logics.go deleted file mode 100644 index 713797e9..00000000 --- a/internal/apps/admin/push/logics.go +++ /dev/null @@ -1,413 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package push - -import ( - "context" - "encoding/json" - "errors" - "strconv" - "strings" - - "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - pkgpush "github.com/Rain-kl/Wavelet/pkg/push" - "gorm.io/gorm" -) - -type smtpConfig struct { - Host string - Port string - Username string - Password string -} - -func loadSMTPConfig(ctx context.Context) smtpConfig { - host, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPHost) - port, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPort) - user, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPUsername) - pass, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPassword) - return smtpConfig{ - Host: host.Value, - Port: port.Value, - Username: user.Value, - Password: pass.Value, - } -} - -func syncBuiltInEvents(ctx context.Context) error { - for _, meta := range BuiltInEvents { - _, err := repository.GetPushEventByKey(ctx, meta.Key) - if errors.Is(err, gorm.ErrRecordNotFound) { - var defaultTemplateStr string - if defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate); err == nil { - defaultTemplateStr = string(defaultTemplateBytes) - } - event := model.PushEvent{ - EventKey: meta.Key, - Name: meta.Name, - Channels: []string{}, - Targets: []string{}, - Template: defaultTemplateStr, - Enabled: false, - } - if err := repository.CreatePushEvent(ctx, &event); err != nil { - return err - } - } else if err != nil { - return err - } - } - return nil -} - -func listPushEvents(ctx context.Context) ([]model.PushEvent, error) { - return repository.ListPushEvents(ctx) -} - -func createPushEvent(ctx context.Context, req CreateEventRequest) (model.PushEvent, error) { - eventKey, eventName, defaultTemplateBytes, err := getEventInfo(req) - if err != nil { - return model.PushEvent{}, err - } - - count, err := repository.CountPushEventsByKey(ctx, eventKey) - if err != nil { - return model.PushEvent{}, err - } - if count > 0 { - return model.PushEvent{}, errors.New("this notification event is already configured") - } - - templateStr := strings.TrimSpace(req.Template) - if templateStr == "" { - templateStr = string(defaultTemplateBytes) - } else { - var tempMap map[string]any - if err := json.Unmarshal([]byte(templateStr), &tempMap); err != nil { - return model.PushEvent{}, errors.New("custom template is not a valid JSON format") - } - } - - channels := req.Channels - if channels == nil { - channels = []string{} - } - targets := req.Targets - if targets == nil { - targets = []string{} - } - - event := model.PushEvent{ - EventKey: eventKey, - Name: eventName, - TaskType: req.TaskType, - Channels: channels, - Targets: targets, - Template: templateStr, - Enabled: req.Enabled, - } - if err := event.Validate(); err != nil { - return model.PushEvent{}, err - } - if err := repository.CreatePushEvent(ctx, &event); err != nil { - return model.PushEvent{}, err - } - return event, nil -} - -func deletePushEvent(ctx context.Context, id uint64) error { - event, err := repository.GetPushEventByID(ctx, id) - if err != nil { - return err - } - return repository.DeletePushEvent(ctx, &event) -} - -func updatePushEvent(ctx context.Context, id uint64, req UpdateEventRequest) error { - event, err := repository.GetPushEventByID(ctx, id) - if err != nil { - return err - } - - event.Channels = req.Channels - event.Targets = req.Targets - event.Template = req.Template - event.Enabled = req.Enabled - if err := event.Validate(); err != nil { - return err - } - return repository.SavePushEvent(ctx, &event) -} - -func togglePushEvent(ctx context.Context, id uint64) (bool, error) { - event, err := repository.GetPushEventByID(ctx, id) - if err != nil { - return false, err - } - - enabled := !event.Enabled - if enabled && len(event.Channels) == 0 { - return false, errors.New("cannot enable event without any push channels configured") - } - if err := repository.UpdatePushEventEnabled(ctx, &event, enabled); err != nil { - return false, err - } - return enabled, nil -} - -func listPushHistories(ctx context.Context, filter repository.PushHistoryListFilter) (int64, []model.PushHistory, error) { - return repository.ListPushHistories(ctx, filter) -} - -func applySMTPFallbackToPushConfig(ctx context.Context, cfg *pkgpush.Config) { - if cfg.Channel != channelEmail || (cfg.URL != "" && cfg.Key != "") { - return - } - smtp := loadSMTPConfig(ctx) - if smtp.Host == "" || smtp.Username == "" { - return - } - port := smtp.Port - if port == "" { - port = "587" - } - cfg.URL = smtp.Host + ":" + port - cfg.Key = smtp.Username - cfg.Secret = smtp.Password -} - -func listPushChannels(ctx context.Context) ([]model.PushChannel, error) { - return repository.ListPushChannels(ctx) -} - -func createPushChannel(ctx context.Context, req CreateChannelRequest) (model.PushChannel, error) { - count, err := repository.CountPushChannelsByName(ctx, req.Name) - if err != nil { - return model.PushChannel{}, err - } - if count > 0 { - return model.PushChannel{}, errors.New("channel name already exists") - } - - channel := model.PushChannel{ - Name: req.Name, - Description: req.Description, - Type: req.Type, - Token: req.Token, - URL: req.URL, - Other: req.Other, - Enabled: req.Enabled, - } - if err := channel.Validate(); err != nil { - return model.PushChannel{}, err - } - if err := repository.CreatePushChannel(ctx, &channel); err != nil { - return model.PushChannel{}, err - } - return channel, nil -} - -func updatePushChannel(ctx context.Context, id uint64, req UpdateChannelRequest) (model.PushChannel, error) { - channel, err := repository.GetPushChannelByID(ctx, id) - if err != nil { - return model.PushChannel{}, err - } - - channel.Description = req.Description - channel.Type = req.Type - channel.Token = req.Token - channel.URL = req.URL - channel.Other = req.Other - channel.Enabled = req.Enabled - if err := channel.Validate(); err != nil { - return model.PushChannel{}, err - } - if err := repository.SavePushChannel(ctx, &channel); err != nil { - return model.PushChannel{}, err - } - return channel, nil -} - -func deletePushChannel(ctx context.Context, id uint64) error { - channel, err := repository.GetPushChannelByID(ctx, id) - if err != nil { - return err - } - return repository.DeletePushChannel(ctx, &channel) -} - -func loadChannelForTest(ctx context.Context, req TestChannelRequest) (string, string, string, string, error) { - if req.Name != "" { - channel, err := repository.GetPushChannelByName(ctx, req.Name) - if err != nil { - return "", "", "", "", errors.New("channel not found") - } - return channel.URL, channel.Token, channel.Other, channel.Type, nil - } - return req.URL, req.Token, req.Other, req.Type, nil -} - -func listActivePushEventsByTaskType(ctx context.Context, taskType string) ([]model.PushEvent, error) { - return repository.ListActivePushEventsByTaskType(ctx, taskType) -} - -func loadUserFromPayload(ctx context.Context, data map[string]any) any { - if u, exists := data["user"]; exists && u != nil { - return u - } - - if userID, ok := extractUserID(data); ok && userID > 0 { - if user, err := repository.GetUserByID(ctx, userID); err == nil { - return &user - } - } - - if username := extractUsername(data); username != "" { - if user, err := repository.GetUserByUsername(ctx, username); err == nil { - return &user - } - } - return nil -} - -func recordPushHistory(ctx context.Context, req SendPayload, status, errMsg string) error { - title := req.Body.Title - content := req.Body.Content - level := req.Body.Level - if title == "" { - title = "系统通知" - } - if level == "" { - level = defaultLevelInfo - } - - target := req.Target - if target == "" { - if req.Config.URL != "" { - target = req.Config.URL - const maxTargetLen = 50 - const truncatedLen = 47 - if len(target) > maxTargetLen { - target = target[:truncatedLen] + "..." - } - } else { - target = "default" - } - } - - history := model.PushHistory{ - EventKey: req.EventKey, - Channel: req.Config.Channel, - Target: target, - Title: title, - Content: content, - Level: level, - Status: status, - ErrorMsg: errMsg, - } - return repository.CreatePushHistory(ctx, &history) -} - -func resolveTargetUser(ctx context.Context, resolved string, _ string) (model.User, bool) { - found := false - var user model.User - - if id, err := strconv.ParseUint(resolved, 10, 64); err == nil { - if u, err := repository.GetUserByID(ctx, id); err == nil { - user = u - found = true - } - } - if !found { - if u, err := repository.GetUserByUsername(ctx, resolved); err == nil { - user = u - found = true - } - } - return user, found -} - -func resolveSystemTarget(ctx context.Context, resolved string, channel string) (string, bool) { - if resolved != "系统" && resolved != "system" && resolved != "0" { - return "", false - } - adminUser, err := repository.GetFirstAdminUser(ctx) - if err != nil { - return resolved, true - } - if channel == channelEmail && adminUser.Email != "" { - return adminUser.Email, true - } - if channel != channelEmail && adminUser.Username != "" { - return adminUser.Username, true - } - return resolved, true -} - -func resolveSMTPConfig(ctx context.Context, url, token, other string) (string, string, string) { - if url != "" && token != "" { - return url, token, other - } - smtp := loadSMTPConfig(ctx) - if smtp.Host == "" || smtp.Username == "" { - return url, token, other - } - port := smtp.Port - if port == "" { - port = "587" - } - if url == "" { - url = smtp.Host + ":" + port - } - if token == "" { - token = smtp.Username - } - if other == "" { - other = smtp.Password - } - return url, token, other -} - -func getSystemUser(ctx context.Context) *model.User { - user := repository.GetSystemUser(ctx) - return &user -} - -func getEventInfo(req CreateEventRequest) (string, string, []byte, error) { - if req.TaskType != "" { - meta := task.GetTaskMetaByAsynqTask(req.TaskType) - if meta == nil { - return "", "", nil, errors.New("unsupported task type") - } - eventKey := "task_completed:" + req.TaskType - eventName := "任务完成: " + meta.Name - defaultTemplate := NotificationMessage{ - Title: "任务完成: " + meta.Name, - Content: "异步任务 {{task_name}} (ID: {{task_id}}) 已完成。状态: {{task_status}},耗时: {{task_duration}} ms。", - Level: defaultLevelInfo, - } - defaultTemplateBytes, err := json.Marshal(defaultTemplate) - if err != nil { - return "", "", nil, err - } - return eventKey, eventName, defaultTemplateBytes, nil - } - - if req.EventKey == "" { - return "", "", nil, errors.New("either event_key or task_type must be provided") - } - - meta, found := findBuiltInEvent(req.EventKey) - if !found { - return "", "", nil, errors.New("unsupported built-in event key") - } - - defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate) - if err != nil { - return "", "", nil, err - } - return req.EventKey, meta.Name, defaultTemplateBytes, nil -} diff --git a/internal/apps/admin/push/push_test.go b/internal/apps/admin/push/push_test.go deleted file mode 100644 index 9174151f..00000000 --- a/internal/apps/admin/push/push_test.go +++ /dev/null @@ -1,760 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package push - -import ( - "bytes" - "context" - "encoding/json" - "net/http" - "net/http/httptest" - "strconv" - "sync" - "testing" - "time" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/shared/response" - "github.com/Rain-kl/Wavelet/internal/testhelper" - pkgpush "github.com/Rain-kl/Wavelet/pkg/push" - "github.com/alicebob/miniredis/v2" - "github.com/gin-gonic/gin" - "github.com/hibiken/asynq" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - "gorm.io/gorm" -) - -var adminLoginEvent = EventMetadata{ - Key: "admin_login", - Name: "管理员登录", - DefaultTemplate: NotificationMessage{ - Title: "管理员登录提醒", - Content: "管理员 {{user.username}} 于 {{time}} 从 IP {{ip}} 登录系统。", - Level: "INFO", - }, - Description: "当管理员成功登录系统时触发此通知", -} - -func init() { - RegisterBuiltInEvent(adminLoginEvent) -} - -// mockPusher mock implementation of pkgpush.Pusher -type mockPusher struct { - mu sync.Mutex - sentBody map[string]any - sentTgt string -} - -func (m *mockPusher) Send(ctx context.Context, cfg pkgpush.Config, target string, body map[string]any, template string, ext map[string]any) (string, error) { - m.mu.Lock() - defer m.mu.Unlock() - m.sentBody = body - m.sentTgt = target - return "", nil -} - -func (m *mockPusher) ValidateConfig(cfg pkgpush.Config) error { - return nil -} - -func setupPushTest(t *testing.T) (*gorm.DB, *miniredis.Miniredis, func()) { - dbConn, mr, cleanup := testhelper.SetupTestEnvironment(t) - - // AutoMigrate push tables in SQLite test environment - err := dbConn.AutoMigrate(&model.PushEvent{}, &model.PushHistory{}, &model.User{}, &model.PushChannel{}, &model.SystemConfig{}) - require.NoError(t, err) - - // 写入数据库系统默认用户 Seed 记录 - sysUser := &model.User{ - ID: 999, - Username: "system", - Nickname: "系统", - Password: "*", - IsActive: true, - } - err = dbConn.Create(sysUser).Error - require.NoError(t, err) - - // Initialize AsynqClient pointing to miniredis - task.AsynqClient = asynq.NewClient(asynq.RedisClientOpt{ - Addr: mr.Addr(), - }) - - // Register the task handler and metadata - task.RegisterHandler(SendNotificationTask, &PushHandler{}) - task.RegisterTaskMeta(SendNotificationMeta) - - return dbConn, mr, func() { - cleanup() - if task.AsynqClient != nil { - task.AsynqClient.Close() - task.AsynqClient = nil - } - } -} - -func setupTestRouter(authUser *model.User) *gin.Engine { - r := testhelper.NewTestGinEngine() - adminGroup := r.Group("/api/v1/admin/push") - - adminGroup.Use(func(c *gin.Context) { - if authUser != nil { - oauth.SetToContext(c, "user_obj", authUser) - } - c.Next() - }) - - adminGroup.GET("/events", ListEvents) - adminGroup.GET("/events/builtin", ListBuiltInEvents) - adminGroup.POST("/events", CreateEvent) - adminGroup.PUT("/events/:id", UpdateEvent) - adminGroup.DELETE("/events/:id", DeleteEvent) - adminGroup.POST("/events/:id/toggle", ToggleEvent) - adminGroup.GET("/histories", ListHistories) - adminGroup.POST("/test", TestPush) - - return r -} - -func TestSyncEvents(t *testing.T) { - dbConn, _, cleanup := setupPushTest(t) - defer cleanup() - - // 1. SyncEvents first time - err := SyncEvents(context.Background()) - require.NoError(t, err) - - // Verify event exists in DB - var event model.PushEvent - err = dbConn.Where("event_key = ?", "admin_login").First(&event).Error - require.NoError(t, err) - assert.Equal(t, "管理员登录", event.Name) - assert.False(t, event.Enabled) - - // Verify DefaultTemplate matches GORM template field - var defaultMsg NotificationMessage - err = json.Unmarshal([]byte(event.Template), &defaultMsg) - require.NoError(t, err) - assert.Equal(t, adminLoginEvent.DefaultTemplate.Title, defaultMsg.Title) - assert.Equal(t, adminLoginEvent.DefaultTemplate.Content, defaultMsg.Content) -} - -func TestEventTrigger(t *testing.T) { - dbConn, _, cleanup := setupPushTest(t) - defer cleanup() - - // SyncEvents - err := SyncEvents(context.Background()) - require.NoError(t, err) - - t.Run("trigger disabled event silently ignored", func(t *testing.T) { - body := map[string]any{ - "user": map[string]any{"username": "test_admin"}, - "ip": "127.0.0.1", - } - DefaultTrigger.Trigger(context.Background(), adminLoginEvent, body) - - // Sleep briefly since Trigger runs in goroutine - time.Sleep(50 * time.Millisecond) - - // Verify no tasks enqueued in TaskExecution GORM table - var count int64 - dbConn.Model(&model.TaskExecution{}).Count(&count) - assert.Equal(t, int64(0), count) - }) - - t.Run("trigger enabled event enqueues task", func(t *testing.T) { - // Create an enabled custom channel in GORM - customChan := &model.PushChannel{ - Name: "mock_channel", - Type: "custom", - URL: "https://webhook.site/trigger", - Other: `{"text": "$content"}`, - Enabled: true, - } - err = dbConn.Create(customChan).Error - require.NoError(t, err) - defer dbConn.Delete(customChan) - - // Enable the push event in DB using struct to trigger JSON serializer - var event model.PushEvent - err = dbConn.Where("event_key = ?", "admin_login").First(&event).Error - require.NoError(t, err) - - event.Enabled = true - event.Channels = []string{"mock_channel"} - event.Targets = []string{"admin_user"} - err = repository.SavePushEvent(context.Background(), &event) - require.NoError(t, err) - - // Trigger - body := map[string]any{ - "user": map[string]any{ - "username": "super_admin", - }, - "ip": "1.1.1.1", - "time": "2026-06-14 18:00:00", - } - DefaultTrigger.Trigger(context.Background(), adminLoginEvent, body) - - // Wait for goroutine execution - time.Sleep(50 * time.Millisecond) - - // Verify TaskExecution enqueued record - var execution model.TaskExecution - err = dbConn.Where("task_type = ?", SendNotificationTask).First(&execution).Error - require.NoError(t, err) - - // Verify enqueued payload structure - var payload SendPayload - err = json.Unmarshal([]byte(execution.Payload), &payload) - require.NoError(t, err) - assert.Equal(t, "admin_login", payload.EventKey) - assert.Equal(t, "custom", payload.Config.Channel) - assert.Equal(t, "https://webhook.site/trigger", payload.Config.URL) - assert.Equal(t, "admin_user", payload.Target) - assert.Equal(t, "管理员登录提醒", payload.Body.Title) - assert.Contains(t, payload.Body.Content, "super_admin") - assert.Contains(t, payload.Body.Content, "1.1.1.1") - }) - - t.Run("trigger without user injects virtual system user", func(t *testing.T) { - // Create an enabled custom channel in GORM - customChan := &model.PushChannel{ - Name: "mock_channel", - Type: "custom", - URL: "https://webhook.site/trigger", - Other: `{"text": "$content"}`, - Enabled: true, - } - err = dbConn.Create(customChan).Error - require.NoError(t, err) - defer dbConn.Delete(customChan) - - // Enable the push event in DB - var event model.PushEvent - err = dbConn.Where("event_key = ?", "admin_login").First(&event).Error - require.NoError(t, err) - - // 清理旧任务执行记录 - dbConn.Where("task_type = ?", SendNotificationTask).Delete(&model.TaskExecution{}) - - event.Enabled = true - event.Channels = []string{"mock_channel"} - event.Targets = []string{"user.username"} - err = repository.SavePushEvent(context.Background(), &event) - require.NoError(t, err) - - // Trigger with empty body (simulates cron scheduler triggering) - DefaultTrigger.Trigger(context.Background(), adminLoginEvent, nil) - - // Wait for goroutine execution - time.Sleep(50 * time.Millisecond) - - // Verify TaskExecution enqueued record - var execution model.TaskExecution - err = dbConn.Where("task_type = ?", SendNotificationTask).First(&execution).Error - require.NoError(t, err) - - var payload SendPayload - err = json.Unmarshal([]byte(execution.Payload), &payload) - require.NoError(t, err) - - // 检查 payload 是否将 target (user.username) 成功替换为 "system" - assert.Equal(t, "system", payload.Target) - // 检查 payload 中的 Content,应当被替换为 "system" 变量 - assert.Contains(t, payload.Body.Content, "system") - }) -} - -func TestPushHandler(t *testing.T) { - dbConn, _, cleanup := setupPushTest(t) - defer cleanup() - - mPusher := &mockPusher{} - pkgpush.Register("mock_channel", mPusher) - - handler := &PushHandler{} - - payload := SendPayload{ - EventKey: "admin_login", - Config: pkgpush.Config{ - Channel: "mock_channel", - URL: "http://mock-url", - }, - Target: "admin_user", - Body: NotificationMessage{ - Title: "Structured Alert", - Content: "Hello World", - Level: "WARNING", - Ext: map[string]any{"extra_val": 42}, - }, - } - payloadBytes, err := json.Marshal(payload) - require.NoError(t, err) - - t.Run("validate payload", func(t *testing.T) { - validated, valErr := handler.ValidatePayload(payloadBytes) - require.NoError(t, valErr) - assert.NotEmpty(t, validated) - }) - - t.Run("execute task successfully", func(t *testing.T) { - res, execErr := handler.Execute(context.Background(), payloadBytes) - require.NoError(t, execErr) - assert.Contains(t, res.Message, "推送成功") - - // Verify mock pusher received flattened variables - mPusher.mu.Lock() - assert.Equal(t, "admin_user", mPusher.sentTgt) - assert.Equal(t, "Structured Alert", mPusher.sentBody["title"]) - assert.Equal(t, "Hello World", mPusher.sentBody["content"]) - assert.Equal(t, "WARNING", mPusher.sentBody["level"]) - assert.Equal(t, float64(42), mPusher.sentBody["extra_val"]) // unmarshaled json numbers are float64 by default - mPusher.mu.Unlock() - - // Verify PushHistory recorded - var history model.PushHistory - err = dbConn.First(&history).Error - require.NoError(t, err) - assert.Equal(t, "admin_login", history.EventKey) - assert.Equal(t, "mock_channel", history.Channel) - assert.Equal(t, "success", history.Status) - assert.Equal(t, "Structured Alert", history.Title) - }) -} - -func TestPushRouters(t *testing.T) { - dbConn, _, cleanup := setupPushTest(t) - defer cleanup() - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - r := setupTestRouter(adminUser) - - // Sync events to populate db - err := SyncEvents(context.Background()) - require.NoError(t, err) - - t.Run("list events", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/push/events", nil) - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - var resp response.Any - err = json.Unmarshal(w.Body.Bytes(), &resp) - require.NoError(t, err) - - dataBytes, _ := json.Marshal(resp.Data) - var events []model.PushEvent - err = json.Unmarshal(dataBytes, &events) - require.NoError(t, err) - - assert.Len(t, events, 1) - assert.Equal(t, "admin_login", events[0].EventKey) - }) - - t.Run("toggle event status", func(t *testing.T) { - var event model.PushEvent - dbConn.First(&event) - - // 1. 未配置任何渠道时开启,应该被拒绝 - req, _ := http.NewRequest("POST", "/api/v1/admin/push/events/"+strconv.FormatUint(event.ID, 10)+"/toggle", nil) - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - assert.Equal(t, http.StatusBadRequest, w.Code) - - // 2. 为该事件关联渠道后,再切换开启,应当成功 - event.Channels = []string{"email"} - _ = repository.SavePushEvent(context.Background(), &event) - - req2, _ := http.NewRequest("POST", "/api/v1/admin/push/events/"+strconv.FormatUint(event.ID, 10)+"/toggle", nil) - w2 := httptest.NewRecorder() - r.ServeHTTP(w2, req2) - assert.Equal(t, http.StatusOK, w2.Code) - - var updated model.PushEvent - dbConn.First(&updated) - assert.True(t, updated.Enabled) - }) - - t.Run("update event", func(t *testing.T) { - var event model.PushEvent - dbConn.First(&event) - - updateReq := UpdateEventRequest{ - Channels: []string{"email"}, - Targets: []string{"user@test.com"}, - Template: `{"title": "Custom Login Alert", "content": "Alert", "level": "WARNING"}`, - Enabled: true, - } - bodyBytes, _ := json.Marshal(updateReq) - req, _ := http.NewRequest("PUT", "/api/v1/admin/push/events/"+strconv.FormatUint(event.ID, 10), bytes.NewBuffer(bodyBytes)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - - var updated model.PushEvent - dbConn.First(&updated) - assert.Equal(t, []string{"email"}, updated.Channels) - assert.Equal(t, []string{"user@test.com"}, updated.Targets) - assert.Contains(t, updated.Template, "Custom Login Alert") - }) - - t.Run("list push histories", func(t *testing.T) { - // Populate history record - hist := model.PushHistory{ - EventKey: "admin_login", - Channel: "email", - Target: "user@test.com", - Title: "Custom Login Alert", - Content: "Alert", - Level: "WARNING", - Status: "success", - } - dbConn.Create(&hist) - - req, _ := http.NewRequest("GET", "/api/v1/admin/push/histories?page=1&page_size=10", nil) - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - var resp response.Any - json.Unmarshal(w.Body.Bytes(), &resp) - - dataMap, ok := resp.Data.(map[string]any) - assert.True(t, ok) - assert.Equal(t, float64(1), dataMap["total"]) - }) - - t.Run("test push endpoint", func(t *testing.T) { - mPusher := &mockPusher{} - pkgpush.Register("test_channel", mPusher) - - testReq := TestPushRequest{ - Config: pkgpush.Config{ - Channel: "test_channel", - URL: "http://test-url", - }, - Target: "test_target", - } - bodyBytes, _ := json.Marshal(testReq) - req, _ := http.NewRequest("POST", "/api/v1/admin/push/test", bytes.NewBuffer(bodyBytes)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - }) - - t.Run("list built-in events", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/push/events/builtin", nil) - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - var resp response.Any - err := json.Unmarshal(w.Body.Bytes(), &resp) - require.NoError(t, err) - - builtins, ok := resp.Data.([]any) - assert.True(t, ok) - assert.NotEmpty(t, builtins) - }) - - t.Run("create and delete push event", func(t *testing.T) { - // Clean up any existing admin_login event first - dbConn.Where("event_key = ?", "admin_login").Delete(&model.PushEvent{}) - - // 1. Create event - createReq := CreateEventRequest{ - EventKey: "admin_login", - Channels: []string{"email"}, - Targets: []string{"admin@test.com"}, - Enabled: true, - } - bodyBytes, _ := json.Marshal(createReq) - req, _ := http.NewRequest("POST", "/api/v1/admin/push/events", bytes.NewBuffer(bodyBytes)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - - // Verify created in DB - var event model.PushEvent - err := dbConn.Where("event_key = ?", "admin_login").First(&event).Error - require.NoError(t, err) - assert.Equal(t, "admin_login", event.EventKey) - assert.Equal(t, "管理员登录", event.Name) - assert.True(t, event.Enabled) - - // 2. Try creating again (should fail) - w2 := httptest.NewRecorder() - req2, _ := http.NewRequest("POST", "/api/v1/admin/push/events", bytes.NewBuffer(bodyBytes)) - req2.Header.Set("Content-Type", "application/json") - r.ServeHTTP(w2, req2) - assert.Equal(t, http.StatusBadRequest, w2.Code) - - // 3. Delete event - w3 := httptest.NewRecorder() - req3, _ := http.NewRequest("DELETE", "/api/v1/admin/push/events/"+strconv.FormatUint(event.ID, 10), nil) - r.ServeHTTP(w3, req3) - assert.Equal(t, http.StatusOK, w3.Code) - - // Verify deleted from DB - var count int64 - dbConn.Model(&model.PushEvent{}).Where("event_key = ?", "admin_login").Count(&count) - assert.Equal(t, int64(0), count) - }) -} - -func TestResolveTarget(t *testing.T) { - dbConn, _, cleanup := setupPushTest(t) - defer cleanup() - - // 1. 创建测试用户与管理员用户 - testUser := &model.User{ - ID: 9999, - Username: "target_user", - Email: "target@test.com", - IsAdmin: false, - } - err := dbConn.Create(testUser).Error - require.NoError(t, err) - - adminUser := &model.User{ - ID: 8888, - Username: "admin_user", - Email: "admin@test.com", - IsAdmin: true, - } - err = dbConn.Create(adminUser).Error - require.NoError(t, err) - - flatBody := map[string]any{ - "user.id": float64(9999), // JSON 反序列化后一般是 float64 - "user.username": "target_user", - "user.email": "target@test.com", - } - - ctx := context.Background() - - t.Run("dynamic user.id resolved and converted for email channel", func(t *testing.T) { - res := resolveTarget(ctx, "user.id", flatBody, "email") - assert.Equal(t, "target@test.com", res) - }) - - t.Run("dynamic user.username resolved and converted for email channel", func(t *testing.T) { - res := resolveTarget(ctx, "user.username", flatBody, "email") - assert.Equal(t, "target@test.com", res) - }) - - t.Run("dynamic user.email resolved directly for email channel", func(t *testing.T) { - res := resolveTarget(ctx, "user.email", flatBody, "email") - assert.Equal(t, "target@test.com", res) - }) - - t.Run("fixed user id resolved and converted for email channel", func(t *testing.T) { - res := resolveTarget(ctx, "9999", flatBody, "email") - assert.Equal(t, "target@test.com", res) - }) - - t.Run("fixed username resolved and converted for email channel", func(t *testing.T) { - res := resolveTarget(ctx, "target_user", flatBody, "email") - assert.Equal(t, "target@test.com", res) - }) - - t.Run("fixed email address resolved directly for email channel", func(t *testing.T) { - res := resolveTarget(ctx, "fixed@example.com", flatBody, "email") - assert.Equal(t, "fixed@example.com", res) - }) - - t.Run("fixed username resolved for non-email channel", func(t *testing.T) { - res := resolveTarget(ctx, "target_user", flatBody, "lark") - assert.Equal(t, "target_user", res) - }) - - t.Run("non-exist user resolved as fallback", func(t *testing.T) { - res := resolveTarget(ctx, "non_exist_user", flatBody, "email") - assert.Equal(t, "non_exist_user", res) - }) - - t.Run("system target resolves to admin email for email channel", func(t *testing.T) { - res := resolveTarget(ctx, "系统", flatBody, "email") - assert.Equal(t, "admin@test.com", res) - - res2 := resolveTarget(ctx, "system", flatBody, "email") - assert.Equal(t, "admin@test.com", res2) - - res3 := resolveTarget(ctx, "0", flatBody, "email") - assert.Equal(t, "admin@test.com", res3) - }) - - t.Run("system target resolves to admin username for lark channel", func(t *testing.T) { - res := resolveTarget(ctx, "系统", flatBody, "lark") - assert.Equal(t, "admin_user", res) - }) -} - -func TestPushChannelAPI(t *testing.T) { - // 1. 模型校验测试 - t.Run("validate push channel model constraints", func(t *testing.T) { - // 校验名称合法性 - c1 := &model.PushChannel{Name: "invalid-name!", URL: "https://hook.com", Other: "{}"} - assert.Error(t, c1.Validate()) - - // 校验 URL 安全前缀 HTTPS - c2 := &model.PushChannel{Name: "custom_channel", URL: "http://insecure-hook.com", Other: "{}"} - assert.Error(t, c2.Validate()) - - // 校验 JSON 格式 - c3 := &model.PushChannel{Name: "custom_channel", URL: "https://hook.com", Other: "{invalid-json}"} - assert.Error(t, c3.Validate()) - - // 正确配置 - c4 := &model.PushChannel{Name: "custom_channel", URL: "https://hook.com", Other: "{\"content\":\"$content\"}"} - assert.NoError(t, c4.Validate()) - - // 飞书渠道校验:非 HTTPS 地址报错 - c5 := &model.PushChannel{Name: "lark_channel", Type: "lark", URL: "http://open.feishu.cn", Other: ""} - assert.Error(t, c5.Validate()) - - // 飞书正确配置 - c6 := &model.PushChannel{Name: "lark_channel", Type: "lark", URL: "https://open.feishu.cn", Other: ""} - assert.NoError(t, c6.Validate()) - - // Telegram 渠道校验 - cTelegramErr := &model.PushChannel{Name: "tg_channel", Type: "telegram", URL: "https://api.telegram.org", Token: "", Other: ""} - assert.Error(t, cTelegramErr.Validate()) - - cTelegramErr2 := &model.PushChannel{Name: "tg_channel", Type: "telegram", URL: "http://api.telegram.org", Token: "123:abc", Other: ""} - assert.Error(t, cTelegramErr2.Validate()) - - cTelegramOk := &model.PushChannel{Name: "tg_channel", Type: "telegram", URL: "", Token: "123:abc", Other: "-100123"} - assert.NoError(t, cTelegramOk.Validate()) - assert.Equal(t, "https://api.telegram.org", cTelegramOk.URL) - - // 邮件配置校验:允许空配置以复用系统全局设置 - c7 := &model.PushChannel{Name: "email_channel", Type: "email", URL: "", Token: "", Other: ""} - assert.NoError(t, c7.Validate()) - - // 邮件正确配置 (非 HTTPS 协议 URL 允许) - c8 := &model.PushChannel{Name: "email_channel", Type: "email", URL: "smtp.exmail.qq.com:465", Token: "user@example.com", Other: "authcode"} - assert.NoError(t, c8.Validate()) - }) - - // 2. HTTP CRUD & 触发鉴权测试 - dbConn, _, cleanup := setupPushTest(t) - defer cleanup() - - // 构建路由以进行 HTTP 模拟请求 - r := testhelper.NewTestGinEngine() - adminGroup := r.Group("/api/v1/admin") - { - adminGroup.GET("/push/channels", ListChannels) - adminGroup.POST("/push/channels", CreateChannel) - adminGroup.PUT("/push/channels/:id", UpdateChannel) - adminGroup.DELETE("/push/channels/:id", DeleteChannel) - adminGroup.POST("/push/channels/test", TestChannel) - } - - var createdID uint64 - - t.Run("admin create channel", func(t *testing.T) { - reqBody := CreateChannelRequest{ - Name: "my_custom_channel", - Description: "My custom channel webhook", - Type: "custom", - Token: "my_chan_token", - URL: "https://webhook.site/test", - Other: `{"title": "$title", "body": "$content"}`, - Enabled: true, - } - bodyBytes, _ := json.Marshal(reqBody) - req, _ := http.NewRequest("POST", "/api/v1/admin/push/channels", bytes.NewBuffer(bodyBytes)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - var resp response.Any - json.Unmarshal(w.Body.Bytes(), &resp) - - dataMap, ok := resp.Data.(map[string]any) - assert.True(t, ok) - assert.Equal(t, "my_custom_channel", dataMap["name"]) - createdID = uint64(dataMap["id"].(float64)) - }) - - t.Run("admin list channels", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/push/channels", nil) - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - var resp response.Any - json.Unmarshal(w.Body.Bytes(), &resp) - list, ok := resp.Data.([]any) - assert.True(t, ok) - assert.Len(t, list, 1) - }) - - t.Run("admin update channel", func(t *testing.T) { - updateReq := UpdateChannelRequest{ - Description: "Updated remark", - Type: "custom", - Token: "new_chan_token", - URL: "https://webhook.site/updated", - Other: `{"text": "$content"}`, - Enabled: true, - } - bodyBytes, _ := json.Marshal(updateReq) - req, _ := http.NewRequest("PUT", "/api/v1/admin/push/channels/"+strconv.FormatUint(createdID, 10), bytes.NewBuffer(bodyBytes)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - - var updated model.PushChannel - dbConn.First(&updated, createdID) - assert.Equal(t, "Updated remark", updated.Description) - assert.Equal(t, "new_chan_token", updated.Token) - assert.Equal(t, `{"text": "$content"}`, updated.Other) - }) - - t.Run("admin test channel endpoint", func(t *testing.T) { - testReq := TestChannelRequest{ - Name: "my_custom_channel", - Target: "test_target", - } - bodyBytes, _ := json.Marshal(testReq) - req, _ := http.NewRequest("POST", "/api/v1/admin/push/channels/test", bytes.NewBuffer(bodyBytes)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - }) - - t.Run("admin delete channel", func(t *testing.T) { - req, _ := http.NewRequest("DELETE", "/api/v1/admin/push/channels/"+strconv.FormatUint(createdID, 10), nil) - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - - var count int64 - dbConn.Model(&model.PushChannel{}).Where("id = ?", createdID).Count(&count) - assert.Equal(t, int64(0), count) - }) -} diff --git a/internal/apps/admin/push/routers.go b/internal/apps/admin/push/routers.go deleted file mode 100644 index 7f937c13..00000000 --- a/internal/apps/admin/push/routers.go +++ /dev/null @@ -1,292 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package push defines push notification HTTP routes. -package push - -import ( - "context" - "errors" - "fmt" - "net/http" - "strconv" - - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/pkg/push" - "github.com/gin-gonic/gin" - "gorm.io/gorm" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -// UpdateEventRequest 更新事件请求参数 -type UpdateEventRequest struct { - Channels []string `json:"channels"` - Targets []string `json:"targets"` - Template string `json:"template" binding:"required"` - Enabled bool `json:"enabled"` -} - -// TestPushRequest 测试推送通道请求参数 -type TestPushRequest struct { - Config push.Config `json:"config" binding:"required"` - Target string `json:"target"` -} - -// SyncEvents automatically registers/updates built-in events in the database. -func SyncEvents(ctx context.Context) error { - return syncBuiltInEvents(ctx) -} - -// ListEvents 获取通知事件列表 -// @Summary 获取所有通知事件 -// @Description 返回系统配置的通知事件列表,包括预置和自定义事件,需要管理员权限 -// @Tags admin-push -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=[]model.PushEvent} "通知事件列表" -// @Router /api/v1/admin/push/events [get] -func ListEvents(c *gin.Context) { - ctx := c.Request.Context() - events, err := listPushEvents(ctx) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(events)) -} - -// CreateEventRequest 创建事件请求参数 -type CreateEventRequest struct { - EventKey string `json:"event_key"` - TaskType string `json:"task_type"` // 关联的异步任务类型 - Channels []string `json:"channels"` - Targets []string `json:"targets"` - Template string `json:"template"` - Enabled bool `json:"enabled"` -} - -func findBuiltInEvent(key string) (EventMetadata, bool) { - for _, meta := range BuiltInEvents { - if meta.Key == key { - return meta, true - } - } - return EventMetadata{}, false -} - -// ListBuiltInEvents 获取内置通知事件列表 -// @Summary 获取所有内置通知事件 -// @Description 返回系统定义的所有内置通知事件元数据,供前端下拉框选择,需要管理员权限 -// @Tags admin-push -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=[]EventMetadata} "内置通知事件列表" -// @Router /api/v1/admin/push/events/builtin [get] -func ListBuiltInEvents(c *gin.Context) { - c.JSON(http.StatusOK, response.OK(BuiltInEvents)) -} - -// CreateEvent 创建通知事件 -// @Summary 创建通知事件 -// @Description 绑定系统内置事件或异步任务、推送渠道、接收目标并创建通知事件配置,需要管理员权限 -// @Tags admin-push -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param request body CreateEventRequest true "创建参数" -// @Success 200 {object} response.Any{data=model.PushEvent} "创建成功" -// @Router /api/v1/admin/push/events [post] -func CreateEvent(c *gin.Context) { - var req CreateEventRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - event, err := createPushEvent(c.Request.Context(), req) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(event)) -} - -// DeleteEvent 删除通知事件配置 -// @Summary 删除通知事件配置 -// @Description 删除数据库中的特定通知事件配置,需要管理员权限 -// @Tags admin-push -// @Produce json -// @Security SessionCookie -// @Param id path int true "事件 ID" -// @Success 200 {object} response.Any{data=string} "删除成功" -// @Router /api/v1/admin/push/events/{id} [delete] -func DeleteEvent(c *gin.Context) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, "invalid event id") - return - } - - if err := deletePushEvent(c.Request.Context(), id); err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - response.AbortNotFound(c, "notification event not found") - return - } - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OKNil()) -} - -// UpdateEvent 更新通知事件 -// @Summary 更新通知事件 -// @Description 更新已有通知事件的推送渠道、接收目标和内容模板,需要管理员权限 -// @Tags admin-push -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param id path int true "事件 ID" -// @Param request body push.UpdateEventRequest true "更新参数" -// @Success 200 {object} response.Any{data=string} "修改成功" -// @Router /api/v1/admin/push/events/{id} [put] -func UpdateEvent(c *gin.Context) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, "invalid event id") - return - } - - var req UpdateEventRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - if err := updatePushEvent(c.Request.Context(), id, req); err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - response.AbortNotFound(c, "notification event not found") - return - } - response.AbortBadRequest(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OKNil()) -} - -// ToggleEvent 快捷切换通知事件启用状态 -// @Summary 快捷切换通知事件启用状态 -// @Description 启用或禁用指定的通知事件 -// @Tags admin-push -// @Produce json -// @Security SessionCookie -// @Param id path int true "事件 ID" -// @Success 200 {object} response.Any{data=string} "切换成功" -// @Router /api/v1/admin/push/events/{id}/toggle [post] -func ToggleEvent(c *gin.Context) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, "invalid event id") - return - } - - enabled, err := togglePushEvent(c.Request.Context(), id) - if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - response.AbortNotFound(c, "notification event not found") - return - } - response.AbortBadRequest(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(enabled)) -} - -// pushHistoriesResponse 推送历史分页响应 -// -//nolint:unused -type pushHistoriesResponse struct { - Total int64 `json:"total"` - Results []model.PushHistory `json:"results"` -} - -// ListHistories 分页获取通知推送历史 -// @Summary 分页获取通知推送历史 -// @Description 返回分页的通知历史日志数据,需要管理员权限 -// @Tags admin-push -// @Produce json -// @Security SessionCookie -// @Param page query int false "当前页码" -// @Param page_size query int false "分页大小" -// @Param event_key query string false "过滤事件名称" -// @Param status query string false "过滤发送状态" -// @Success 200 {object} response.Any{data=pushHistoriesResponse} "推送历史列表" -// @Router /api/v1/admin/push/histories [get] -func ListHistories(c *gin.Context) { - page, _ := strconv.Atoi(c.DefaultQuery("page", "1")) - pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20")) - if page < 1 { - page = 1 - } - if pageSize < 1 { - pageSize = 20 - } - - total, results, err := listPushHistories(c.Request.Context(), repository.PushHistoryListFilter{ - EventKey: c.Query("event_key"), - Status: c.Query("status"), - Page: page, - PageSize: pageSize, - }) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - - c.JSON(http.StatusOK, response.OK(map[string]any{ - "total": total, - "results": results, - })) -} - -// TestPush 测试推送通道发送 -// @Summary 测试推送通道发送 -// @Description 接收临时通知渠道配置并在本地同步调用 Pusher.Send 发送测试消息 -// @Tags admin-push -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param request body push.TestPushRequest true "测试请求体" -// @Success 200 {object} response.Any{data=string} "测试成功" -// @Router /api/v1/admin/push/test [post] -func TestPush(c *gin.Context) { - var req TestPushRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - pusher, err := push.GetPusher(req.Config.Channel) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - if err := pusher.ValidateConfig(req.Config); err != nil { - response.AbortBadRequest(c, fmt.Sprintf("validation failed: %v", err)) - return - } - - applySMTPFallbackToPushConfig(c.Request.Context(), &req.Config) - - testBody := map[string]any{ - keyTitle: "测试通道推送", - keyContent: "当您收到这条消息,说明当前渠道连通性测试通过。", - keyLevel: defaultLevelInfo, - } - if _, err := pusher.Send(c.Request.Context(), req.Config, req.Target, testBody, "", nil); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OKNil()) -} diff --git a/internal/apps/admin/push/task_listener.go b/internal/apps/admin/push/task_listener.go deleted file mode 100644 index ec119710..00000000 --- a/internal/apps/admin/push/task_listener.go +++ /dev/null @@ -1,124 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package push - -import ( - "context" - "encoding/json" - "strconv" - "time" - - "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/pkg/logger" -) - -// RegisterTaskListeners subscribes push notification handlers to task completion events. -func RegisterTaskListeners() { - task.OnTaskCompleted(handleTaskCompleted) -} - -func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, result *task.TaskResult, execErr error) { - events, err := listActivePushEventsByTaskType(ctx, execution.TaskType) - if err != nil { - logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", execution.TaskType, err) - return - } - if len(events) == 0 { - return - } - - body := map[string]any{ - "task_id": execution.TaskID, - "task_name": execution.TaskName, - "task_type": execution.TaskType, - "task_status": string(execution.Status), - "task_duration": execution.Duration, - "time": time.Now().Format("2006-01-02 15:04:05"), - } - if execErr != nil { - body["task_error"] = execErr.Error() - } else { - body["task_error"] = "" - } - if result != nil { - body["task_result"] = result.Message - } else { - body["task_result"] = "" - } - - var payloadMap map[string]any - if execution.Payload != "" { - if err := json.Unmarshal([]byte(execution.Payload), &payloadMap); err == nil { - body["payload"] = payloadMap - extractUserFromMap(ctx, payloadMap, body) - } - } - if result != nil && result.Detail != "" { - var detailMap map[string]any - if err := json.Unmarshal([]byte(result.Detail), &detailMap); err == nil { - body["detail"] = detailMap - extractUserFromMap(ctx, detailMap, body) - } - } - - for _, event := range events { - meta := EventMetadata{ - Key: event.EventKey, - Name: event.Name, - Description: "异步任务执行完毕触发的自动通知", - } - DefaultTrigger.Trigger(ctx, meta, body) - } -} - -func extractUserFromMap(ctx context.Context, data map[string]any, body map[string]any) { - if u, exists := body["user"]; exists && u != nil { - return - } - if user := loadUserFromPayload(ctx, data); user != nil { - body["user"] = user - } -} - -func extractUserID(data map[string]any) (uint64, bool) { - for _, k := range []string{"user_id", "userId", "uid"} { - val, ok := data[k] - if !ok || val == nil { - continue - } - switch v := val.(type) { - case float64: - if v >= 0 { - return uint64(v), true - } - case int: - if v >= 0 { - return uint64(v), true - } - case int64: - if v >= 0 { - return uint64(v), true - } - case uint64: - return v, true - case string: - if id, err := strconv.ParseUint(v, 10, 64); err == nil { - return id, true - } - } - } - return 0, false -} - -func extractUsername(data map[string]any) string { - for _, k := range []string{"username", "user_name"} { - if val, ok := data[k]; ok && val != nil { - if s, ok := val.(string); ok && s != "" { - return s - } - } - } - return "" -} diff --git a/internal/apps/admin/push/tasks.go b/internal/apps/admin/push/tasks.go deleted file mode 100644 index de310e87..00000000 --- a/internal/apps/admin/push/tasks.go +++ /dev/null @@ -1,127 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package push defines push notification HTTP routes and background tasks. -package push - -import ( - "context" - "encoding/json" - "errors" - "fmt" - - "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/pkg/push" -) - -const ( - // SendNotificationTask 发送推送通知任务标识 - SendNotificationTask = "push:send" - // TaskTypeSendNotification 推送通知管理类型 - TaskTypeSendNotification = "send_notification" -) - -// SendNotificationMeta represents the task metadata. -var SendNotificationMeta = task.TaskMeta{ - Type: TaskTypeSendNotification, - AsynqTask: SendNotificationTask, - Name: "推送通知", - Description: "异步执行系统通知的多渠道派发与推送", - SupportsTime: false, - MaxRetry: task.DefaultMaxRetry, - Queue: task.QueueDefault, - Retryable: true, - Params: []task.TaskParam{ - { - Name: "event_key", - Label: "事件标识", - Type: "string", - Required: true, - Placeholder: "admin_login", - }, - { - Name: "target", - Label: "目标接收者", - Type: "string", - Required: false, - }, - }, -} - -// PushHandler 通知推送异步任务处理器 -// -//nolint:revive -type PushHandler struct{} - -// ValidatePayload 校验并标准化推送参数 -func (h *PushHandler) ValidatePayload(payload []byte) ([]byte, error) { - if len(payload) == 0 { - return nil, errors.New("payload is required") - } - - var req SendPayload - if err := json.Unmarshal(payload, &req); err != nil { - return nil, fmt.Errorf("invalid json format: %w", err) - } - - if req.Config.Channel == "" { - return nil, errors.New("channel type is required") - } - - return json.Marshal(req) -} - -// Execute 异步执行推送操作并记录推送历史审计 -func (h *PushHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) { - var req SendPayload - if err := json.Unmarshal(payload, &req); err != nil { - task.AppendLog(ctx, "解析推送参数失败: %v", err) - return nil, fmt.Errorf("parse payload failed: %w", err) - } - - task.AppendLog(ctx, "开始推送通知: 事件 = %s, 渠道 = %s, 接收目标 = %s", req.EventKey, req.Config.Channel, req.Target) - - pusher, err := push.GetPusher(req.Config.Channel) - if err != nil { - errWrap := fmt.Errorf("get pusher failed: %w", err) - task.AppendLog(ctx, "推送失败: %v", errWrap) - if task.IsFinalAttempt(ctx) { - h.recordHistory(ctx, req, "failed", errWrap.Error()) - } - return nil, errWrap - } - - // 执行真正的消息推送,扁平化为原始 json 格式 - flatBody := req.Body.Flatten() - upstreamResp, err := pusher.Send(ctx, req.Config, req.Target, flatBody, req.Template, nil) - - title := req.Body.Title - content := req.Body.Content - - if err != nil { - task.AppendLog(ctx, "消息推送失败 (标题: %s): %v", title, err) - if upstreamResp != "" { - task.AppendLog(ctx, "上游返回: %s", upstreamResp) - } - if task.IsFinalAttempt(ctx) { - h.recordHistory(ctx, req, "failed", err.Error()) - } - return nil, fmt.Errorf("pusher.Send failed: %w", err) - } - - task.AppendLog(ctx, "消息推送成功 (标题: %s, 内容摘要: %s)", title, content) - if upstreamResp != "" { - task.AppendLog(ctx, "上游返回: %s", upstreamResp) - } - h.recordHistory(ctx, req, "success", "") - - return &task.TaskResult{ - Message: fmt.Sprintf("推送成功: [%s] -> %s", req.Config.Channel, req.Target), - }, nil -} - -func (h *PushHandler) recordHistory(ctx context.Context, req SendPayload, status string, errMsg string) { - if dbErr := recordPushHistory(ctx, req, status, errMsg); dbErr != nil { - task.AppendLog(ctx, "写入推送历史审计记录失败: %v", dbErr) - } -} diff --git a/internal/apps/admin/status/log_database.go b/internal/apps/admin/status/log_database.go deleted file mode 100644 index 230ca476..00000000 --- a/internal/apps/admin/status/log_database.go +++ /dev/null @@ -1,102 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package status - -import ( - "context" - "errors" - "net/http" - - "github.com/Rain-kl/Wavelet/internal/infra/config" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/repository/logstore" - "github.com/Rain-kl/Wavelet/internal/shared/response" - "github.com/Rain-kl/Wavelet/pkg/logger" - "github.com/gin-gonic/gin" - "gorm.io/gorm" -) - -const ( - logDBNamePostgres = "postgres" - logDBNameSQLite = "sqlite" - logDBNameClickHouse = "clickhouse" - defaultLogRetentionDays = 30 -) - -// LogDatabaseStatus 日志库状态。 -type LogDatabaseStatus struct { - ActiveDatabase string `json:"active_database"` - Migration string `json:"migration"` - RetentionDays map[string]int `json:"retention_days"` - AvailableTargets []string `json:"available_targets"` -} - -// GetLogDatabaseStatus 返回当前日志库状态。 -// @Summary 获取日志数据库状态 -// @Description 返回当前日志主库、迁移状态、各库保留天数与合法迁移目标,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=status.LogDatabaseStatus} "获取成功" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/status/log-database [get] -func GetLogDatabaseStatus(c *gin.Context) { - ctx := c.Request.Context() - store, err := logstore.Active(ctx) - if err != nil { - logger.ErrorF(ctx, "获取日志存储实例失败: %v", err) - response.AbortInternal(c, "日志存储初始化失败") - return - } - activeDB, err := store.Status.ActiveDatabase(ctx) - if err != nil { - logger.ErrorF(ctx, "获取日志库状态失败: %v", err) - response.AbortInternal(c, "获取日志库状态失败") - return - } - migration := "idle" - if logstore.Migrating(ctx) { - migration = "migrating" - } - c.JSON(http.StatusOK, response.OK(LogDatabaseStatus{ - ActiveDatabase: activeDB, - Migration: migration, - RetentionDays: map[string]int{ - logDBNamePostgres: retentionOr(ctx, model.ConfigKeyLogRetentionDaysPostgres), - logDBNameSQLite: retentionOr(ctx, model.ConfigKeyLogRetentionDaysSQLite), - logDBNameClickHouse: retentionOr(ctx, model.ConfigKeyLogRetentionDaysClickHouse), - }, - AvailableTargets: availableLogTargets(activeDB), - })) -} - -func retentionOr(ctx context.Context, key string) int { - v, err := repository.GetIntByKey(ctx, key) - if err != nil { - if !errors.Is(err, gorm.ErrRecordNotFound) { - logger.ErrorF(ctx, "读取日志保留天数配置失败 key=%s: %v", key, err) - } - return defaultLogRetentionDays - } - if v < 1 { - return defaultLogRetentionDays - } - return v -} - -func availableLogTargets(active string) []string { - if active == logDBNameClickHouse { - if config.Config.Database.Enabled { - return []string{logDBNamePostgres} - } - return []string{logDBNameSQLite} - } - if config.Config.ClickHouse.Enabled { - return []string{logDBNameClickHouse} - } - return []string{} -} diff --git a/internal/apps/admin/status/routers.go b/internal/apps/admin/status/routers.go deleted file mode 100644 index 30a0c5c5..00000000 --- a/internal/apps/admin/status/routers.go +++ /dev/null @@ -1,355 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package status 提供系统状态查询接口 -package status - -import ( - "context" - "fmt" - "log" - "math" - "net/http" - "os" - "os/exec" - "runtime" - "time" - - "github.com/Rain-kl/Wavelet/internal/infra/config" - "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -// startTime 记录服务启动时间 -var startTime = time.Now() - -const ( - hoursInDay = 24 - minutesInHour = 60 - secondsInMinute = 60 - nanosPerSecond = 1e9 - binaryKB = 0 - binaryMB = 1 - binaryGB = 2 - valueThreshold = 10 // 格式化时区分整数显示的阈值 -) - -// SystemStatusResponse 系统状态响应结构体 -type SystemStatusResponse struct { - Uptime string `json:"uptime"` - NumGoroutine int `json:"num_goroutine"` - Alloc string `json:"alloc"` - TotalAlloc string `json:"total_alloc"` - Sys string `json:"sys"` - Lookups uint64 `json:"lookups"` - Mallocs uint64 `json:"mallocs"` - Frees uint64 `json:"frees"` - HeapAlloc string `json:"heap_alloc"` - HeapSys string `json:"heap_sys"` - HeapIdle string `json:"heap_idle"` - HeapInuse string `json:"heap_inuse"` - HeapReleased string `json:"heap_released"` - HeapObjects uint64 `json:"heap_objects"` - StackInuse string `json:"stack_inuse"` - StackSys string `json:"stack_sys"` - MSpanInuse string `json:"mspan_inuse"` - MSpanSys string `json:"mspan_sys"` - MCacheInuse string `json:"mcache_inuse"` - MCacheSys string `json:"mcache_sys"` - BuckHashSys string `json:"buck_hash_sys"` - GCSys string `json:"gc_sys"` - OtherSys string `json:"other_sys"` - NextGC string `json:"next_gc"` - LastGCTime string `json:"last_gc_time"` - PauseTotalNs string `json:"pause_total_ns"` - LastPause string `json:"last_pause"` - NumGC uint32 `json:"num_gc"` -} - -// formatBytes 格式化字节大小 -func formatBytes(bytes uint64) string { - const unit = 1024 - if bytes < unit { - return fmt.Sprintf("%d B", bytes) - } - div, exp := int64(unit), 0 - for n := bytes / unit; n >= unit; n /= unit { - div *= unit - exp++ - } - value := float64(bytes) / float64(div) - var suffix string - switch exp { - case binaryKB: - suffix = "KiB" - case binaryMB: - suffix = "MiB" - case binaryGB: - suffix = "GiB" - default: - suffix = "TiB" - } - - // 格式化规则: - // - 如果是整数(如 16, 73, 105, 986, 112): - // - 如果 >= 10,则格式化为 "%.0f" (e.g. "16 KiB") - // - 如果 < 10,则格式化为 "%.1f" (e.g. "9.0 KiB") - // - 如果不是整数(如 5.8, 9.1, 7.6, 4.8):格式化为 "%.1f" - if value == math.Trunc(value) { - if value >= valueThreshold { - return fmt.Sprintf("%.0f %s", value, suffix) - } - return fmt.Sprintf("%.1f %s", value, suffix) - } - return fmt.Sprintf("%.1f %s", value, suffix) -} - -// formatDuration 格式化时间持续时间 -func formatDuration(d time.Duration) string { - days := int(d.Hours()) / hoursInDay - hours := int(d.Hours()) % hoursInDay - minutes := int(d.Minutes()) % minutesInHour - seconds := int(d.Seconds()) % secondsInMinute - - var res string - if days > 0 { - res += fmt.Sprintf("%d天", days) - } - if hours > 0 { - res += fmt.Sprintf("%d小时", hours) - } - if minutes > 0 { - res += fmt.Sprintf("%d分钟", minutes) - } - if seconds > 0 || res == "" { - res += fmt.Sprintf("%d秒钟", seconds) - } - return res -} - -// GetSystemStatus 获取系统状态信息 -// @Summary 获取系统状态信息 -// @Description 获取后端服务运行状态、Goroutine、内存指标等详细统计数据,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=status.SystemStatusResponse} "获取成功" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Router /api/v1/admin/status [get] -func GetSystemStatus(c *gin.Context) { - var m runtime.MemStats - runtime.ReadMemStats(&m) - - uptime := formatDuration(time.Since(startTime)) - numGoroutine := runtime.NumGoroutine() - - var lastGCTime string - switch { - case m.LastGC > 0 && m.LastGC <= math.MaxInt64: - lastGCTime = formatDuration(time.Since(time.Unix(0, int64(m.LastGC)))) - case m.LastGC > 0: - lastGCTime = "未知" - default: - lastGCTime = "无" - } - - var lastPause string - if m.NumGC > 0 { - lastPause = fmt.Sprintf("%.3fs", float64(m.PauseNs[(m.NumGC-1)%256])/nanosPerSecond) - } else { - lastPause = "0.000s" - } - - res := SystemStatusResponse{ - Uptime: uptime, - NumGoroutine: numGoroutine, - Alloc: formatBytes(m.Alloc), - TotalAlloc: formatBytes(m.TotalAlloc), - Sys: formatBytes(m.Sys), - Lookups: m.Lookups, - Mallocs: m.Mallocs, - Frees: m.Frees, - HeapAlloc: formatBytes(m.HeapAlloc), - HeapSys: formatBytes(m.HeapSys), - HeapIdle: formatBytes(m.HeapIdle), - HeapInuse: formatBytes(m.HeapInuse), - HeapReleased: formatBytes(m.HeapReleased), - HeapObjects: m.HeapObjects, - StackInuse: formatBytes(m.StackInuse), - StackSys: formatBytes(m.StackSys), - MSpanInuse: formatBytes(m.MSpanInuse), - MSpanSys: formatBytes(m.MSpanSys), - MCacheInuse: formatBytes(m.MCacheInuse), - MCacheSys: formatBytes(m.MCacheSys), - BuckHashSys: formatBytes(m.BuckHashSys), - GCSys: formatBytes(m.GCSys), - OtherSys: formatBytes(m.OtherSys), - NextGC: formatBytes(m.NextGC), - LastGCTime: lastGCTime, - PauseTotalNs: fmt.Sprintf("%.1fs", float64(m.PauseTotalNs)/nanosPerSecond), - LastPause: lastPause, - NumGC: m.NumGC, - } - - c.JSON(http.StatusOK, response.OK(res)) -} - -// DatabaseInfoResponse 数据库信息响应结构体 -type DatabaseInfoResponse struct { - Type string `json:"type"` - Name string `json:"name"` - Version string `json:"version"` -} - -// getSQLiteInfo 返回 SQLite 数据库信息 -func getSQLiteInfo(ctx context.Context) DatabaseInfoResponse { - info := DatabaseInfoResponse{ - Type: "sqlite", - Name: config.Config.Database.SQLitePath, - Version: "SQLite", - } - if info.Name == "" { - info.Name = "./data/wavelet.db" - } - gormDB := db.DB(ctx) - if gormDB == nil { - return info - } - var ver string - if err := gormDB.Raw("SELECT sqlite_version()").Scan(&ver).Error; err == nil && ver != "" { - info.Version = "SQLite " + ver - } - return info -} - -// getPostgresInfo 返回 PostgreSQL 数据库信息 -func getPostgresInfo(ctx context.Context) DatabaseInfoResponse { - info := DatabaseInfoResponse{ - Type: "postgres", - Name: config.Config.Database.Database, - Version: "PostgreSQL", - } - gormDB := db.DB(ctx) - if gormDB == nil { - return info - } - var ver string - if err := gormDB.Raw("SELECT version()").Scan(&ver).Error; err == nil && ver != "" { - info.Version = ver - } - return info -} - -// GetDatabaseInfo 获取当前数据库类型及版本信息 -// @Summary 获取数据库信息 -// @Description 返回当前使用的数据库类型(sqlite/postgres)、名称/路径及版本字符串,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=status.DatabaseInfoResponse} "获取成功" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Router /api/v1/admin/db-info [get] -func GetDatabaseInfo(c *gin.Context) { - var info DatabaseInfoResponse - if !config.Config.Database.Enabled { - info = getSQLiteInfo(c.Request.Context()) - } else { - info = getPostgresInfo(c.Request.Context()) - } - c.JSON(http.StatusOK, response.OK(info)) -} - -// ExportDatabase 导出数据库 -// @Summary 导出数据库 -// @Description SQLite 时直接下载 .db 文件;PostgreSQL 时执行 pg_dump 并流式下载 .sql 文件,需要管理员权限 -// @Tags admin -// @Produce application/octet-stream -// @Security SessionCookie -// @Success 200 {file} binary "数据库文件" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "导出失败" -// @Router /api/v1/admin/db-export [get] -func ExportDatabase(c *gin.Context) { - if !config.Config.Database.Enabled { - exportSQLite(c) - } else { - exportPostgres(c) - } -} - -// exportSQLite 以 HTTP 附件方式下载 SQLite .db 文件 -func exportSQLite(c *gin.Context) { - path := config.Config.Database.SQLitePath - if path == "" { - path = "./data/wavelet.db" - } - - f, err := os.Open(path) //nolint:gosec // path is loaded from server startup configuration, not user input - if err != nil { - response.AbortInternal(c, "无法打开数据库文件: "+err.Error()) - return - } - defer func() { - if closeErr := f.Close(); closeErr != nil { - _ = closeErr - } - }() - - fi, err := f.Stat() - if err != nil { - response.AbortInternal(c, "无法读取数据库文件信息: "+err.Error()) - return - } - - c.Header("Content-Disposition", `attachment; filename="wavelet.db"`) - c.Header("Content-Type", "application/octet-stream") - c.Header("Content-Length", fmt.Sprintf("%d", fi.Size())) - c.Status(http.StatusOK) - http.ServeContent(c.Writer, c.Request, "wavelet.db", fi.ModTime(), f) -} - -// exportPostgres 执行 pg_dump 并将输出流式传输给客户端 -func exportPostgres(c *gin.Context) { - dbCfg := config.Config.Database - - // 检查 pg_dump 是否可用 - pgDumpPath, err := exec.LookPath("pg_dump") - if err != nil { - response.AbortInternal(c, "pg_dump 不可用,请确保服务器已安装 PostgreSQL 客户端工具") - return - } - - args := []string{ - "--no-password", - "-h", dbCfg.Host, - "-p", fmt.Sprintf("%d", dbCfg.Port), - "-U", dbCfg.Username, - dbCfg.Database, - } - - cmd := exec.CommandContext(c.Request.Context(), pgDumpPath, args...) //nolint:gosec // pgDumpPath is a looked up command path, args are from database configuration - if dbCfg.Password != "" { - cmd.Env = append(os.Environ(), "PGPASSWORD="+dbCfg.Password) - } else { - cmd.Env = os.Environ() - } - - fileName := fmt.Sprintf("wavelet_%s.sql", time.Now().Format("20060102_150405")) - c.Header("Content-Disposition", `attachment; filename="`+fileName+`"`) - c.Header("Content-Type", "application/octet-stream") - c.Status(http.StatusOK) - - cmd.Stdout = c.Writer - cmd.Stderr = nil // 忽略 stderr 以避免污染输出流 - - if err := cmd.Run(); err != nil { - // 响应头已发出,无法再写 JSON 错误,记录到服务器日志 - log.Printf("[db-export] pg_dump failed: %v\n", err) - } -} diff --git a/internal/apps/admin/system_config/errs.go b/internal/apps/admin/system_config/errs.go deleted file mode 100644 index 312e4372..00000000 --- a/internal/apps/admin/system_config/errs.go +++ /dev/null @@ -1,16 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package system_config 提供系统配置管理功能 -package system_config - -// 系统配置错误消息常量 -const ( - SystemConfigNotFound = "系统配置不存在" - ConfigKeyRequired = "配置键不能为空" - ConfigValueRequired = "配置值不能为空" - ConfigKeyExists = "配置键已存在" - protectedConfigKeyMessage = "该配置项由系统任务管理,禁止手动修改" - StorageDriverSwitchRequiresMigration = "存在存量文件,请通过存储迁移任务切换存储引擎" -) diff --git a/internal/apps/admin/system_config/logics.go b/internal/apps/admin/system_config/logics.go deleted file mode 100644 index bab02077..00000000 --- a/internal/apps/admin/system_config/logics.go +++ /dev/null @@ -1,137 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package system_config - -import ( - "context" - "encoding/json" - "errors" - "time" - - "github.com/Rain-kl/Wavelet/internal/infra/objectstore" - "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/pkg/logger" - "gorm.io/gorm" -) - -func isProtectedConfigKey(key string) bool { - return key == model.ConfigKeyLogDatabase || key == model.ConfigKeyLogDBMigration -} - -func createSystemConfig(ctx context.Context, req CreateSystemConfigRequest) error { - if isProtectedConfigKey(req.Key) { - return errors.New(protectedConfigKeyMessage) - } - exists, err := repository.SystemConfigExists(ctx, req.Key) - if err != nil { - return err - } - if exists { - return errors.New(ConfigKeyExists) - } - - config := model.SystemConfig{ - Key: req.Key, - Value: req.Value, - Type: req.Type, - Visibility: req.Visibility, - Description: req.Description, - } - if err := repository.CreateSystemConfig(ctx, &config); err != nil { - return err - } - - invalidateSystemConfigCaches(ctx, req.Key) - if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil { - logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err) - } - return nil -} - -func listSystemConfigs(ctx context.Context, configType string) ([]model.SystemConfig, error) { - return repository.ListAdminSystemConfigs(ctx, configType) -} - -func getSystemConfig(ctx context.Context, key string) (model.SystemConfig, error) { - return repository.GetAdminSystemConfigByKey(ctx, key) -} - -func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigRequest) error { - if isProtectedConfigKey(key) { - return errors.New(protectedConfigKeyMessage) - } - config, err := repository.GetAdminSystemConfigByKey(ctx, key) - if err != nil { - return err - } - - var originalDriver objectstore.Driver - if key == model.ConfigKeyStorageConfig { - var currentCfg objectstore.Config - if err := json.Unmarshal([]byte(config.Value), ¤tCfg); err == nil { - originalDriver = currentCfg.Driver - } - - validatedVal, err := validateAndMergeStorageConfig(ctx, req.Value, config.Value) - if err != nil { - return err - } - req.Value = validatedVal - } - - if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { - updates := map[string]any{ - "description": req.Description, - } - if req.Visibility != nil { - updates["visibility"] = *req.Visibility - config.Visibility = *req.Visibility - } - if key != model.ConfigKeySMTPPassword || req.Value != maskedConfigValue { - updates["value"] = req.Value - config.Value = req.Value - } - if err := tx.Model(&config).Updates(updates).Error; err != nil { - return err - } - resolveStorageMigrationTasksOnDirectDriverUpdate(ctx, tx, key, originalDriver, req.Value) - return nil - }); err != nil { - return err - } - - invalidateCachesAfterConfigUpdate(ctx, key) - return nil -} - -func resolveStorageMigrationTasksOnDirectDriverUpdate( - ctx context.Context, - tx *gorm.DB, - key string, - originalDriver objectstore.Driver, - newValue string, -) { - if key != model.ConfigKeyStorageConfig || originalDriver == "" { - return - } - - var newCfg objectstore.Config - if err := json.Unmarshal([]byte(newValue), &newCfg); err != nil { - return - } - if newCfg.Driver != originalDriver { - return - } - - if err := repository.MarkFailedTaskExecutionsSucceededTx( - tx, - "storage:migrate", - "存储配置直接更新,故障迁移任务自动标记为已解决", - time.Now(), - ); err != nil { - logger.ErrorF(ctx, "自动更新迁移任务状态失败: %v", err) - } -} diff --git a/internal/apps/admin/system_config/routers.go b/internal/apps/admin/system_config/routers.go deleted file mode 100644 index 2be447d1..00000000 --- a/internal/apps/admin/system_config/routers.go +++ /dev/null @@ -1,373 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package system_config - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "net/http" - "strings" - - "github.com/gin-gonic/gin" - "gorm.io/gorm" - - "github.com/Rain-kl/Wavelet/internal/apps/cap" - "github.com/Rain-kl/Wavelet/internal/apps/upload" - "github.com/Rain-kl/Wavelet/internal/infra/objectstore" - "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/shared/response" - "github.com/Rain-kl/Wavelet/pkg/logger" - mail "github.com/Rain-kl/Wavelet/pkg/mail" -) - -const maskedConfigValue = "******" - -// CreateSystemConfigRequest 创建系统配置请求 -type CreateSystemConfigRequest struct { - Key string `json:"key" binding:"required,max=64"` - Value string `json:"value" binding:"required"` - Type string `json:"type" binding:"required,oneof=system business"` - Visibility int `json:"visibility" binding:"oneof=0 1"` - Description string `json:"description" binding:"max=255"` -} - -// UpdateSystemConfigRequest 更新系统配置请求 -type UpdateSystemConfigRequest struct { - Value string `json:"value" binding:"required"` - Visibility *int `json:"visibility" binding:"omitempty,oneof=0 1"` - Description string `json:"description" binding:"max=255"` -} - -// CreateSystemConfig 创建系统配置 -// @Summary 创建系统配置 -// @Description 创建一条新的系统配置项,配置键不可重复,同时将新配置同步到 Redis,需要管理员权限 -// @Tags admin -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param request body system_config.CreateSystemConfigRequest true "创建请求参数" -// @Success 200 {object} response.Any{data=string} "创建成功" -// @Failure 400 {object} response.Any "参数错误或配置键已存在" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/system-configs [post] -func CreateSystemConfig(c *gin.Context) { - var req CreateSystemConfigRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - if isProtectedConfigKey(req.Key) { - response.AbortBadRequest(c, protectedConfigKeyMessage) - return - } - - if err := createSystemConfig(c.Request.Context(), req); err != nil { - if err.Error() == ConfigKeyExists { - response.AbortBadRequest(c, ConfigKeyExists) - return - } - response.AbortInternal(c, err.Error()) - return - } - - c.JSON(http.StatusOK, response.OKNil()) -} - -// ListSystemConfigs 获取系统配置列表 -// @Summary 获取系统配置列表 -// @Description 返回所有系统配置列表,支持按配置类型(system/business)过滤,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Param type query string false "配置类型(system/business)" -// @Success 200 {object} response.Any{data=[]model.SystemConfig} "系统配置列表" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/system-configs [get] -func ListSystemConfigs(c *gin.Context) { - configs, err := listSystemConfigs(c.Request.Context(), c.Query("type")) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - - for i := range configs { - configs[i].Value = maskSensitiveConfig(configs[i].Key, configs[i].Value) - } - - c.JSON(http.StatusOK, response.OK(configs)) -} - -// GetSystemConfig 获取单个系统配置 -// @Summary 获取单个系统配置 -// @Description 根据配置键获取对应的系统配置详情,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Param key path string true "配置键" -// @Success 200 {object} response.Any{data=model.SystemConfig} "系统配置详情" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 404 {object} response.Any "配置不存在" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/system-configs/{key} [get] -func GetSystemConfig(c *gin.Context) { - config, err := getSystemConfig(c.Request.Context(), c.Param("key")) - if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - response.AbortNotFound(c, SystemConfigNotFound) - } else { - response.AbortInternal(c, err.Error()) - } - return - } - - config.Value = maskSensitiveConfig(config.Key, config.Value) - - c.JSON(http.StatusOK, response.OK(config)) -} - -// UpdateSystemConfig 更新系统配置 -// @Summary 更新系统配置 -// @Description 根据配置键更新对应的配置内容,同时将更新同步到 Redis,需要管理员权限 -// @Tags admin -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param key path string true "配置键" -// @Param request body system_config.UpdateSystemConfigRequest true "更新请求参数" -// @Success 200 {object} response.Any{data=string} "更新成功" -// @Failure 400 {object} response.Any "参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 404 {object} response.Any "配置不存在" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/system-configs/{key} [put] -func UpdateSystemConfig(c *gin.Context) { - var req UpdateSystemConfigRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - key := c.Param("key") - if isProtectedConfigKey(key) { - response.AbortBadRequest(c, protectedConfigKeyMessage) - return - } - if err := updateSystemConfig(c.Request.Context(), key, req); err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - response.AbortNotFound(c, SystemConfigNotFound) - return - } - if isStorageConfigValidationError(err) { - response.AbortBadRequest(c, err.Error()) - return - } - response.AbortInternal(c, err.Error()) - return - } - - c.JSON(http.StatusOK, response.OKNil()) -} - -func invalidateSystemConfigCaches(ctx context.Context, key string) { - if err := repository.InvalidateSystemConfigCache(ctx, key); err != nil { - logger.WarnF(ctx, "清理系统配置缓存失败: %v", err) - } - if cap.IsRuntimeConfigKey(key) { - cap.InvalidateRuntimeSettings() - } -} - -func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) { - invalidateSystemConfigCaches(ctx, key) - - if key == model.ConfigKeyStorageConfig { - upload.ResetAccessCaches() - upload.PublishAccessCacheInvalidation(ctx) - objectstore.ResetCache() - objectstore.PublishCacheInvalidation(ctx) - } - if key == model.ConfigKeyFileAccessWhitelist { - upload.ResetAccessCaches() - upload.PublishAccessCacheInvalidation(ctx) - } - - if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil { - logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err) - } -} - -// TestSMTPRequest 测试 SMTP 配置请求 -type TestSMTPRequest struct { - SMTPHost string `json:"smtp_host" binding:"required,max=255"` - SMTPPort int `json:"smtp_port" binding:"required"` - SMTPUsername string `json:"smtp_username" binding:"required,max=255"` - SMTPPassword string `json:"smtp_password" binding:"required,max=255"` - To string `json:"to" binding:"required,email"` -} - -// TestSMTPResponse 测试 SMTP 配置响应 -type TestSMTPResponse struct { - Success bool `json:"success"` - Log string `json:"log"` - Error string `json:"error"` -} - -// TestSMTP 测试 SMTP 邮件发送 -// @Summary 测试 SMTP 邮件发送 -// @Description 使用传入的配置进行 SMTP 邮件发送测试,支持使用 ****** 占位符使用保存的数据库密码 -// @Tags admin -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param request body system_config.TestSMTPRequest true "测试请求参数" -// @Success 200 {object} response.Any{data=system_config.TestSMTPResponse} "测试执行完毕" -// @Failure 400 {object} response.Any "参数错误" -// @Router /api/v1/admin/system-configs/smtp/test [post] -func TestSMTP(c *gin.Context) { - var req TestSMTPRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - password := req.SMTPPassword - if password == maskedConfigValue { - if sc, err := repository.GetSystemConfigByKey(c.Request.Context(), model.ConfigKeySMTPPassword); err == nil { - password = sc.Value - } - } - - cfg := mail.Config{ - Host: req.SMTPHost, - Port: req.SMTPPort, - Username: req.SMTPUsername, - Password: password, - } - - subject := "Wavelet SMTP Test Mail" - body := `

SMTP Mail Connection Test

-

If you received this message, your SMTP configuration is correct and mail sending is working properly.

-

Sent from Wavelet.

` - - logs, err := mail.SendMailWithLog(c.Request.Context(), cfg, req.To, subject, body) - resp := TestSMTPResponse{ - Success: err == nil, - Log: logs, - } - if err != nil { - resp.Error = err.Error() - } - - c.JSON(http.StatusOK, response.OK(resp)) -} - -func isStorageConfigValidationError(err error) bool { - msg := err.Error() - return msg == StorageDriverSwitchRequiresMigration || - strings.HasPrefix(msg, "解析") || - strings.HasPrefix(msg, "验证") || - strings.HasPrefix(msg, "初始化测试") || - strings.HasPrefix(msg, "存储连通性") || - strings.HasPrefix(msg, "序列化") || - strings.HasPrefix(msg, "检查存量文件") -} - -func maskSensitiveConfig(key, value string) string { - if value == "" { - return value - } - switch key { - case model.ConfigKeySMTPPassword: - return maskedConfigValue - case model.ConfigKeyStorageConfig: - var cfg objectstore.Config - if err := json.Unmarshal([]byte(value), &cfg); err == nil { - masked := objectstore.MaskSecrets(cfg) - if val, err := json.Marshal(masked); err == nil { - return string(val) - } - } - } - return value -} - -// validateAndMergeStorageConfig parses, merges unmasked secrets, validates parameter values, -// and tests connectivity of the new storage configuration. -func validateAndMergeStorageConfig(ctx context.Context, value string, currentConfig string) (string, error) { - var currentCfg objectstore.Config - if err := json.Unmarshal([]byte(currentConfig), ¤tCfg); err != nil { - return "", fmt.Errorf("解析当前存储配置失败: %w", err) - } - - var newCfg objectstore.Config - if err := json.Unmarshal([]byte(value), &newCfg); err != nil { - return "", fmt.Errorf("解析目标存储配置失败: %w", err) - } - - // 合并被掩码屏蔽的敏感信息,获取完整的真实配置 - targetCfg := objectstore.MergeMaskedSecrets(newCfg, currentCfg) - if err := validateMergedStorageConfig(ctx, currentCfg, newCfg, targetCfg); err != nil { - return "", err - } - - // 序列化为最终保存的真实明文配置,防止保存屏蔽的 ****** 字符 - unmaskedVal, err := json.Marshal(targetCfg) - if err != nil { - return "", fmt.Errorf("序列化存储配置失败: %w", err) - } - - return string(unmaskedVal), nil -} - -func validateMergedStorageConfig(ctx context.Context, currentCfg, newCfg, targetCfg objectstore.Config) error { - if newCfg.Driver != "" && newCfg.Driver != currentCfg.Driver { - var uploadCount int64 - if err := db.DB(ctx).Model(&model.Upload{}). - Where("status != ?", model.UploadStatusDeleted). - Count(&uploadCount).Error; err != nil { - return fmt.Errorf("检查存量文件失败: %w", err) - } - if uploadCount > 0 { - return errors.New(StorageDriverSwitchRequiresMigration) - } - if err := validateDriverConfig(targetCfg, newCfg.Driver); err != nil { - return fmt.Errorf("验证目标存储配置参数失败: %w", err) - } - pendingCfg := targetCfg - pendingCfg.Driver = newCfg.Driver - return testStorageBackend(ctx, pendingCfg, newCfg.Driver) - } - - if err := objectstore.ValidateConfig(targetCfg); err != nil { - return fmt.Errorf("验证存储配置参数失败: %w", err) - } - return testStorageBackend(ctx, targetCfg, targetCfg.Driver) -} - -func validateDriverConfig(cfg objectstore.Config, driver objectstore.Driver) error { - cfg.Driver = driver - return objectstore.ValidateConfig(cfg) -} - -func testStorageBackend(ctx context.Context, cfg objectstore.Config, driver objectstore.Driver) error { - testBackend, err := objectstore.NewBackend(ctx, cfg, driver) - if err != nil { - return fmt.Errorf("初始化测试存储实例失败: %w", err) - } - if err := testBackend.Test(ctx); err != nil { - return fmt.Errorf("存储连通性测试失败: %w", err) - } - return nil -} diff --git a/internal/apps/admin/system_config/routers_test.go b/internal/apps/admin/system_config/routers_test.go deleted file mode 100644 index befef2b5..00000000 --- a/internal/apps/admin/system_config/routers_test.go +++ /dev/null @@ -1,550 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package system_config - -import ( - "bufio" - "bytes" - "context" - "encoding/json" - "net" - "net/http" - "net/http/httptest" - "net/textproto" - "strings" - "testing" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/infra/objectstore" - "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -const expectedDefaultConfigsCount = 35 - -func setupTestRouter(authUser *model.User) *gin.Engine { - r := testhelper.NewTestGinEngine() - adminGroup := r.Group("/api/v1/admin") - - // Mock authentication middleware - adminGroup.Use(func(c *gin.Context) { - if authUser != nil { - oauth.SetToContext(c, oauth.UserObjKey, authUser) - } - c.Next() - }) - - adminGroup.POST("/system-configs", CreateSystemConfig) - adminGroup.GET("/system-configs", ListSystemConfigs) - - systemConfigRouter := adminGroup.Group("/system-configs/:key") - { - systemConfigRouter.GET("", GetSystemConfig) - systemConfigRouter.PUT("", UpdateSystemConfig) - } - - return r -} - -func TestCreateSystemConfig(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("create successfully", func(t *testing.T) { - payload := CreateSystemConfigRequest{ - Key: "custom_key", - Value: "custom_value", - Type: "system", - Visibility: model.ConfigVisibilityVisible, - Description: "desc", - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("POST", "/api/v1/admin/system-configs", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - // Verify database - var cfg model.SystemConfig - err := dbConn.Where("key = ?", "custom_key").First(&cfg).Error - if err != nil { - t.Fatalf("failed to find system config in DB: %v", err) - } - - // Verify caches are invalidated after create and repopulate on read - _, err = db.Redis.HGet( - context.Background(), - db.PrefixedKey(repository.SystemConfigRedisHashKey), - "custom_key", - ).Result() - if err == nil { - t.Fatal("expected redis cache miss immediately after create") - } - - loaded, err := repository.GetSystemConfigByKey(context.Background(), "custom_key") - if err != nil { - t.Fatalf("GetSystemConfigByKey(custom_key) error = %v", err) - } - if loaded.Value != "custom_value" { - t.Errorf("GetSystemConfigByKey(custom_key).Value = %q, want %q", loaded.Value, "custom_value") - } - if loaded.Visibility != model.ConfigVisibilityVisible { - t.Errorf("GetSystemConfigByKey(custom_key).Visibility = %d, want %d", loaded.Visibility, model.ConfigVisibilityVisible) - } - }) - - t.Run("create duplicate key error", func(t *testing.T) { - // Key "custom_key" already exists from previous test - payload := CreateSystemConfigRequest{ - Key: "custom_key", - Value: "another_value", - Type: "system", - Description: "desc", - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("POST", "/api/v1/admin/system-configs", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusBadRequest { - t.Errorf("expected 400 Bad Request on duplicate key, got %d", w.Code) - } - }) -} - -func TestListSystemConfigs(t *testing.T) { - _, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("list all seeded configurations", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/system-configs", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK, got %d", w.Code) - } - - var resp response.Any - _ = json.Unmarshal(w.Body.Bytes(), &resp) - - dataBytes, _ := json.Marshal(resp.Data) - var configs []model.SystemConfig - _ = json.Unmarshal(dataBytes, &configs) - - // Defaults seed configurations - if len(configs) != expectedDefaultConfigsCount { - t.Errorf("expected %d default configs, got %d", expectedDefaultConfigsCount, len(configs)) - } - }) - - t.Run("filter by type business", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/system-configs?type=business", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - var resp response.Any - _ = json.Unmarshal(w.Body.Bytes(), &resp) - - dataBytes, _ := json.Marshal(resp.Data) - var configs []model.SystemConfig - _ = json.Unmarshal(dataBytes, &configs) - - if len(configs) != 4 { - t.Errorf("expected 4 business configs, got %d: %v", len(configs), configs) - } - keys := make(map[string]struct{}, len(configs)) - for _, cfg := range configs { - keys[cfg.Key] = struct{}{} - } - if _, ok := keys[model.ConfigKeyMaxAPIKeysPerUser]; !ok { - t.Errorf("missing business config %s", model.ConfigKeyMaxAPIKeysPerUser) - } - }) -} - -func TestGetSystemConfig(t *testing.T) { - _, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("get existing configuration", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/system-configs/"+model.ConfigKeySiteName, nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Errorf("expected 200 OK, got %d", w.Code) - } - - var resp response.Any - _ = json.Unmarshal(w.Body.Bytes(), &resp) - - dataBytes, _ := json.Marshal(resp.Data) - var cfg model.SystemConfig - _ = json.Unmarshal(dataBytes, &cfg) - - if cfg.Value != "Wavelet" { - t.Errorf("expected 'Wavelet', got '%s'", cfg.Value) - } - }) - - t.Run("get non-existent config", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/system-configs/non_existent_key", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusNotFound { - t.Errorf("expected 404 Not Found, got %d", w.Code) - } - }) -} - -func TestUpdateSystemConfig(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("update successfully", func(t *testing.T) { - hidden := model.ConfigVisibilityHidden - payload := UpdateSystemConfigRequest{ - Value: "Super Site Name", - Visibility: &hidden, - Description: "Updated Description", - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("PUT", "/api/v1/admin/system-configs/"+model.ConfigKeySiteName, bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - // Verify database - var cfg model.SystemConfig - dbConn.Where("key = ?", model.ConfigKeySiteName).First(&cfg) - if cfg.Value != "Super Site Name" || cfg.Description != "Updated Description" || cfg.Visibility != model.ConfigVisibilityHidden { - t.Errorf("database values not updated: %+v", cfg) - } - - // Verify caches are invalidated after update and repopulate on read - _, err := db.Redis.HGet( - context.Background(), - db.PrefixedKey(repository.SystemConfigRedisHashKey), - model.ConfigKeySiteName, - ).Result() - if err == nil { - t.Fatal("expected redis cache miss immediately after update") - } - - loaded, err := repository.GetSystemConfigByKey(context.Background(), model.ConfigKeySiteName) - if err != nil { - t.Fatalf("GetSystemConfigByKey(site_name) error = %v", err) - } - if loaded.Value != "Super Site Name" { - t.Errorf("GetSystemConfigByKey(site_name).Value = %q, want %q", loaded.Value, "Super Site Name") - } - if loaded.Visibility != model.ConfigVisibilityHidden { - t.Errorf("GetSystemConfigByKey(site_name).Visibility = %d, want %d", loaded.Visibility, model.ConfigVisibilityHidden) - } - }) - - t.Run("update non-existent config", func(t *testing.T) { - payload := UpdateSystemConfigRequest{ - Value: "New Value", - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("PUT", "/api/v1/admin/system-configs/invalid_key", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusNotFound { - t.Errorf("expected 404 Not Found, got %d", w.Code) - } - }) -} - -func TestTestSMTP(t *testing.T) { - _, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - r := setupTestRouter(adminUser) - r.POST("/api/v1/admin/system-configs/smtp/test", TestSMTP) - - // Start a mock SMTP server - l, err := net.Listen("tcp", "127.0.0.1:0") - if err != nil { - t.Fatalf("failed to start mock smtp server: %v", err) - } - defer func() { _ = l.Close() }() - - port := l.Addr().(*net.TCPAddr).Port - - go func() { - conn, err := l.Accept() - if err != nil { - return - } - defer func() { _ = conn.Close() }() - - writer := bufio.NewWriter(conn) - reader := bufio.NewReader(conn) - tp := textproto.NewReader(reader) - - // 220 Ready - _, _ = writer.WriteString("220 mock.smtp.com SMTP Ready\r\n") - _ = writer.Flush() - - // Read HELO/EHLO - _, _ = tp.ReadLine() - _, _ = writer.WriteString("250-mock.smtp.com\r\n250 AUTH PLAIN\r\n") - _ = writer.Flush() - - // Read AUTH PLAIN - _, _ = tp.ReadLine() - _, _ = writer.WriteString("235 Authentication successful\r\n") - _ = writer.Flush() - - // Read MAIL FROM - _, _ = tp.ReadLine() - _, _ = writer.WriteString("250 OK\r\n") - _ = writer.Flush() - - // Read RCPT TO - _, _ = tp.ReadLine() - _, _ = writer.WriteString("250 OK\r\n") - _ = writer.Flush() - - // Read DATA - _, _ = tp.ReadLine() - _, _ = writer.WriteString("354 Start mail input\r\n") - _ = writer.Flush() - - // Read body lines until dot - for { - line, err := tp.ReadLine() - if err != nil || line == "." { - break - } - } - _, _ = writer.WriteString("250 OK\r\n") - _ = writer.Flush() - - // Read QUIT - _, _ = tp.ReadLine() - _, _ = writer.WriteString("221 Bye\r\n") - _ = writer.Flush() - }() - - payload := TestSMTPRequest{ - SMTPHost: "127.0.0.1", - SMTPPort: port, - SMTPUsername: "sender@example.com", - SMTPPassword: "password", - To: "recipient@example.com", - } - - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("POST", "/api/v1/admin/system-configs/smtp/test", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - var resp response.Any - json.Unmarshal(w.Body.Bytes(), &resp) - - dataBytes, _ := json.Marshal(resp.Data) - var testResp TestSMTPResponse - json.Unmarshal(dataBytes, &testResp) - - if !testResp.Success { - t.Errorf("expected test success, got failed: %s. Log: %s", testResp.Error, testResp.Log) - } -} - -func TestUpdateStorageConfigValidation(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("update storage config successfully", func(t *testing.T) { - tempDir := t.TempDir() - cfg := objectstore.DefaultConfig() - cfg.Local.Root = tempDir - - cfgBytes, _ := json.Marshal(cfg) - payload := UpdateSystemConfigRequest{ - Value: string(cfgBytes), - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("PUT", "/api/v1/admin/system-configs/storage_config", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - // Verify database - var dbCfg model.SystemConfig - dbConn.Where("key = ?", "storage_config").First(&dbCfg) - var savedCfg objectstore.Config - _ = json.Unmarshal([]byte(dbCfg.Value), &savedCfg) - if savedCfg.Local.Root != tempDir { - t.Errorf("expected local root to be updated to %s, got %s", tempDir, savedCfg.Local.Root) - } - }) - - t.Run("update storage config failed connectivity check", func(t *testing.T) { - cfg := objectstore.DefaultConfig() - cfg.Driver = objectstore.DriverS3 - cfg.S3.Bucket = "non-existent-bucket" - cfg.S3.Endpoint = "http://127.0.0.1:9999" // Will fail connectivity check - - cfgBytes, _ := json.Marshal(cfg) - payload := UpdateSystemConfigRequest{ - Value: string(cfgBytes), - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("PUT", "/api/v1/admin/system-configs/storage_config", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusBadRequest { - t.Errorf("expected 400 Bad Request, got %d. Body: %s", w.Code, w.Body.String()) - } - }) - - t.Run("reject driver switch when uploads exist", func(t *testing.T) { - upload := model.Upload{ - ID: 88001, - UserID: 1, - FileName: "keep.txt", - FilePath: "uploads/keep.txt", - FileSize: 4, - MimeType: "text/plain", - Extension: "txt", - Type: "attachment", - Status: model.UploadStatusUsed, - } - if err := dbConn.Create(&upload).Error; err != nil { - t.Fatalf("seed upload failed: %v", err) - } - - tempDir := t.TempDir() - cfg := objectstore.DefaultConfig() - cfg.Driver = objectstore.DriverS3 - cfg.S3.Endpoint = "http://127.0.0.1:19998" - cfg.S3.Region = "us-east-1" - cfg.S3.Bucket = "wavelet" - cfg.S3.AccessKeyID = "test" - cfg.S3.SecretAccessKey = "test" - cfg.Local.Root = tempDir - - cfgBytes, _ := json.Marshal(cfg) - payload := UpdateSystemConfigRequest{Value: string(cfgBytes)} - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("PUT", "/api/v1/admin/system-configs/storage_config", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusBadRequest { - t.Fatalf("expected 400 Bad Request, got %d. Body: %s", w.Code, w.Body.String()) - } - if !strings.Contains(w.Body.String(), StorageDriverSwitchRequiresMigration) { - t.Fatalf("expected migration-required error, got: %s", w.Body.String()) - } - }) - - t.Run("switch to local while active s3 is unreachable", func(t *testing.T) { - if err := dbConn.Where("1 = 1").Delete(&model.Upload{}).Error; err != nil { - t.Fatalf("clear uploads failed: %v", err) - } - - activeCfg := objectstore.DefaultConfig() - activeCfg.Driver = objectstore.DriverS3 - activeCfg.S3.Endpoint = "http://127.0.0.1:9999" - activeCfg.S3.Region = "us-east-1" - activeCfg.S3.Bucket = "wavelet" - activeCfg.S3.AccessKeyID = "test" - activeCfg.S3.SecretAccessKey = "test" - activeBytes, _ := json.Marshal(activeCfg) - seedCfg := model.SystemConfig{ - Key: "storage_config", - Value: string(activeBytes), - Type: "system", - } - if err := dbConn.Where("key = ?", "storage_config"). - Assign(map[string]any{"value": seedCfg.Value, "type": seedCfg.Type}). - FirstOrCreate(&seedCfg).Error; err != nil { - t.Fatalf("seed active storage config failed: %v", err) - } - - tempDir := t.TempDir() - stagedCfg := activeCfg - stagedCfg.Driver = objectstore.DriverLocal - stagedCfg.Local.Root = tempDir - - cfgBytes, _ := json.Marshal(stagedCfg) - payload := UpdateSystemConfigRequest{Value: string(cfgBytes)} - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("PUT", "/api/v1/admin/system-configs/storage_config", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - var dbCfg model.SystemConfig - if err := dbConn.Where("key = ?", "storage_config").First(&dbCfg).Error; err != nil { - t.Fatalf("load saved storage config failed: %v", err) - } - var savedCfg objectstore.Config - if err := json.Unmarshal([]byte(dbCfg.Value), &savedCfg); err != nil { - t.Fatalf("parse saved storage config failed: %v", err) - } - if savedCfg.Driver != objectstore.DriverLocal { - t.Fatalf("active driver = %q, want %q after save", savedCfg.Driver, objectstore.DriverLocal) - } - if savedCfg.Local.Root != tempDir { - t.Fatalf("staged local root = %q, want %q", savedCfg.Local.Root, tempDir) - } - }) -} diff --git a/internal/apps/admin/task/errs.go b/internal/apps/admin/task/errs.go deleted file mode 100644 index e45edf79..00000000 --- a/internal/apps/admin/task/errs.go +++ /dev/null @@ -1,23 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package task 提供任务管理接口 -package task - -// 任务管理相关错误消息 -const ( - InvalidTaskType = "无效的任务类型" - InvalidTimeRange = "无效的时间范围" - TaskDispatchFailed = "任务下发失败" - UserIDRequired = "用户ID必填" - TaskNotFound = "任务执行记录不存在" - TaskNotRetryable = "该任务不支持重试" - TaskNotFailed = "只有失败的任务才能重试" - TaskMaxRetryExceeded = "已达到最大重试次数" - TaskRetryFailed = "任务重试失败" - InvalidCronExpression = "无效的 Cron 表达式" - ScheduleNotFound = "定时任务不存在" - ScheduleSaveFailed = "保存定时任务失败" - ScheduleDeleteFailed = "删除定时任务失败" -) diff --git a/internal/apps/admin/task/routers.go b/internal/apps/admin/task/routers.go deleted file mode 100644 index 6703b839..00000000 --- a/internal/apps/admin/task/routers.go +++ /dev/null @@ -1,417 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package task - -import ( - "fmt" - "net/http" - "strconv" - "strings" - "time" - - "github.com/Rain-kl/Wavelet/internal/apps/admin" - "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/internal/infra/task/scheduler" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/pkg/logger" - "github.com/gin-gonic/gin" - "github.com/robfig/cron/v3" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -// ListTaskTypes 获取支持的任务类型列表 -// @Summary 获取支持的任务类型 -// @Description 返回系统支持的所有可调度任务类型列表,包括任务名称、描述、是否支持时间范围等元数据,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=[]task.TaskMeta} "任务类型列表" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Router /api/v1/admin/tasks/types [get] -func ListTaskTypes(c *gin.Context) { - c.JSON(http.StatusOK, response.OK(task.GetDispatchableTasks())) -} - -// DispatchTaskRequest 下发任务请求 -type DispatchTaskRequest struct { - TaskType string `json:"task_type" binding:"required"` - StartTime *time.Time `json:"start_time"` - EndTime *time.Time `json:"end_time"` - UserID *uint64 `json:"user_id"` - Payload string `json:"payload"` -} - -// DispatchTask 下发任务 -// @Summary 下发异步任务 -// @Description 手动触发指定类型的异步任务,支持指定时间范围和用户,需要管理员权限 -// @Tags admin -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param request body DispatchTaskRequest true "任务请求参数" -// @Success 200 {object} response.Any{data=string} "任务已入队" -// @Failure 400 {object} response.Any "任务类型不存在或参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "任务入队失败" -// @Router /api/v1/admin/tasks/dispatch [post] -func DispatchTask(c *gin.Context) { - var req DispatchTaskRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - meta := task.GetTaskMeta(req.TaskType) - if meta == nil { - response.AbortBadRequest(c, InvalidTaskType) - return - } - - var payloadBytes []byte - if strings.TrimSpace(req.Payload) != "" { - payloadBytes = []byte(req.Payload) - } - - validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - taskID, err := task.DispatchTask(c.Request.Context(), req.TaskType, validated, "manual") - if err != nil { - response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskDispatchFailed, err)) - return - } - - c.JSON(http.StatusOK, response.OK(taskID)) -} - -// ListTaskExecutions 查询任务执行记录列表 -// @Summary 查询任务执行记录 -// @Description 分页查询任务执行记录,支持按状态和任务类型筛选,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Param status query string false "状态筛选 (pending/running/succeeded/failed)" -// @Param task_type query string false "任务类型筛选" -// @Param page query int false "页码" default(1) -// @Param page_size query int false "每页条数" default(20) -// @Success 200 {object} response.Any{data=object} "任务执行记录列表" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Router /api/v1/admin/tasks/executions [get] -func ListTaskExecutions(c *gin.Context) { - var req model.ListTaskExecutionsRequest - if err := c.ShouldBindQuery(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - if req.TaskType != "" { - if meta := task.GetTaskMeta(req.TaskType); meta != nil { - req.TaskType = meta.AsynqTask - } - } - - executions, total, err := repository.ListTaskExecutions(c.Request.Context(), req) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - - c.JSON(http.StatusOK, response.OK(gin.H{ - "items": executions, - "total": total, - "page": req.Page, - "page_size": req.PageSize, - })) -} - -// GetTaskExecution 查询单条任务执行详情 -// @Summary 查询任务执行详情 -// @Description 根据 ID 查询任务执行记录详情,包含完整执行日志,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Param id path int true "任务执行记录 ID" -// @Success 200 {object} response.Any{data=model.TaskExecution} "任务执行详情" -// @Failure 400 {object} response.Any "参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 404 {object} response.Any "记录不存在" -// @Router /api/v1/admin/tasks/executions/{id} [get] -func GetTaskExecution(c *gin.Context) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, admin.InvalidTaskExecutionID) - return - } - - execution, err := repository.GetTaskExecutionByID(c.Request.Context(), id) - if err != nil { - response.AbortNotFound(c, TaskNotFound) - return - } - - c.JSON(http.StatusOK, response.OK(execution)) -} - -// RetryTask 重试失败的任务 -// @Summary 重试失败任务 -// @Description 重新下发一条失败的任务,创建新的执行记录,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Param id path int true "任务执行记录 ID" -// @Success 200 {object} response.Any{data=string} "新任务的 TaskID" -// @Failure 400 {object} response.Any "任务不支持重试或参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 404 {object} response.Any "记录不存在" -// @Failure 500 {object} response.Any "重试失败" -// @Router /api/v1/admin/tasks/executions/{id}/retry [post] -func RetryTask(c *gin.Context) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, admin.InvalidTaskExecutionID) - return - } - - newTaskID, err := task.RetryTask(c.Request.Context(), id) - if err != nil { - errMsg := err.Error() - switch { - case strings.Contains(errMsg, "不存在"): - response.AbortNotFound(c, errMsg) - case strings.Contains(errMsg, "只有失败的任务") || strings.Contains(errMsg, "不支持重试") || strings.Contains(errMsg, "已达到最大重试"): - response.AbortBadRequest(c, errMsg) - default: - response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskRetryFailed, err)) - } - return - } - - c.JSON(http.StatusOK, response.OK(newTaskID)) -} - -// ListSchedules 获取定时任务列表 -// @Summary 获取定时任务列表 -// @Description 返回系统所有的定时任务配置列表,包括名称、关联的异步任务类型、Cron 表达式和启用状态,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=[]model.Schedule} "定时任务列表" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Router /api/v1/admin/tasks/schedules [get] -func ListSchedules(c *gin.Context) { - schedules, err := repository.ListSchedules(c.Request.Context()) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(schedules)) -} - -// CreateScheduleRequest 创建定时任务请求 -type CreateScheduleRequest struct { - Name string `json:"name" binding:"required"` - TaskType string `json:"task_type" binding:"required"` - Cron string `json:"cron" binding:"required"` - Payload string `json:"payload"` - IsActive *bool `json:"is_active" binding:"required"` -} - -// CreateSchedule 创建定时任务 -// @Summary 创建定时任务 -// @Description 新增一个动态定时任务配置,关联已有的异步任务,配置 Cron 表达式和执行参数,并触发调度器热加载,需要管理员权限 -// @Tags admin -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param request body CreateScheduleRequest true "创建定时任务请求参数" -// @Success 200 {object} response.Any{data=model.Schedule} "创建成功的定时任务信息" -// @Failure 400 {object} response.Any "Cron 表达式无效、异步任务类型不存在或参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "保存定时任务失败" -// @Router /api/v1/admin/tasks/schedules [post] -func CreateSchedule(c *gin.Context) { - var req CreateScheduleRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - // 校验 Cron 表达式 - if _, err := cron.ParseStandard(req.Cron); err != nil { - response.AbortBadRequest(c, InvalidCronExpression) - return - } - - // 校验关联的异步任务类型 - meta := task.GetTaskMeta(req.TaskType) - if meta == nil { - response.AbortBadRequest(c, InvalidTaskType) - return - } - - // 校验并规范化 Payload - var payloadBytes []byte - if strings.TrimSpace(req.Payload) != "" { - payloadBytes = []byte(req.Payload) - } - validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - schedule := &model.Schedule{ - Name: req.Name, - TaskType: req.TaskType, - Cron: req.Cron, - Payload: string(validated), - IsActive: *req.IsActive, - } - - if err := repository.CreateSchedule(c.Request.Context(), schedule); err != nil { - response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err)) - return - } - - // 触发调度服务重载 - if err := scheduler.ReloadScheduler(); err != nil { - logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err) - } - - c.JSON(http.StatusOK, response.OK(schedule)) -} - -// UpdateScheduleRequest 修改定时任务请求 -type UpdateScheduleRequest struct { - Name string `json:"name" binding:"required"` - TaskType string `json:"task_type" binding:"required"` - Cron string `json:"cron" binding:"required"` - Payload string `json:"payload"` - IsActive *bool `json:"is_active" binding:"required"` -} - -// UpdateSchedule 修改定时任务 -// @Summary 修改定时任务 -// @Description 修改一个定时任务的配置(名称、Cron 表达式、异步任务参数和是否启用等),并触发调度器热加载,需要管理员权限 -// @Tags admin -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param id path int true "定时任务 ID" -// @Param request body UpdateScheduleRequest true "修改定时任务请求参数" -// @Success 200 {object} response.Any{data=model.Schedule} "修改后的定时任务信息" -// @Failure 400 {object} response.Any "Cron 表达式无效、参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 404 {object} response.Any "定时任务不存在" -// @Failure 500 {object} response.Any "修改定时任务失败" -// @Router /api/v1/admin/tasks/schedules/{id} [put] -func UpdateSchedule(c *gin.Context) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, "无效的定时任务ID") - return - } - - var req UpdateScheduleRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - // 检查定时任务是否存在 - schedule, err := repository.GetScheduleByID(c.Request.Context(), id) - if err != nil { - response.AbortNotFound(c, ScheduleNotFound) - return - } - - // 校验 Cron 表达式 - if _, err := cron.ParseStandard(req.Cron); err != nil { - response.AbortBadRequest(c, InvalidCronExpression) - return - } - - // 校验关联的异步任务类型 - meta := task.GetTaskMeta(req.TaskType) - if meta == nil { - response.AbortBadRequest(c, InvalidTaskType) - return - } - - // 校验并规范化 Payload - var payloadBytes []byte - if strings.TrimSpace(req.Payload) != "" { - payloadBytes = []byte(req.Payload) - } - validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - schedule.Name = req.Name - schedule.TaskType = req.TaskType - schedule.Cron = req.Cron - schedule.Payload = string(validated) - schedule.IsActive = *req.IsActive - - if err := repository.UpdateSchedule(c.Request.Context(), schedule); err != nil { - response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err)) - return - } - - // 触发调度服务重载 - if err := scheduler.ReloadScheduler(); err != nil { - logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err) - } - - c.JSON(http.StatusOK, response.OK(schedule)) -} - -// DeleteSchedule 删除定时任务 -// @Summary 删除定时任务 -// @Description 删除指定的定时任务配置,并触发调度器热加载,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Param id path int true "定时任务 ID" -// @Success 200 {object} response.Any{data=string} "删除结果" -// @Failure 400 {object} response.Any "参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "删除定时任务失败" -// @Router /api/v1/admin/tasks/schedules/{id} [delete] -func DeleteSchedule(c *gin.Context) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, "无效的定时任务ID") - return - } - - if err := repository.DeleteSchedule(c.Request.Context(), id); err != nil { - response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err)) - return - } - - // 触发调度服务重载 - if err := scheduler.ReloadScheduler(); err != nil { - logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err) - } - - c.JSON(http.StatusOK, response.OKNil()) -} diff --git a/internal/apps/admin/task/routers_test.go b/internal/apps/admin/task/routers_test.go deleted file mode 100644 index 6dd07aa6..00000000 --- a/internal/apps/admin/task/routers_test.go +++ /dev/null @@ -1,496 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package task - -import ( - "bytes" - "context" - "encoding/json" - "fmt" - "net/http" - "net/http/httptest" - "testing" - "time" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - uploadtask "github.com/Rain-kl/Wavelet/internal/apps/upload/task" - "github.com/Rain-kl/Wavelet/internal/apps/user" - "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/platform/bootstrap" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/gin-gonic/gin" - "github.com/hibiken/asynq" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -func setupTaskTestEnvironment(t *testing.T) func() { - _, mr, cleanup := testhelper.SetupTestEnvironment(t) - bootstrap.RegisterTasks() - task.AsynqClient = asynq.NewClient(asynq.RedisClientOpt{ - Addr: mr.Addr(), - }) - return func() { - if task.AsynqClient != nil { - _ = task.AsynqClient.Close() - task.AsynqClient = nil - } - cleanup() - } -} - -func setupTestRouter(authUser *model.User) *gin.Engine { - r := testhelper.NewTestGinEngine() - adminGroup := r.Group("/api/v1/admin") - - // Mock authentication middleware - adminGroup.Use(func(c *gin.Context) { - if authUser != nil { - oauth.SetToContext(c, oauth.UserObjKey, authUser) - } - c.Next() - }) - - adminGroup.GET("/tasks/types", ListTaskTypes) - adminGroup.POST("/tasks/dispatch", DispatchTask) - adminGroup.GET("/tasks/executions", ListTaskExecutions) - adminGroup.GET("/tasks/executions/:id", GetTaskExecution) - adminGroup.POST("/tasks/executions/:id/retry", RetryTask) - return r -} - -func TestListTaskTypes(t *testing.T) { - cleanup := setupTaskTestEnvironment(t) - defer cleanup() - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/types", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Errorf("expected 200 OK, got %d", w.Code) - } - - var resp response.Any - _ = json.Unmarshal(w.Body.Bytes(), &resp) - - dataBytes, _ := json.Marshal(resp.Data) - var taskMetas []task.TaskMeta - _ = json.Unmarshal(dataBytes, &taskMetas) - - if len(taskMetas) == 0 { - t.Error("expected at least one dispatchable task type") - } - - foundCleanup := false - foundWarmImageCache := false - for _, m := range taskMetas { - if m.Type == uploadtask.TaskTypeSystemCleanup { - foundCleanup = true - } - if m.Type == uploadtask.TaskTypeWarmImageCache { - foundWarmImageCache = true - } - } - if !foundCleanup { - t.Errorf("expected task type %s to be listed", uploadtask.TaskTypeSystemCleanup) - } - if !foundWarmImageCache { - t.Errorf("expected task type %s to be listed", uploadtask.TaskTypeWarmImageCache) - } -} - -func TestDispatchTask(t *testing.T) { - cleanup := setupTaskTestEnvironment(t) - defer cleanup() - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("dispatch valid task successfully", func(t *testing.T) { - payload := DispatchTaskRequest{ - TaskType: uploadtask.TaskTypeSystemCleanup, - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code, "Body: %s", w.Body.String()) - - var resp response.Any - _ = json.Unmarshal(w.Body.Bytes(), &resp) - assert.Empty(t, resp.ErrorMsg) - assert.NotNil(t, resp.Data) - - // 返回的 data 应该是 taskID - taskID, ok := resp.Data.(string) - assert.True(t, ok) - assert.NotEmpty(t, taskID) - }) - - t.Run("dispatch send_email task successfully with valid payload", func(t *testing.T) { - payload := DispatchTaskRequest{ - TaskType: user.TaskTypeSendEmail, - Payload: `{"to":"receiver@example.com","subject":"Test Subject","body":"Test Body"}`, - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code, "Body: %s", w.Body.String()) - - var resp response.Any - json.Unmarshal(w.Body.Bytes(), &resp) - assert.Empty(t, resp.ErrorMsg) - assert.NotNil(t, resp.Data) - }) - - t.Run("dispatch send_email task failure with invalid payload json", func(t *testing.T) { - payload := DispatchTaskRequest{ - TaskType: user.TaskTypeSendEmail, - Payload: `{"to":`, - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusBadRequest, w.Code) - var resp response.Any - json.Unmarshal(w.Body.Bytes(), &resp) - assert.Contains(t, resp.ErrorMsg, "无效的 JSON 格式") - }) - - t.Run("dispatch send_email task failure with missing fields", func(t *testing.T) { - payload := DispatchTaskRequest{ - TaskType: user.TaskTypeSendEmail, - Payload: `{"to":"","subject":"Test","body":"Test"}`, - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusBadRequest, w.Code) - var resp response.Any - json.Unmarshal(w.Body.Bytes(), &resp) - assert.Contains(t, resp.ErrorMsg, "不能为空") - }) - - t.Run("dispatch invalid task type failure", func(t *testing.T) { - payload := DispatchTaskRequest{ - TaskType: "invalid_task_type", - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusBadRequest, w.Code) - - var resp response.Any - json.Unmarshal(w.Body.Bytes(), &resp) - assert.Equal(t, InvalidTaskType, resp.ErrorMsg) - }) - - t.Run("dispatch with empty body failure", func(t *testing.T) { - req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer([]byte("{}"))) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusBadRequest, w.Code) - }) -} - -func TestListTaskExecutions(t *testing.T) { - cleanup := setupTaskTestEnvironment(t) - defer cleanup() - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - ctx := context.Background() - - // 准备测试数据 - now := time.Now() - records := []*model.TaskExecution{ - {TaskID: "exec_001", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusSucceeded, TriggeredBy: "manual", Duration: 1500, Result: "清理完成", StartedAt: &now, FinishedAt: &now}, - {TaskID: "exec_002", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusFailed, TriggeredBy: "system", ErrorMessage: "连接超时", StartedAt: &now, FinishedAt: &now}, - {TaskID: "exec_003", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusPending, TriggeredBy: "manual", Retryable: true, MaxRetry: 3}, - } - for _, r := range records { - err := repository.CreateTaskExecution(ctx, r) - require.NoError(t, err) - } - - t.Run("list all executions", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - - var resp response.Any - json.Unmarshal(w.Body.Bytes(), &resp) - - dataBytes, _ := json.Marshal(resp.Data) - var data map[string]interface{} - json.Unmarshal(dataBytes, &data) - - assert.Equal(t, float64(3), data["total"]) - }) - - t.Run("filter by status", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions?status=failed", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - - var resp response.Any - json.Unmarshal(w.Body.Bytes(), &resp) - - dataBytes, _ := json.Marshal(resp.Data) - var data map[string]interface{} - json.Unmarshal(dataBytes, &data) - - assert.Equal(t, float64(1), data["total"]) - }) - - t.Run("filter by task_type (asynq task name)", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions?task_type=system:cleanup", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - - var resp response.Any - json.Unmarshal(w.Body.Bytes(), &resp) - - dataBytes, _ := json.Marshal(resp.Data) - var data map[string]interface{} - json.Unmarshal(dataBytes, &data) - - assert.Equal(t, float64(3), data["total"]) - }) - - t.Run("filter by task_type (management task type)", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions?task_type=system_cleanup", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - - var resp response.Any - json.Unmarshal(w.Body.Bytes(), &resp) - - dataBytes, _ := json.Marshal(resp.Data) - var data map[string]interface{} - json.Unmarshal(dataBytes, &data) - - assert.Equal(t, float64(3), data["total"]) - }) - - t.Run("pagination", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions?page=1&page_size=2", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - - var resp response.Any - json.Unmarshal(w.Body.Bytes(), &resp) - - dataBytes, _ := json.Marshal(resp.Data) - var data map[string]interface{} - json.Unmarshal(dataBytes, &data) - - assert.Equal(t, float64(3), data["total"]) - }) -} - -func TestGetTaskExecution(t *testing.T) { - cleanup := setupTaskTestEnvironment(t) - defer cleanup() - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - ctx := context.Background() - - // 创建测试记录 - execution := &model.TaskExecution{ - TaskID: "detail_001", - TaskType: "system:cleanup", - TaskName: "系统垃圾清理", - Status: model.TaskExecutionStatusSucceeded, - Log: "[10:00:01] 开始扫描\n[10:00:02] 找到 50 个文件\n[10:00:03] 清理完成", - Result: "共清理 50 个文件", - Duration: 2000, - Retryable: true, - MaxRetry: 3, - TriggeredBy: "manual", - } - err := repository.CreateTaskExecution(ctx, execution) - require.NoError(t, err) - - t.Run("get existing execution", func(t *testing.T) { - url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d", execution.ID) - req, _ := http.NewRequest("GET", url, nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - - var resp response.Any - json.Unmarshal(w.Body.Bytes(), &resp) - - dataBytes, _ := json.Marshal(resp.Data) - var detail model.TaskExecution - json.Unmarshal(dataBytes, &detail) - - assert.Equal(t, "detail_001", detail.TaskID) - assert.Equal(t, model.TaskExecutionStatusSucceeded, detail.Status) - assert.Contains(t, detail.Log, "开始扫描") - assert.Contains(t, detail.Log, "清理完成") - assert.Equal(t, int64(2000), detail.Duration) - }) - - t.Run("get non-existent execution", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions/99999999", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusNotFound, w.Code) - }) - - t.Run("invalid ID format", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions/invalid", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusBadRequest, w.Code) - }) -} - -func TestRetryTask(t *testing.T) { - cleanup := setupTaskTestEnvironment(t) - defer cleanup() - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - ctx := context.Background() - - t.Run("retry failed task successfully", func(t *testing.T) { - now := time.Now() - execution := &model.TaskExecution{ - TaskID: "retry_api_001", - TaskType: "system:cleanup", - TaskName: "系统垃圾清理", - Status: model.TaskExecutionStatusFailed, - ErrorMessage: "S3 连接超时", - Retryable: true, - MaxRetry: 3, - RetryCount: 0, - TriggeredBy: "manual", - StartedAt: &now, - FinishedAt: &now, - } - err := repository.CreateTaskExecution(ctx, execution) - require.NoError(t, err) - - url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID) - req, _ := http.NewRequest("POST", url, nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - - var resp response.Any - json.Unmarshal(w.Body.Bytes(), &resp) - assert.Empty(t, resp.ErrorMsg) - assert.NotNil(t, resp.Data) - - // 验证新记录 - newTaskID, ok := resp.Data.(string) - assert.True(t, ok) - assert.NotEmpty(t, newTaskID) - - newExecution, err := repository.GetTaskExecutionByTaskID(ctx, newTaskID) - require.NoError(t, err) - assert.Equal(t, 1, newExecution.RetryCount) - assert.Equal(t, "retry", newExecution.TriggeredBy) - }) - - t.Run("retry succeeded task fails", func(t *testing.T) { - execution := &model.TaskExecution{ - TaskID: "retry_succeeded_001", - TaskType: "system:cleanup", - TaskName: "系统垃圾清理", - Status: model.TaskExecutionStatusSucceeded, - Retryable: true, - MaxRetry: 3, - TriggeredBy: "manual", - } - err := repository.CreateTaskExecution(ctx, execution) - require.NoError(t, err) - - url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID) - req, _ := http.NewRequest("POST", url, nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusBadRequest, w.Code) - }) - - t.Run("retry non-retryable task fails", func(t *testing.T) { - execution := &model.TaskExecution{ - TaskID: "retry_not_allowed_001", - TaskType: "system:cleanup", - TaskName: "系统垃圾清理", - Status: model.TaskExecutionStatusFailed, - Retryable: false, - TriggeredBy: "manual", - } - err := repository.CreateTaskExecution(ctx, execution) - require.NoError(t, err) - - url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID) - req, _ := http.NewRequest("POST", url, nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusBadRequest, w.Code) - }) - - t.Run("retry non-existent task", func(t *testing.T) { - req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/executions/99999999/retry", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusNotFound, w.Code) - }) - - t.Run("retry with invalid ID", func(t *testing.T) { - req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/executions/invalid/retry", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - assert.Equal(t, http.StatusBadRequest, w.Code) - }) -} diff --git a/internal/apps/admin/template/errs.go b/internal/apps/admin/template/errs.go deleted file mode 100644 index d8dda622..00000000 --- a/internal/apps/admin/template/errs.go +++ /dev/null @@ -1,16 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package template 提供模板管理功能 -package template - -// 模板管理相关错误消息 -const ( - TemplateNotFound = "模板不存在" - TemplateKeyRequired = "模板标识符不能为空" - TemplateNameRequired = "模板名称不能为空" - TemplateContentRequired = "模板内容不能为空" - TemplateKeyExists = "模板标识符已存在" - SystemTemplateCannotDelete = "系统预置模板不可删除" - SystemTemplateCannotModifyKey = "系统预置模板不可修改标识符" -) diff --git a/internal/apps/admin/template/logics.go b/internal/apps/admin/template/logics.go deleted file mode 100644 index 9cae21cf..00000000 --- a/internal/apps/admin/template/logics.go +++ /dev/null @@ -1,78 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package template - -import ( - "context" - "errors" - - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" -) - -func createTemplate(ctx context.Context, req CreateTemplateRequest) (model.Template, error) { - exists, err := repository.TemplateExistsByKey(ctx, req.Key) - if err != nil { - return model.Template{}, err - } - if exists { - return model.Template{}, errors.New(TemplateKeyExists) - } - - tmpl := model.Template{ - Key: req.Key, - Name: req.Name, - Type: req.Type, - Subject: req.Subject, - Content: req.Content, - Description: req.Description, - IsSystem: false, - } - if err := tmpl.Validate(); err != nil { - return model.Template{}, err - } - if err := repository.CreateTemplate(ctx, &tmpl); err != nil { - return model.Template{}, err - } - return tmpl, nil -} - -func listTemplates(ctx context.Context) ([]model.Template, error) { - return repository.ListTemplates(ctx) -} - -func getTemplate(ctx context.Context, key string) (model.Template, error) { - return repository.GetTemplateByKey(ctx, key) -} - -func updateTemplate(ctx context.Context, key string, req UpdateTemplateRequest) (model.Template, error) { - tmpl, err := repository.GetTemplateByKey(ctx, key) - if err != nil { - return model.Template{}, err - } - - tmpl.Name = req.Name - tmpl.Type = req.Type - tmpl.Subject = req.Subject - tmpl.Content = req.Content - tmpl.Description = req.Description - if err := tmpl.Validate(); err != nil { - return model.Template{}, err - } - if err := repository.SaveTemplate(ctx, &tmpl); err != nil { - return model.Template{}, err - } - return tmpl, nil -} - -func deleteTemplate(ctx context.Context, key string) error { - tmpl, err := repository.GetTemplateByKey(ctx, key) - if err != nil { - return err - } - if tmpl.IsSystem { - return errors.New(SystemTemplateCannotDelete) - } - return repository.DeleteTemplate(ctx, &tmpl) -} diff --git a/internal/apps/admin/template/routers.go b/internal/apps/admin/template/routers.go deleted file mode 100644 index d4958dd6..00000000 --- a/internal/apps/admin/template/routers.go +++ /dev/null @@ -1,176 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package template - -import ( - "errors" - "net/http" - - "github.com/gin-gonic/gin" - "gorm.io/gorm" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -// CreateTemplateRequest 创建模板请求 -type CreateTemplateRequest struct { - Key string `json:"key" binding:"required,max=80"` - Name string `json:"name" binding:"required,max=100"` - Type string `json:"type" binding:"required,max=20"` - Subject string `json:"subject" binding:"max=255"` - Content string `json:"content" binding:"required"` - Description string `json:"description" binding:"max=255"` -} - -// UpdateTemplateRequest 更新模板请求 -type UpdateTemplateRequest struct { - Name string `json:"name" binding:"required,max=100"` - Type string `json:"type" binding:"required,max=20"` - Subject string `json:"subject" binding:"max=255"` - Content string `json:"content" binding:"required"` - Description string `json:"description" binding:"max=255"` -} - -func abortTemplateLogicError(c *gin.Context, err error) bool { - if err == nil { - return false - } - if errors.Is(err, gorm.ErrRecordNotFound) { - response.AbortNotFound(c, TemplateNotFound) - return true - } - msg := err.Error() - switch msg { - case TemplateKeyExists, SystemTemplateCannotDelete: - response.AbortBadRequest(c, msg) - return true - } - response.AbortInternal(c, msg) - return true -} - -// CreateTemplate 创建模板 -// @Summary 创建模板 -// @Description 创建一条新的自定义通知模板,模板标识符(Key)不可重复,需要管理员权限 -// @Tags admin -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param request body template.CreateTemplateRequest true "创建请求参数" -// @Success 200 {object} response.Any{data=string} "创建成功" -// @Failure 400 {object} response.Any "参数错误或模板标识符已存在" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/templates [post] -func CreateTemplate(c *gin.Context) { - var req CreateTemplateRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - tmpl, err := createTemplate(c.Request.Context(), req) - if abortTemplateLogicError(c, err) { - return - } - - c.JSON(http.StatusOK, response.OK(tmpl)) -} - -// ListTemplates 获取模板列表 -// @Summary 获取模板列表 -// @Description 返回所有通知模板列表,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=[]model.Template} "模板列表" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/templates [get] -func ListTemplates(c *gin.Context) { - templates, err := listTemplates(c.Request.Context()) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - - c.JSON(http.StatusOK, response.OK(templates)) -} - -// GetTemplate 获取单个模板 -// @Summary 获取单个模板 -// @Description 根据模板标识符获取对应的模板详情,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Param key path string true "模板标识符" -// @Success 200 {object} response.Any{data=model.Template} "模板详情" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 404 {object} response.Any "模板不存在" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/templates/{key} [get] -func GetTemplate(c *gin.Context) { - tmpl, err := getTemplate(c.Request.Context(), c.Param("key")) - if abortTemplateLogicError(c, err) { - return - } - - c.JSON(http.StatusOK, response.OK(tmpl)) -} - -// UpdateTemplate 更新模板 -// @Summary 更新模板 -// @Description 根据模板标识符更新对应的模板内容,需要管理员权限 -// @Tags admin -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param key path string true "模板标识符" -// @Param request body template.UpdateTemplateRequest true "更新请求参数" -// @Success 200 {object} response.Any{data=model.Template} "更新成功" -// @Failure 400 {object} response.Any "参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 404 {object} response.Any "模板不存在" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/templates/{key} [put] -func UpdateTemplate(c *gin.Context) { - var req UpdateTemplateRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - tmpl, err := updateTemplate(c.Request.Context(), c.Param("key"), req) - if abortTemplateLogicError(c, err) { - return - } - - c.JSON(http.StatusOK, response.OK(tmpl)) -} - -// DeleteTemplate 删除模板 -// @Summary 删除模板 -// @Description 根据模板标识符删除对应模板,系统预置模板不可删除,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Param key path string true "模板标识符" -// @Success 200 {object} response.Any{data=string} "删除成功" -// @Failure 400 {object} response.Any "不可删除系统模板" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 404 {object} response.Any "模板不存在" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/templates/{key} [delete] -func DeleteTemplate(c *gin.Context) { - if err := deleteTemplate(c.Request.Context(), c.Param("key")); abortTemplateLogicError(c, err) { - return - } - - c.JSON(http.StatusOK, response.OKNil()) -} diff --git a/internal/apps/admin/template/routers_test.go b/internal/apps/admin/template/routers_test.go deleted file mode 100644 index e615a8a5..00000000 --- a/internal/apps/admin/template/routers_test.go +++ /dev/null @@ -1,242 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package template - -import ( - "bytes" - "encoding/json" - "net/http" - "net/http/httptest" - "testing" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -func setupTestRouter(authUser *model.User) *gin.Engine { - r := testhelper.NewTestGinEngine() - adminGroup := r.Group("/api/v1/admin") - - // Mock authentication middleware - adminGroup.Use(func(c *gin.Context) { - if authUser != nil { - oauth.SetToContext(c, oauth.UserObjKey, authUser) - } - c.Next() - }) - - adminGroup.GET("/templates", ListTemplates) - adminGroup.POST("/templates", CreateTemplate) - - templateRouter := adminGroup.Group("/templates/:key") - { - templateRouter.GET("", GetTemplate) - templateRouter.PUT("", UpdateTemplate) - templateRouter.DELETE("", DeleteTemplate) - } - - return r -} - -func TestCreateTemplate(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("create successfully", func(t *testing.T) { - payload := CreateTemplateRequest{ - Key: "test_template", - Name: "Test Template", - Type: "email", - Subject: "Test Subject", - Content: "Hello {{.Name}}", - Description: "Test Desc", - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("POST", "/api/v1/admin/templates", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - var tmpl model.Template - err := dbConn.Where("key = ?", "test_template").First(&tmpl).Error - if err != nil { - t.Fatalf("failed to find template in DB: %v", err) - } - if tmpl.Name != "Test Template" { - t.Errorf("expected Name 'Test Template', got '%s'", tmpl.Name) - } - }) - - t.Run("create duplicate key error", func(t *testing.T) { - payload := CreateTemplateRequest{ - Key: "test_template", - Name: "Another Name", - Type: "email", - Subject: "Another Subject", - Content: "Hello", - Description: "desc", - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("POST", "/api/v1/admin/templates", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusBadRequest { - t.Errorf("expected 400 Bad Request on duplicate key, got %d", w.Code) - } - }) -} - -func TestListTemplates(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - // Seed system templates manually for testing - t1 := model.Template{Key: "login_email", Name: "Login Code", Type: "email", Content: "code {{.Code}}", IsSystem: true} - t2 := model.Template{Key: "register_email", Name: "Register Code", Type: "email", Content: "code {{.Code}}", IsSystem: true} - dbConn.Create(&t1) - dbConn.Create(&t2) - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("list templates", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/templates", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK, got %d", w.Code) - } - - var resp response.Any - _ = json.Unmarshal(w.Body.Bytes(), &resp) - - dataBytes, _ := json.Marshal(resp.Data) - var templates []model.Template - _ = json.Unmarshal(dataBytes, &templates) - - if len(templates) != 2 { - t.Errorf("expected 2 templates, got %d", len(templates)) - } - }) -} - -func TestGetTemplate(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - t1 := model.Template{Key: "login_email", Name: "Login Code", Type: "email", Content: "code {{.Code}}", IsSystem: true} - dbConn.Create(&t1) - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("get existing", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/templates/login_email", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Errorf("expected 200 OK, got %d", w.Code) - } - }) - - t.Run("get non-existent", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/templates/non_existent", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusNotFound { - t.Errorf("expected 404 Not Found, got %d", w.Code) - } - }) -} - -func TestUpdateTemplate(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - t1 := model.Template{Key: "login_email", Name: "Login Code", Type: "email", Content: "code {{.Code}}", IsSystem: true} - dbConn.Create(&t1) - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("update successfully", func(t *testing.T) { - payload := UpdateTemplateRequest{ - Name: "Updated Login Code", - Type: "email", - Subject: "New Subject", - Content: "new code {{.Code}}", - Description: "new desc", - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("PUT", "/api/v1/admin/templates/login_email", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - var tmpl model.Template - dbConn.Where("key = ?", "login_email").First(&tmpl) - if tmpl.Name != "Updated Login Code" || tmpl.Subject != "New Subject" { - t.Errorf("database values not updated: %+v", tmpl) - } - }) -} - -func TestDeleteTemplate(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - t1 := model.Template{Key: "login_email", Name: "Login Code", Type: "email", Content: "code {{.Code}}", IsSystem: true} - t2 := model.Template{Key: "custom_tmpl", Name: "Custom", Type: "email", Content: "hi", IsSystem: false} - dbConn.Create(&t1) - dbConn.Create(&t2) - - adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("delete system template should fail", func(t *testing.T) { - req, _ := http.NewRequest("DELETE", "/api/v1/admin/templates/login_email", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusBadRequest { - t.Errorf("expected 400 Bad Request when deleting system template, got %d", w.Code) - } - }) - - t.Run("delete custom template should succeed", func(t *testing.T) { - req, _ := http.NewRequest("DELETE", "/api/v1/admin/templates/custom_tmpl", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK, got %d", w.Code) - } - - var count int64 - dbConn.Model(&model.Template{}).Where("key = ?", "custom_tmpl").Count(&count) - if count != 0 { - t.Error("custom template was not deleted from DB") - } - }) -} diff --git a/internal/apps/admin/updater/errs.go b/internal/apps/admin/updater/errs.go deleted file mode 100644 index ecdbbecd..00000000 --- a/internal/apps/admin/updater/errs.go +++ /dev/null @@ -1,17 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package updater manages GitHub Release checks and in-place application upgrades. -package updater - -const ( - errInvalidRepository = "上游仓库地址无效" - errReleaseRequestFailed = "获取上游版本失败" - errReleaseResponseInvalid = "上游版本响应无效" - errNoCompatibleRelease = "未找到兼容的 Release" - errNoCompatibleAsset = "未找到当前系统对应的 Release 资产" - errDevelopmentBuild = "开发版本无法执行自动升级" - errAlreadyUpToDate = "当前已是最新版本" - errUpgradeAlreadyRunning = "已有升级任务正在执行" - errAutomaticUpgradeBlocked = "当前平台暂不支持自动替换二进制" -) diff --git a/internal/apps/admin/updater/logics.go b/internal/apps/admin/updater/logics.go deleted file mode 100644 index 8a5ac9d4..00000000 --- a/internal/apps/admin/updater/logics.go +++ /dev/null @@ -1,647 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package updater - -import ( - "archive/tar" - "archive/zip" - "compress/gzip" - "context" - "encoding/json" - "errors" - "fmt" - "io" - "net/http" - "net/url" - "os" - "path/filepath" - "runtime" - "strings" - "sync" - "time" - - "github.com/Rain-kl/Wavelet/internal/buildinfo" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/pkg/logger" - "golang.org/x/mod/semver" -) - -const ( - githubAPIBaseURL = "https://api.github.com" - maxArchiveSize = int64(1024 * 1024 * 1024) - maxReleaseSize = int64(4 * 1024 * 1024) - repositoryParts = 2 - windowsOS = "windows" - archiveFileMode = 0o600 - stagedBinaryMode = 0o700 -) - -type releaseAsset struct { - Name string `json:"name"` - BrowserDownloadURL string `json:"browser_download_url"` - Size int64 `json:"size"` - State string `json:"state"` -} - -type githubRelease struct { - TagName string `json:"tag_name"` - Name string `json:"name"` - Body string `json:"body"` - HTMLURL string `json:"html_url"` - Draft bool `json:"draft"` - Prerelease bool `json:"prerelease"` - Published time.Time `json:"published_at"` - Assets []releaseAsset `json:"assets"` -} - -// Status describes the current build and the newest compatible upstream release. -type Status struct { - CurrentVersion string `json:"current_version"` - BuildTime string `json:"build_time"` - LatestVersion string `json:"latest_version"` - UpdateAvailable bool `json:"update_available"` - CanUpgrade bool `json:"can_upgrade"` - Prerelease bool `json:"prerelease"` - ReleaseName string `json:"release_name"` - ReleaseNotes string `json:"release_notes"` - ReleaseURL string `json:"release_url"` - PublishedAt string `json:"published_at"` - UpstreamRepository string `json:"upstream_repository"` - AssetName string `json:"asset_name"` - Platform string `json:"platform"` -} - -type releaseClient interface { - Do(req *http.Request) (*http.Response, error) -} - -type manager struct { - client releaseClient - mu sync.Mutex - upgrading bool -} - -var defaultManager = &manager{ - client: &http.Client{Timeout: 10 * time.Minute}, -} - -func normalizeVersion(version string) string { - version = strings.TrimSpace(version) - if version == "" || version == "dev" { - return "" - } - if !strings.HasPrefix(version, "v") { - version = "v" + version - } - if !semver.IsValid(version) { - return "" - } - return version -} - -func parseRepository(raw string) (string, error) { - raw = strings.TrimSpace(raw) - if raw == "" { - return "", errors.New(errInvalidRepository) - } - - if !strings.Contains(raw, "://") { - repo := strings.TrimSuffix(strings.Trim(raw, "/"), ".git") - if len(strings.Split(repo, "/")) == repositoryParts { - return repo, nil - } - return "", errors.New(errInvalidRepository) - } - - parsed, err := url.Parse(raw) - if err != nil || !strings.EqualFold(parsed.Hostname(), "github.com") { - return "", errors.New(errInvalidRepository) - } - repo := strings.TrimSuffix(strings.Trim(parsed.Path, "/"), ".git") - if len(strings.Split(repo, "/")) != repositoryParts { - return "", errors.New(errInvalidRepository) - } - return repo, nil -} - -func expectedAssetName(tag string) string { - extension := "tar.gz" - if runtime.GOOS == windowsOS { - extension = "zip" - } - return fmt.Sprintf("wavelet_%s_%s_%s.%s", tag, runtime.GOOS, runtime.GOARCH, extension) -} - -func expectedAssetNames(repository, tag string) []string { - names := []string{expectedAssetName(tag)} - if parts := strings.Split(repository, "/"); len(parts) == repositoryParts { - repoName := parts[1] - if repoName != "wavelet" { - extension := "tar.gz" - if runtime.GOOS == windowsOS { - extension = "zip" - } - names = append(names, fmt.Sprintf("%s_%s_%s_%s.%s", repoName, tag, runtime.GOOS, runtime.GOARCH, extension)) - } - } - return names -} - -func selectLatestRelease(repository string, releases []githubRelease) (githubRelease, releaseAsset, error) { - var selected githubRelease - var selectedAsset releaseAsset - selectedVersion := "" - - for _, release := range releases { - version := normalizeVersion(release.TagName) - if release.Draft || version == "" { - continue - } - expectedNames := expectedAssetNames(repository, release.TagName) - for _, asset := range release.Assets { - matched := false - for _, name := range expectedNames { - if asset.Name == name { - matched = true - break - } - } - if !matched || asset.BrowserDownloadURL == "" || asset.State != "uploaded" { - continue - } - if selectedVersion == "" || semver.Compare(version, selectedVersion) > 0 { - selected = release - selectedAsset = asset - selectedVersion = version - } - } - } - - if selectedVersion == "" { - return githubRelease{}, releaseAsset{}, errors.New(errNoCompatibleRelease) - } - return selected, selectedAsset, nil -} - -func (m *manager) fetchRelease(ctx context.Context, repository string) (githubRelease, releaseAsset, error) { - req, err := http.NewRequestWithContext( - ctx, - http.MethodGet, - fmt.Sprintf("%s/repos/%s/releases?per_page=30", githubAPIBaseURL, repository), - nil, - ) - if err != nil { - return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errReleaseRequestFailed, err) - } - req.Header.Set("Accept", "application/vnd.github+json") - req.Header.Set("User-Agent", "Wavelet-Updater") - req.Header.Set("X-GitHub-Api-Version", "2022-11-28") - - resp, err := m.client.Do(req) - if err != nil { - return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errReleaseRequestFailed, err) - } - defer func() { - // The response body is read-only; close errors cannot affect the parsed result. - _ = resp.Body.Close() - }() - if resp.StatusCode != http.StatusOK { - return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: HTTP %d", errReleaseRequestFailed, resp.StatusCode) - } - - var releases []githubRelease - decoder := json.NewDecoder(io.LimitReader(resp.Body, maxReleaseSize)) - if err := decoder.Decode(&releases); err != nil { - return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errReleaseResponseInvalid, err) - } - - release, asset, err := selectLatestRelease(repository, releases) - if err != nil { - return githubRelease{}, releaseAsset{}, err - } - logger.InfoF(ctx, "[Updater] Selected latest compatible release: %s (Asset: %s)", release.TagName, asset.Name) - return release, asset, nil -} - -func loadRepository(ctx context.Context) (string, error) { - config, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUpdateUpstreamRepository) - if err != nil { - return "", fmt.Errorf("%s: %w", errInvalidRepository, err) - } - return parseRepository(config.Value) -} - -func (m *manager) status(ctx context.Context) (Status, releaseAsset, error) { - upstreamRepo, err := loadRepository(ctx) - if err != nil { - return Status{}, releaseAsset{}, err - } - release, asset, err := m.fetchRelease(ctx, upstreamRepo) - if err != nil { - return Status{}, releaseAsset{}, err - } - - currentVersion := normalizeVersion(buildinfo.Version) - latestVersion := normalizeVersion(release.TagName) - updateAvailable := currentVersion != "" && semver.Compare(latestVersion, currentVersion) > 0 - - logger.InfoF(ctx, "[Updater] Check update complete. current: %s, latest: %s, update_available: %t", buildinfo.Version, release.TagName, updateAvailable) - - return Status{ - CurrentVersion: buildinfo.Version, - BuildTime: buildinfo.BuildTime, - LatestVersion: release.TagName, - UpdateAvailable: updateAvailable, - CanUpgrade: updateAvailable && runtime.GOOS != windowsOS, - Prerelease: release.Prerelease, - ReleaseName: release.Name, - ReleaseNotes: release.Body, - ReleaseURL: release.HTMLURL, - PublishedAt: release.Published.Format(time.RFC3339), - UpstreamRepository: upstreamRepo, - AssetName: asset.Name, - Platform: runtime.GOOS + "/" + runtime.GOARCH, - }, asset, nil -} - -func downloadArchive(ctx context.Context, client releaseClient, asset releaseAsset, destination string) error { - if asset.Size <= 0 || asset.Size > maxArchiveSize { - return fmt.Errorf("release 资产大小无效: %d", asset.Size) - } - logger.InfoF(ctx, "[Updater] Downloading release asset: %s", asset.Name) - req, err := http.NewRequestWithContext(ctx, http.MethodGet, asset.BrowserDownloadURL, nil) - if err != nil { - return fmt.Errorf("创建升级下载请求失败: %w", err) - } - req.Header.Set("User-Agent", "Wavelet-Updater") - - resp, err := client.Do(req) - if err != nil { - return fmt.Errorf("下载升级资产失败: %w", err) - } - defer func() { - // The downloaded body has already been validated by size before use. - _ = resp.Body.Close() - }() - if resp.StatusCode != http.StatusOK { - return fmt.Errorf("下载升级资产失败: HTTP %d", resp.StatusCode) - } - - file, err := os.OpenFile(destination, os.O_CREATE|os.O_EXCL|os.O_WRONLY, archiveFileMode) //nolint:gosec // destination is created inside the verified executable directory. - if err != nil { - return fmt.Errorf("创建升级归档失败: %w", err) - } - - written, err := io.Copy(file, io.LimitReader(resp.Body, maxArchiveSize+1)) - if err != nil { - _ = file.Close() - return fmt.Errorf("写入升级归档失败: %w", err) - } - if err := file.Close(); err != nil { - return fmt.Errorf("关闭升级归档失败: %w", err) - } - if written > maxArchiveSize || written != asset.Size { - return fmt.Errorf("升级归档大小不匹配: got %d, want %d", written, asset.Size) - } - logger.InfoF(ctx, "[Updater] Successfully downloaded release asset to %s", destination) - return nil -} - -func safeArchivePath(destination, name string) (string, error) { - cleanName := filepath.Clean(name) - if filepath.IsAbs(cleanName) || cleanName == "." || strings.HasPrefix(cleanName, ".."+string(filepath.Separator)) { - return "", fmt.Errorf("归档包含非法路径: %s", name) - } - target := filepath.Join(destination, cleanName) - relative, err := filepath.Rel(destination, target) - if err != nil || relative == ".." || strings.HasPrefix(relative, ".."+string(filepath.Separator)) { - return "", fmt.Errorf("归档路径越界: %s", name) - } - return target, nil -} - -func matchBinaryName(name string, candidates []string) bool { - for _, candidate := range candidates { - if runtime.GOOS == windowsOS { - if strings.EqualFold(name, candidate) { - return true - } - } else { - if name == candidate { - return true - } - } - } - return false -} - -func getCandidateBinaryNames(executable string, repository string) []string { - execName := filepath.Base(executable) - names := []string{execName} - - addName := func(base string) { - name := base - if runtime.GOOS == windowsOS && !strings.HasSuffix(strings.ToLower(name), ".exe") { - name += ".exe" - } - for _, existing := range names { - if existing == name { - return - } - } - names = append(names, name) - } - - if parts := strings.Split(repository, "/"); len(parts) == repositoryParts { - addName(parts[1]) - } - addName("wavelet") - - return names -} - -func isLikelyBinary(name string, isDir bool, mode os.FileMode) bool { - if isDir { - return false - } - base := strings.ToLower(filepath.Base(name)) - - // Exclude typical non-binary metadata files - exclusions := []string{ - "license", "licence", "copying", "notice", "readme", "changelog", - } - for _, excl := range exclusions { - if strings.HasPrefix(base, excl) { - return false - } - } - - if runtime.GOOS == windowsOS { - return filepath.Ext(base) == ".exe" - } - - // On Unix, it should either have the executable permission bit set, OR have no extension - return (mode.Perm()&0111 != 0) || (filepath.Ext(base) == "") -} - -func findBinaryInTarGz(archivePath string, candidates []string) (string, error) { - file, err := os.Open(archivePath) //nolint:gosec // archivePath is created by prepareUpgrade in the executable directory. - if err != nil { - return "", err - } - defer func() { - _ = file.Close() - }() - - gzipReader, err := gzip.NewReader(file) - if err != nil { - return "", err - } - defer func() { - _ = gzipReader.Close() - }() - - reader := tar.NewReader(gzipReader) - var binaries []string - for { - header, err := reader.Next() - if errors.Is(err, io.EOF) { - break - } - if err != nil { - return "", err - } - if header.Typeflag == tar.TypeReg && isLikelyBinary(header.Name, false, header.FileInfo().Mode()) { - binaries = append(binaries, header.Name) - } - } - - if len(binaries) == 1 { - return binaries[0], nil - } - - // Fallback to candidate match if multiple or zero likely binaries found - for _, name := range binaries { - if matchBinaryName(filepath.Base(name), candidates) { - return name, nil - } - } - - return "", errors.New(errNoCompatibleAsset) -} - -func findBinaryInZip(archivePath string, candidates []string) (string, error) { - reader, err := zip.OpenReader(archivePath) - if err != nil { - return "", err - } - defer func() { - _ = reader.Close() - }() - - var binaries []string - for _, file := range reader.File { - if !file.FileInfo().IsDir() && isLikelyBinary(file.Name, false, file.FileInfo().Mode()) { - binaries = append(binaries, file.Name) - } - } - - if len(binaries) == 1 { - return binaries[0], nil - } - - // Fallback to candidate match if multiple or zero likely binaries found - for _, name := range binaries { - if matchBinaryName(filepath.Base(name), candidates) { - return name, nil - } - } - - return "", errors.New(errNoCompatibleAsset) -} - -func extractTarGz(ctx context.Context, archivePath, destination, targetName string, candidates []string) (string, error) { - binaryPathInArchive, err := findBinaryInTarGz(archivePath, candidates) - if err != nil { - return "", err - } - - logger.InfoF(ctx, "[Updater] Extracting tar.gz archive: %s (extracting: %s)", archivePath, binaryPathInArchive) - file, err := os.Open(archivePath) //nolint:gosec // archivePath is created by prepareUpgrade in the executable directory. - if err != nil { - return "", err - } - defer func() { - // Read-only archive close errors do not change extraction validity. - _ = file.Close() - }() - gzipReader, err := gzip.NewReader(file) - if err != nil { - return "", err - } - defer func() { - // The gzip checksum is verified while reading the selected file. - _ = gzipReader.Close() - }() - - reader := tar.NewReader(gzipReader) - for { - header, err := reader.Next() - if errors.Is(err, io.EOF) { - break - } - if err != nil { - return "", err - } - if header.Name != binaryPathInArchive { - continue - } - target, err := safeArchivePath(destination, targetName) - if err != nil { - return "", err - } - output, err := os.OpenFile(target, os.O_CREATE|os.O_EXCL|os.O_WRONLY, stagedBinaryMode) //nolint:gosec // target is constrained by safeArchivePath. - if err != nil { - return "", err - } - written, copyErr := io.Copy(output, io.LimitReader(reader, maxArchiveSize+1)) - closeErr := output.Close() - if copyErr != nil { - return "", copyErr - } - if closeErr != nil { - return "", closeErr - } - if written > maxArchiveSize { - return "", errors.New("解压后的程序文件超过大小限制") - } - logger.InfoF(ctx, "[Updater] Successfully extracted binary to %s", target) - return target, nil - } - return "", errors.New(errNoCompatibleAsset) -} - -func extractZip(ctx context.Context, archivePath, destination, targetName string, candidates []string) (string, error) { - binaryPathInArchive, err := findBinaryInZip(archivePath, candidates) - if err != nil { - return "", err - } - - logger.InfoF(ctx, "[Updater] Extracting zip archive: %s (extracting: %s)", archivePath, binaryPathInArchive) - reader, err := zip.OpenReader(archivePath) - if err != nil { - return "", err - } - defer func() { - // Read-only archive close errors do not change extraction validity. - _ = reader.Close() - }() - for _, file := range reader.File { - if file.Name != binaryPathInArchive { - continue - } - target, err := safeArchivePath(destination, targetName) - if err != nil { - return "", err - } - input, err := file.Open() - if err != nil { - return "", err - } - output, err := os.OpenFile(target, os.O_CREATE|os.O_EXCL|os.O_WRONLY, stagedBinaryMode) //nolint:gosec // target is constrained by safeArchivePath. - if err != nil { - _ = input.Close() - return "", err - } - written, copyErr := io.Copy(output, io.LimitReader(input, maxArchiveSize+1)) - inputCloseErr := input.Close() - outputCloseErr := output.Close() - if copyErr != nil { - return "", copyErr - } - if inputCloseErr != nil { - return "", inputCloseErr - } - if outputCloseErr != nil { - return "", outputCloseErr - } - if written > maxArchiveSize { - return "", errors.New("解压后的程序文件超过大小限制") - } - logger.InfoF(ctx, "[Updater] Successfully extracted binary to %s", target) - return target, nil - } - return "", errors.New(errNoCompatibleAsset) -} - -func (m *manager) prepareUpgrade(ctx context.Context) (string, string, error) { - if runtime.GOOS == windowsOS { - return "", "", errors.New(errAutomaticUpgradeBlocked) - } - if normalizeVersion(buildinfo.Version) == "" { - return "", "", errors.New(errDevelopmentBuild) - } - - m.mu.Lock() - defer m.mu.Unlock() - if m.upgrading { - return "", "", errors.New(errUpgradeAlreadyRunning) - } - - status, asset, err := m.status(ctx) - if err != nil { - return "", "", err - } - if !status.UpdateAvailable { - return "", "", errors.New(errAlreadyUpToDate) - } - - logger.InfoF(ctx, "[Updater] Preparing upgrade. current: %s, latest: %s", status.CurrentVersion, status.LatestVersion) - - executable, err := os.Executable() - if err != nil { - return "", "", fmt.Errorf("定位当前程序失败: %w", err) - } - executable, err = filepath.EvalSymlinks(executable) - if err != nil { - return "", "", fmt.Errorf("解析当前程序路径失败: %w", err) - } - - tempDir, err := os.MkdirTemp(filepath.Dir(executable), ".wavelet-update-*") - if err != nil { - return "", "", fmt.Errorf("创建升级目录失败: %w", err) - } - - archivePath := filepath.Join(tempDir, asset.Name) - if err := downloadArchive(ctx, m.client, asset, archivePath); err != nil { - // Cleanup is best effort because the download error is the actionable failure. - _ = os.RemoveAll(tempDir) - return "", "", err - } - - targetName := filepath.Base(executable) - candidates := getCandidateBinaryNames(executable, status.UpstreamRepository) - - var stagedBinary string - if strings.HasSuffix(asset.Name, ".zip") { - stagedBinary, err = extractZip(ctx, archivePath, tempDir, targetName, candidates) - } else { - stagedBinary, err = extractTarGz(ctx, archivePath, tempDir, targetName, candidates) - } - if err != nil { - // Cleanup is best effort because the extraction error is the actionable failure. - _ = os.RemoveAll(tempDir) - return "", "", fmt.Errorf("解压升级资产失败: %w", err) - } - logger.InfoF(ctx, "[Updater] Staged binary successfully prepared: %s", stagedBinary) - m.upgrading = true - return executable, stagedBinary, nil -} - -func (m *manager) finishUpgrade() { - m.mu.Lock() - defer m.mu.Unlock() - m.upgrading = false -} diff --git a/internal/apps/admin/updater/logics_test.go b/internal/apps/admin/updater/logics_test.go deleted file mode 100644 index b6fe09f0..00000000 --- a/internal/apps/admin/updater/logics_test.go +++ /dev/null @@ -1,130 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package updater - -import ( - "runtime" - "testing" - "time" -) - -func TestParseRepository(t *testing.T) { - tests := []struct { - name string - input string - want string - wantErr bool - }{ - {name: "short form", input: "Rain-kl/Wavelet", want: "Rain-kl/Wavelet"}, - {name: "GitHub URL", input: "https://github.com/Rain-kl/Wavelet.git", want: "Rain-kl/Wavelet"}, - {name: "unsupported host", input: "https://example.com/Rain-kl/Wavelet", wantErr: true}, - {name: "missing owner", input: "Wavelet", wantErr: true}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got, err := parseRepository(tt.input) - if gotErr := err != nil; gotErr != tt.wantErr { - t.Errorf("parseRepository(%q) error = %v, want error presence = %t", tt.input, err, tt.wantErr) - } - if got != tt.want { - t.Errorf("parseRepository(%q) = %q, want %q", tt.input, got, tt.want) - } - }) - } -} - -func TestSelectLatestRelease(t *testing.T) { - assetNameV1 := expectedAssetName("v1.0.0") - assetNameV2 := expectedAssetName("v2.0.0") - releases := []githubRelease{ - { - TagName: "v1.0.0", - Published: time.Date(2026, time.June, 1, 0, 0, 0, 0, time.UTC), - Assets: []releaseAsset{{ - Name: assetNameV1, - BrowserDownloadURL: "https://example.com/v1", - State: "uploaded", - }}, - }, - { - TagName: "v2.0.0", - Published: time.Date(2026, time.June, 2, 0, 0, 0, 0, time.UTC), - Assets: []releaseAsset{{ - Name: assetNameV2, - BrowserDownloadURL: "https://example.com/v2", - State: "uploaded", - }}, - }, - { - TagName: "v3.0.0", - Assets: []releaseAsset{{ - Name: "wavelet_v3.0.0_other_platform.tar.gz", - BrowserDownloadURL: "https://example.com/v3", - State: "uploaded", - }}, - }, - } - - release, asset, err := selectLatestRelease("Rain-kl/Wavelet", releases) - if err != nil { - t.Fatalf("selectLatestRelease() error = %v", err) - } - if release.TagName != "v2.0.0" { - t.Errorf("selectLatestRelease() tag = %q, want %q", release.TagName, "v2.0.0") - } - if asset.Name != assetNameV2 { - t.Errorf("selectLatestRelease() asset = %q, want %q", asset.Name, assetNameV2) - } -} - -func TestSelectLatestReleaseWithCustomRepo(t *testing.T) { - extension := "tar.gz" - if runtime.GOOS == "windows" { - extension = "zip" - } - releases := []githubRelease{ - { - TagName: "v1.0.0", - Published: time.Date(2026, time.June, 1, 0, 0, 0, 0, time.UTC), - Assets: []releaseAsset{{ - Name: "wavelet_v1.0.0_" + runtime.GOOS + "_" + runtime.GOARCH + "." + extension, - BrowserDownloadURL: "https://example.com/v1", - State: "uploaded", - }}, - }, - { - TagName: "v2.0.0", - Published: time.Date(2026, time.June, 2, 0, 0, 0, 0, time.UTC), - Assets: []releaseAsset{{ - Name: "PixezSync_v2.0.0_" + runtime.GOOS + "_" + runtime.GOARCH + "." + extension, - BrowserDownloadURL: "https://example.com/v2", - State: "uploaded", - }}, - }, - } - - release, asset, err := selectLatestRelease("Rain-kl/PixezSync", releases) - if err != nil { - t.Fatalf("selectLatestRelease() error = %v", err) - } - if release.TagName != "v2.0.0" { - t.Errorf("selectLatestRelease() tag = %q, want %q", release.TagName, "v2.0.0") - } - expectedName := "PixezSync_v2.0.0_" + runtime.GOOS + "_" + runtime.GOARCH + "." + extension - if asset.Name != expectedName { - t.Errorf("selectLatestRelease() asset = %q, want %q", asset.Name, expectedName) - } -} - -func TestExpectedAssetName(t *testing.T) { - extension := "tar.gz" - if runtime.GOOS == "windows" { - extension = "zip" - } - want := "wavelet_v1.2.3_" + runtime.GOOS + "_" + runtime.GOARCH + "." + extension - if got := expectedAssetName("v1.2.3"); got != want { - t.Errorf("expectedAssetName(%q) = %q, want %q", "v1.2.3", got, want) - } -} diff --git a/internal/apps/admin/updater/restart_unix.go b/internal/apps/admin/updater/restart_unix.go deleted file mode 100644 index 24d94c6c..00000000 --- a/internal/apps/admin/updater/restart_unix.go +++ /dev/null @@ -1,50 +0,0 @@ -//go:build !windows - -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package updater - -import ( - "context" - "fmt" - "os" - "path/filepath" - "syscall" - - "github.com/Rain-kl/Wavelet/pkg/logger" -) - -const installedBinaryMode = 0o755 - -func replaceAndRestart(executable, stagedBinary string) error { - ctx := context.Background() - logger.InfoF(ctx, "[Updater] Swapping executable: %s -> %s", executable, stagedBinary) - backup := executable + ".old" - - if err := os.Remove(backup); err != nil && !os.IsNotExist(err) { - return fmt.Errorf("删除旧备份失败: %w", err) - } - - if err := os.Rename(executable, backup); err != nil { - return fmt.Errorf("备份当前程序失败: %w", err) - } - - if err := os.Rename(stagedBinary, executable); err != nil { - _ = os.Rename(backup, executable) - return fmt.Errorf("替换当前程序失败: %w", err) - } - - if err := os.Chmod(executable, installedBinaryMode); err != nil { //nolint:gosec // the installed application binary must be executable. - _ = os.Remove(executable) - _ = os.Rename(backup, executable) - return fmt.Errorf("设置程序执行权限失败: %w", err) - } - - stagingDir := filepath.Dir(stagedBinary) - // Cleanup is best effort; a leftover staging directory must not block restart. - _ = os.RemoveAll(stagingDir) - - logger.InfoF(ctx, "[Updater] Executing syscall.Exec to restart service: %s %v", executable, os.Args) - return syscall.Exec(executable, os.Args, os.Environ()) //nolint:gosec // executable is resolved from os.Executable and never supplied by a request. -} diff --git a/internal/apps/admin/updater/restart_windows.go b/internal/apps/admin/updater/restart_windows.go deleted file mode 100644 index a454cd30..00000000 --- a/internal/apps/admin/updater/restart_windows.go +++ /dev/null @@ -1,12 +0,0 @@ -//go:build windows - -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package updater - -import "errors" - -func replaceAndRestart(_, _ string) error { - return errors.New(errAutomaticUpgradeBlocked) -} diff --git a/internal/apps/admin/updater/routers.go b/internal/apps/admin/updater/routers.go deleted file mode 100644 index 91cd999e..00000000 --- a/internal/apps/admin/updater/routers.go +++ /dev/null @@ -1,69 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package updater - -import ( - "context" - "net/http" - "time" - - "github.com/Rain-kl/Wavelet/pkg/logger" - "github.com/Rain-kl/Wavelet/pkg/util" - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -// GetUpdateStatus 获取应用更新状态 -// @Summary 获取应用更新状态 -// @Description 从系统配置指定的 GitHub 上游仓库查询最新兼容 Release,并与当前服务版本比较 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=updater.Status} "更新状态" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "查询失败" -// @Router /api/v1/admin/update [get] -func GetUpdateStatus(c *gin.Context) { - status, _, err := defaultManager.status(c.Request.Context()) - if err != nil { - logger.ErrorF(c.Request.Context(), "[Updater] check release failed: %v", err) - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(status)) -} - -// ApplyUpdate 下载并应用应用更新 -// @Summary 下载并应用应用更新 -// @Description 下载当前平台对应的 GitHub Actions Release 资产,替换当前二进制并重启进程 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any "升级已准备并即将重启" -// @Failure 400 {object} response.Any "当前版本不可升级" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "升级准备失败" -// @Router /api/v1/admin/update/apply [post] -func ApplyUpdate(c *gin.Context) { - executable, stagedBinary, err := defaultManager.prepareUpgrade(c.Request.Context()) - if err != nil { - logger.ErrorF(c.Request.Context(), "[Updater] prepare upgrade failed: %v", err) - response.AbortBadRequest(c, err.Error()) - return - } - - logger.InfoF(c.Request.Context(), "[Updater] upgrade prepared; restarting with %s", stagedBinary) - c.JSON(http.StatusOK, response.OKNil()) - - util.Go(func() { - time.Sleep(time.Second) - if err := replaceAndRestart(executable, stagedBinary); err != nil { - defaultManager.finishUpgrade() - logger.ErrorF(context.Background(), "[Updater] replace and restart failed: %v", err) - } - }) -} diff --git a/internal/apps/admin/user/errs.go b/internal/apps/admin/user/errs.go deleted file mode 100644 index 936eb5a2..00000000 --- a/internal/apps/admin/user/errs.go +++ /dev/null @@ -1,23 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package user 提供用户管理功能 -package user - -const ( - userNotFound = "用户不存在" - cannotDisable = "不能禁用管理员用户" - cannotDelete = "不能删除管理员用户" - cannotDeleteSelf = "不能删除当前登录用户" - updateUserFailed = "更新用户状态失败" - deleteUserFailed = "删除用户失败" - usernameExists = "用户名已存在" - usernameRequired = "用户名不能为空" - passwordTooShort = "密码长度不能少于 8 位" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - createUserFailed = "创建用户失败" - emailRequired = "邮箱不能为空" - emailExists = "邮箱已被注册" - cannotRevokeSelfAdmin = "不能撤销当前登录用户的管理员权限" - updateUserInfoFailed = "更新用户信息失败" -) diff --git a/internal/apps/admin/user/logics.go b/internal/apps/admin/user/logics.go deleted file mode 100644 index 20bd6904..00000000 --- a/internal/apps/admin/user/logics.go +++ /dev/null @@ -1,210 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package user - -import ( - "context" - "errors" - "strings" - "time" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" -) - -func listUsers(ctx context.Context, req listUsersRequest) (int64, []model.User, error) { - return repository.ListAdminUsers(ctx, repository.AdminUserListFilter{ - UserID: req.UserID, - Username: strings.TrimSpace(req.Username), - Email: strings.TrimSpace(req.Email), - Page: req.Page, - PageSize: req.PageSize, - }) -} - -func getUserDetail(ctx context.Context, id uint64) (model.User, error) { - return repository.GetAdminUserDetail(ctx, id) -} - -func updateUserStatus(ctx context.Context, id uint64, active bool) error { - flags, err := repository.GetUserAdminFlags(ctx, id) - if err != nil { - return err - } - if !active && flags.IsAdmin { - return errors.New(cannotDisable) - } - - var tokens []model.AccessToken - if !active { - tokens, _ = repository.ListAccessTokensByUserID(ctx, id) - } - - err = repository.UpdateUserActive(ctx, id, active) - if err == nil { - oauth.InvalidateCachedUser(ctx, id) - if !active { - for _, token := range tokens { - oauth.InvalidateCachedToken(ctx, token.TokenHash) - } - } - } - return err -} - -func deleteUser(ctx context.Context, currentUserID, targetID uint64) error { - if currentUserID == targetID { - return errors.New(cannotDeleteSelf) - } - flags, err := repository.GetUserAdminFlags(ctx, targetID) - if err != nil { - return err - } - if flags.IsAdmin { - return errors.New(cannotDelete) - } - - tokens, _ := repository.ListAccessTokensByUserID(ctx, targetID) - - err = repository.DeleteUserWithRelations(ctx, targetID) - if err == nil { - oauth.InvalidateCachedUser(ctx, targetID) - for _, token := range tokens { - oauth.InvalidateCachedToken(ctx, token.TokenHash) - } - } - return err -} - -func createUser(ctx context.Context, req createUserRequest) (model.User, error) { - req.Username = strings.TrimSpace(req.Username) - req.Nickname = strings.TrimSpace(req.Nickname) - req.Password = strings.TrimSpace(req.Password) - req.Email = strings.TrimSpace(req.Email) - - if req.Username == "" { - return model.User{}, errors.New(usernameRequired) - } - if req.Email == "" { - return model.User{}, errors.New(emailRequired) - } - if len(req.Password) < minPasswordLength { - return model.User{}, errors.New(passwordTooShort) - } - - count, err := repository.CountUsersByUsername(ctx, req.Username) - if err != nil { - return model.User{}, err - } - if count > 0 { - return model.User{}, errors.New(usernameExists) - } - - emailCount, err := repository.CountUsersByEmail(ctx, req.Email) - if err != nil { - return model.User{}, err - } - if emailCount > 0 { - return model.User{}, errors.New(emailExists) - } - - newUser := model.User{ - ID: idgen.NextUint64ID(), - Username: req.Username, - Nickname: req.Nickname, - Email: req.Email, - IsActive: req.IsActive, - IsAdmin: req.IsAdmin, - LastLoginAt: time.Time{}, - } - if newUser.Nickname == "" { - newUser.Nickname = req.Username - } - if err := newUser.SetEncryptedPassword(req.Password); err != nil { - return model.User{}, err - } - if err := repository.CreateUser(ctx, &newUser); err != nil { - return model.User{}, err - } - return newUser, nil -} - -type updateUserParam struct { - ID uint64 - Nickname string - Email string - IsAdmin bool - Password string -} - -func updateUser(ctx context.Context, currentUserID uint64, param updateUserParam) error { - param.Nickname = strings.TrimSpace(param.Nickname) - param.Email = strings.TrimSpace(param.Email) - param.Password = strings.TrimSpace(param.Password) - - if param.Email == "" { - return errors.New(emailRequired) - } - - targetUser, err := repository.GetAdminUserDetail(ctx, param.ID) - if err != nil { - return err - } - - // 不能撤销当前登录用户的管理员权限 - if currentUserID == param.ID && !param.IsAdmin && targetUser.IsAdmin { - return errors.New(cannotRevokeSelfAdmin) - } - - // 如果修改了邮箱,检查邮箱是否被其他用户占用 - if targetUser.Email != param.Email { - count, err := repository.CountUsersByEmail(ctx, param.Email) - if err != nil { - return err - } - if count > 0 { - return errors.New(emailExists) - } - } - - // 密码强度校验(如果输入了新密码) - if param.Password != "" && len(param.Password) < minPasswordLength { - return errors.New(passwordTooShort) - } - - // 是否需要撤销 Token (重置密码或取消管理员) - needRevokeTokens := (param.Password != "") || (targetUser.IsAdmin && !param.IsAdmin) - var tokens []model.AccessToken - if needRevokeTokens { - tokens, _ = repository.ListAccessTokensByUserID(ctx, param.ID) - } - - // 更新字段 - targetUser.Nickname = param.Nickname - if targetUser.Nickname == "" { - targetUser.Nickname = targetUser.Username - } - targetUser.Email = param.Email - targetUser.IsAdmin = param.IsAdmin - - if param.Password != "" { - if err := targetUser.SetEncryptedPassword(param.Password); err != nil { - return err - } - } - - // 执行更新 - err = repository.UpdateUser(ctx, &targetUser) - if err == nil { - oauth.InvalidateCachedUser(ctx, param.ID) - if needRevokeTokens { - for _, token := range tokens { - oauth.InvalidateCachedToken(ctx, token.TokenHash) - } - } - } - return err -} diff --git a/internal/apps/admin/user/routers.go b/internal/apps/admin/user/routers.go deleted file mode 100644 index c4009424..00000000 --- a/internal/apps/admin/user/routers.go +++ /dev/null @@ -1,348 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package user - -import ( - "errors" - "net/http" - "strconv" - "time" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/pkg/logger" - "github.com/gin-gonic/gin" - "gorm.io/gorm" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -// minPasswordLength 密码最小长度 -const minPasswordLength = 8 - -// listUsersRequest 用户列表查询请求 -type listUsersRequest struct { - Page int `form:"page" binding:"min=1"` - PageSize int `form:"page_size" binding:"min=1,max=100"` - UserID *uint64 `form:"user_id" binding:"omitempty,gt=0"` - Username string `form:"username"` - Email string `form:"email"` -} - -type user struct { - ID uint64 `json:"id,string"` - Username string `json:"username"` - Nickname string `json:"nickname"` - Email string `json:"email"` - AvatarURL string `json:"avatar_url"` - IsActive bool `json:"is_active"` - IsAdmin bool `json:"is_admin"` - Bio string `json:"bio"` - Phone string `json:"phone"` - Gender string `json:"gender"` - Website string `json:"website"` - Location string `json:"location"` - LastLoginAt time.Time `json:"last_login_at"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at"` -} - -// listUsersResponse 用户列表响应 -type listUsersResponse struct { - Users []user `json:"users"` - Total int64 `json:"total"` -} - -func parseUserID(c *gin.Context) (uint64, bool) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil || id == 0 { - response.AbortBadRequest(c, userNotFound) - return 0, false - } - return id, true -} - -func toUser(u model.User) user { - return user{ - ID: u.ID, - Username: u.Username, - Nickname: u.Nickname, - Email: u.Email, - AvatarURL: u.AvatarURL, - IsActive: u.IsActive, - IsAdmin: u.IsAdmin, - Bio: u.Bio, - Phone: u.Phone, - Gender: u.Gender, - Website: u.Website, - Location: u.Location, - LastLoginAt: u.LastLoginAt, - CreatedAt: u.CreatedAt, - UpdatedAt: u.UpdatedAt, - } -} - -func abortUserLogicError(c *gin.Context, err error, notFoundMsg string, forbiddenMsgs, badRequestMsgs []string) bool { - if err == nil { - return false - } - if errors.Is(err, gorm.ErrRecordNotFound) { - response.AbortNotFound(c, notFoundMsg) - return true - } - msg := err.Error() - for _, m := range badRequestMsgs { - if msg == m { - response.AbortBadRequest(c, msg) - return true - } - } - for _, m := range forbiddenMsgs { - if msg == m { - response.AbortForbidden(c, msg) - return true - } - } - logger.ErrorF(c.Request.Context(), "Admin user error: %v", err) - response.AbortInternal(c, "内部服务器错误") - return true -} - -// ListUsers 获取用户列表 -// @Summary 获取用户列表 -// @Description 分页返回用户列表,支持按用户 ID 和用户名筛选,需要管理员权限 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Param request query listUsersRequest true "查询参数" -// @Success 200 {object} response.Any{data=user.listUsersResponse} "用户列表" -// @Failure 400 {object} response.Any "参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/users [get] -func ListUsers(c *gin.Context) { - var req listUsersRequest - if err := c.ShouldBindQuery(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - total, modelUsers, err := listUsers(c.Request.Context(), req) - if err != nil { - logger.ErrorF(c.Request.Context(), "List admin users failed: %v", err) - response.AbortInternal(c, "获取用户列表失败") - return - } - - users := make([]user, 0, len(modelUsers)) - for _, modelUser := range modelUsers { - users = append(users, toUser(modelUser)) - } - - c.JSON(http.StatusOK, response.OK(listUsersResponse{ - Users: users, - Total: total, - })) -} - -// GetUser 获取用户详情 -// @Summary 获取用户详情 -// @Description 返回指定用户的完整个人资料和系统状态,需要管理员权限,不返回密码等敏感字段 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Param id path int true "用户 ID" -// @Success 200 {object} response.Any{data=user.user} "用户详情" -// @Failure 400 {object} response.Any "参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 404 {object} response.Any "用户不存在" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/users/{id} [get] -func GetUser(c *gin.Context) { - id, ok := parseUserID(c) - if !ok { - return - } - - targetUser, err := getUserDetail(c.Request.Context(), id) - if abortUserLogicError(c, err, userNotFound, nil, nil) { - return - } - - c.JSON(http.StatusOK, response.OK(toUser(targetUser))) -} - -// updateUserStatusRequest 更新用户状态请求 -type updateUserStatusRequest struct { - IsActive bool `json:"is_active"` -} - -// UpdateUserStatus 更新用户状态(启用/禁用) -// @Summary 更新用户状态 -// @Description 启用或禁用指定用户,管理员账号无法被禁用,需要管理员权限 -// @Tags admin -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param id path int true "用户 ID" -// @Param request body updateUserStatusRequest true "状态参数" -// @Success 200 {object} response.Any{data=string} "更新成功" -// @Failure 400 {object} response.Any "参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限或尝试禁用管理员" -// @Failure 404 {object} response.Any "用户不存在" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/users/{id}/status [put] -func UpdateUserStatus(c *gin.Context) { - var req updateUserStatusRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - id, ok := parseUserID(c) - if !ok { - return - } - - if err := updateUserStatus(c.Request.Context(), id, req.IsActive); err != nil { - if abortUserLogicError(c, err, userNotFound, []string{cannotDisable}, nil) { - return - } - response.AbortInternal(c, updateUserFailed) - return - } - - c.JSON(http.StatusOK, response.OKNil()) -} - -// DeleteUser 删除用户 -// @Summary 删除用户 -// @Description 删除指定非管理员用户,需要管理员权限,不能删除当前登录用户 -// @Tags admin -// @Produce json -// @Security SessionCookie -// @Param id path int true "用户 ID" -// @Success 200 {object} response.Any{data=string} "删除成功" -// @Failure 400 {object} response.Any "参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限、尝试删除管理员或当前用户" -// @Failure 404 {object} response.Any "用户不存在" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/users/{id} [delete] -func DeleteUser(c *gin.Context) { - id, ok := parseUserID(c) - if !ok { - return - } - - currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) - if err := deleteUser(c.Request.Context(), currUser.ID, id); err != nil { - if abortUserLogicError(c, err, userNotFound, []string{cannotDelete, cannotDeleteSelf}, nil) { - return - } - response.AbortInternal(c, deleteUserFailed) - return - } - - c.JSON(http.StatusOK, response.OKNil()) -} - -// createUserRequest 创建用户请求 -type createUserRequest struct { - Username string `json:"username" binding:"required,min=3,max=64"` - Password string `json:"password" binding:"required,min=8,max=64"` - Nickname string `json:"nickname" binding:"omitempty,max=64"` - Email string `json:"email" binding:"required,email,max=255"` - IsActive bool `json:"is_active"` - IsAdmin bool `json:"is_admin"` -} - -// CreateUser 创建用户 -// @Summary 创建用户 -// @Description 创建一个本地密码登录的新用户,需要管理员权限 -// @Tags admin -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param request body user.createUserRequest true "创建用户参数" -// @Success 200 {object} response.Any{data=user.user} "创建成功" -// @Failure 400 {object} response.Any "参数错误或用户名已存在" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/users [post] -func CreateUser(c *gin.Context) { - var req createUserRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - newUser, err := createUser(c.Request.Context(), req) - if abortUserLogicError(c, err, "", nil, []string{usernameRequired, emailRequired, passwordTooShort, usernameExists, emailExists}) { - return - } - - c.JSON(http.StatusOK, response.OK(toUser(newUser))) -} - -// updateUserRequest 更新用户信息请求 -type updateUserRequest struct { - Nickname string `json:"nickname" binding:"max=64"` - Email string `json:"email" binding:"required,email,max=255"` - IsAdmin bool `json:"is_admin"` - Password string `json:"password" binding:"omitempty,min=8,max=64"` -} - -// UpdateUser 更新用户信息 -// @Summary 更新用户信息 -// @Description 更新指定用户的昵称、邮箱、管理员权限,并可选重置密码,需要管理员权限 -// @Tags admin -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param id path int true "用户 ID" -// @Param request body user.updateUserRequest true "更新参数" -// @Success 200 {object} response.Any{data=string} "更新成功" -// @Failure 400 {object} response.Any "参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Failure 403 {object} response.Any "无管理员权限或尝试修改自身权限" -// @Failure 404 {object} response.Any "用户不存在" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/admin/users/{id} [put] -func UpdateUser(c *gin.Context) { - var req updateUserRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - id, ok := parseUserID(c) - if !ok { - return - } - - currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) - err := updateUser(c.Request.Context(), currUser.ID, updateUserParam{ - ID: id, - Nickname: req.Nickname, - Email: req.Email, - IsAdmin: req.IsAdmin, - Password: req.Password, - }) - - if err != nil { - if abortUserLogicError(c, err, userNotFound, []string{cannotRevokeSelfAdmin}, []string{emailRequired, emailExists, passwordTooShort}) { - return - } - response.AbortInternal(c, updateUserInfoFailed) - return - } - - c.JSON(http.StatusOK, response.OKNil()) -} diff --git a/internal/apps/admin/user/routers_test.go b/internal/apps/admin/user/routers_test.go deleted file mode 100644 index 91e8a880..00000000 --- a/internal/apps/admin/user/routers_test.go +++ /dev/null @@ -1,582 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package user - -import ( - "bytes" - "encoding/json" - "net/http" - "net/http/httptest" - "testing" - "time" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -func setupTestRouter(authUser *model.User) *gin.Engine { - r := testhelper.NewTestGinEngine() - adminGroup := r.Group("/api/v1/admin") - - // Mock authentication middleware - adminGroup.Use(func(c *gin.Context) { - if authUser != nil { - oauth.SetToContext(c, oauth.UserObjKey, authUser) - } - c.Next() - }) - - adminGroup.GET("/users", ListUsers) - adminGroup.POST("/users", CreateUser) - adminGroup.GET("/users/:id", GetUser) - adminGroup.PUT("/users/:id/status", UpdateUserStatus) - adminGroup.DELETE("/users/:id", DeleteUser) - return r -} - -func TestListUsers(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - // Seed users - users := []model.User{ - { - ID: 1001, - Username: "alice", - Nickname: "Alice Nickname", - IsActive: true, - IsAdmin: false, - LastLoginAt: time.Now(), - }, - { - ID: 1002, - Username: "bob", - Nickname: "Bob Nickname", - IsActive: true, - IsAdmin: false, - LastLoginAt: time.Now(), - }, - { - ID: 1003, - Username: "charlie", - Nickname: "Charlie Nickname", - IsActive: false, - IsAdmin: true, - LastLoginAt: time.Now(), - }, - } - - for _, u := range users { - if err := dbConn.Create(&u).Error; err != nil { - t.Fatalf("failed to seed user: %v", err) - } - } - - adminUser := &model.User{ID: 1003, Username: "charlie", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("basic pagination list", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/users?page=1&page_size=2", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - var resp response.Any - if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil { - t.Fatalf("failed to parse response: %v", err) - } - - // Parse data map to our structure - dataBytes, _ := json.Marshal(resp.Data) - var listResp listUsersResponse - if err := json.Unmarshal(dataBytes, &listResp); err != nil { - t.Fatalf("failed to parse list response: %v", err) - } - - if len(listResp.Users) != 2 { - t.Errorf("expected 2 users, got %d", len(listResp.Users)) - } - if listResp.Total != 3 { - t.Errorf("expected total 3, got %d", listResp.Total) - } - // Ordered by ID ASC - if listResp.Users[0].ID != 1001 || listResp.Users[1].ID != 1002 { - t.Errorf("expected ordered ASC, got first ID %d, second ID %d", listResp.Users[0].ID, listResp.Users[1].ID) - } - }) - - t.Run("filter by user_id", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/users?page=1&page_size=10&user_id=1001", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - var resp response.Any - _ = json.Unmarshal(w.Body.Bytes(), &resp) - - dataBytes, _ := json.Marshal(resp.Data) - var listResp listUsersResponse - _ = json.Unmarshal(dataBytes, &listResp) - - if len(listResp.Users) != 1 || listResp.Users[0].ID != 1001 { - t.Errorf("expected 1 user with ID 1001, got total %d", len(listResp.Users)) - } - }) - - t.Run("filter by username prefix", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/users?page=1&page_size=10&username=bo", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - var resp response.Any - _ = json.Unmarshal(w.Body.Bytes(), &resp) - - dataBytes, _ := json.Marshal(resp.Data) - var listResp listUsersResponse - _ = json.Unmarshal(dataBytes, &listResp) - - if len(listResp.Users) != 1 || listResp.Users[0].Username != "bob" { - t.Errorf("expected bob, got %v", listResp.Users) - } - }) - - t.Run("invalid pagination parameter", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/users?page=0&page_size=10", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusBadRequest { - t.Errorf("expected 400 Bad Request, got %d", w.Code) - } - }) -} - -func TestGetUser(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - targetUser := model.User{ - ID: 1001, - Username: "alice", - Password: "secret-hash", - Nickname: "Alice Nickname", - Email: "alice@example.com", - AvatarURL: "https://example.com/avatar.png", - IsActive: true, - IsAdmin: false, - Bio: "hello", - Phone: "123456", - Gender: "female", - Website: "https://example.com", - Location: "Shanghai", - } - if err := dbConn.Create(&targetUser).Error; err != nil { - t.Fatalf("failed to seed user: %v", err) - } - - adminUser := &model.User{ID: 1003, Username: "charlie", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("get full user profile", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/users/1001", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - var resp response.Any - if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil { - t.Fatalf("failed to parse response: %v", err) - } - - dataBytes, _ := json.Marshal(resp.Data) - var resUser user - if err := json.Unmarshal(dataBytes, &resUser); err != nil { - t.Fatalf("failed to parse response data: %v", err) - } - - if resUser.Email != targetUser.Email || resUser.Bio != targetUser.Bio || resUser.Phone != targetUser.Phone || - resUser.Gender != targetUser.Gender || resUser.Website != targetUser.Website || resUser.Location != targetUser.Location { - t.Errorf("profile fields were not returned correctly: %+v", resUser) - } - if bytes.Contains(dataBytes, []byte("secret-hash")) { - t.Error("response should not include password") - } - }) - - t.Run("get non-existent user", func(t *testing.T) { - req, _ := http.NewRequest("GET", "/api/v1/admin/users/9999", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusNotFound { - t.Errorf("expected 404 Not Found, got %d. Body: %s", w.Code, w.Body.String()) - } - }) -} - -func TestUpdateUserStatus(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - // Seed users - regularUser := model.User{ - ID: 1001, - Username: "alice", - IsActive: true, - IsAdmin: false, - } - adminUser := model.User{ - ID: 1002, - Username: "bob", - IsActive: true, - IsAdmin: true, - } - - dbConn.Create(®ularUser) - dbConn.Create(&adminUser) - - router := setupTestRouter(&adminUser) - - t.Run("disable regular user successfully", func(t *testing.T) { - payload := updateUserStatusRequest{IsActive: false} - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("PUT", "/api/v1/admin/users/1001/status", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - // Verify DB status - var u model.User - dbConn.First(&u, 1001) - if u.IsActive { - t.Error("user should be deactivated in the database") - } - }) - - t.Run("cannot disable admin user", func(t *testing.T) { - payload := updateUserStatusRequest{IsActive: false} - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("PUT", "/api/v1/admin/users/1002/status", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusForbidden { - t.Errorf("expected 403 Forbidden, got %d. Body: %s", w.Code, w.Body.String()) - } - - var resp response.Any - _ = json.Unmarshal(w.Body.Bytes(), &resp) - if resp.ErrorMsg != cannotDisable { - t.Errorf("expected error message '%s', got '%s'", cannotDisable, resp.ErrorMsg) - } - }) - - t.Run("cannot enable/disable non-existent user", func(t *testing.T) { - payload := updateUserStatusRequest{IsActive: false} - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("PUT", "/api/v1/admin/users/9999/status", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusNotFound { - t.Errorf("expected 404 Not Found, got %d. Body: %s", w.Code, w.Body.String()) - } - }) -} - -func TestCreateUser(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - adminUser := &model.User{ID: 1003, Username: "charlie", IsAdmin: true} - router := setupTestRouter(adminUser) - - t.Run("create user successfully", func(t *testing.T) { - payload := createUserRequest{ - Username: "newuser", - Password: "newpassword123", - Nickname: "New Nickname", - Email: "newuser@example.com", - IsActive: true, - IsAdmin: false, - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("POST", "/api/v1/admin/users", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - var resp response.Any - if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil { - t.Fatalf("failed to parse response: %v", err) - } - - if resp.ErrorMsg != "" { - t.Errorf("expected empty error message, got '%s'", resp.ErrorMsg) - } - - dataBytes, _ := json.Marshal(resp.Data) - var resUser user - if err := json.Unmarshal(dataBytes, &resUser); err != nil { - t.Fatalf("failed to parse response data: %v", err) - } - - if resUser.Username != "newuser" || resUser.Nickname != "New Nickname" || !resUser.IsActive || resUser.IsAdmin { - t.Errorf("unexpected user values: %+v", resUser) - } - - // Verify in DB - var dbUser model.User - if err := dbConn.Where("username = ?", "newuser").First(&dbUser).Error; err != nil { - t.Fatalf("failed to find user in db: %v", err) - } - if dbUser.Email != "newuser@example.com" { - t.Errorf("expected email 'newuser@example.com', got '%s'", dbUser.Email) - } - if !dbUser.CheckPassword("newpassword123") { - t.Error("password was not hashed correctly") - } - }) - - t.Run("create user with duplicate username", func(t *testing.T) { - // Create the first user - existing := model.User{ - ID: 2001, - Username: "dupuser", - Nickname: "Dup User", - Email: "dupuser@example.com", - } - dbConn.Create(&existing) - - payload := createUserRequest{ - Username: "dupuser", - Password: "password123", - Nickname: "Another Nick", - Email: "another@example.com", - IsActive: true, - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("POST", "/api/v1/admin/users", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusBadRequest { - t.Errorf("expected 400 Bad Request, got %d. Body: %s", w.Code, w.Body.String()) - } - - var resp response.Any - _ = json.Unmarshal(w.Body.Bytes(), &resp) - if resp.ErrorMsg != usernameExists { - t.Errorf("expected error '%s', got '%s'", usernameExists, resp.ErrorMsg) - } - }) - - t.Run("create user with duplicate email", func(t *testing.T) { - existing := model.User{ - ID: 2002, - Username: "existingemail", - Nickname: "Existing Email", - Email: "dupemail@example.com", - } - dbConn.Create(&existing) - - payload := createUserRequest{ - Username: "newuser2", - Password: "password123", - Nickname: "New User 2", - Email: "dupemail@example.com", - IsActive: true, - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("POST", "/api/v1/admin/users", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusBadRequest { - t.Errorf("expected 400 Bad Request, got %d. Body: %s", w.Code, w.Body.String()) - } - - var resp response.Any - _ = json.Unmarshal(w.Body.Bytes(), &resp) - if resp.ErrorMsg != emailExists { - t.Errorf("expected error '%s', got '%s'", emailExists, resp.ErrorMsg) - } - }) - - t.Run("validation error - password too short", func(t *testing.T) { - payload := createUserRequest{ - Username: "shortpass", - Password: "123", - Email: "shortpass@example.com", - IsActive: true, - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("POST", "/api/v1/admin/users", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusBadRequest { - t.Errorf("expected 400 Bad Request, got %d. Body: %s", w.Code, w.Body.String()) - } - }) - - t.Run("validation error - invalid email format", func(t *testing.T) { - payload := map[string]interface{}{ - "username": "bademail", - "password": "password123", - "email": "not-an-email", - "is_active": true, - } - body, _ := json.Marshal(payload) - req, _ := http.NewRequest("POST", "/api/v1/admin/users", bytes.NewBuffer(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusBadRequest { - t.Errorf("expected 400 Bad Request, got %d. Body: %s", w.Code, w.Body.String()) - } - }) -} - -func TestDeleteUser(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - if err := dbConn.AutoMigrate(&model.AccessToken{}, &model.ExternalAccount{}); err != nil { - t.Fatalf("failed to migrate delete-related tables: %v", err) - } - - regularUser := model.User{ - ID: 1001, - Username: "alice", - IsActive: true, - IsAdmin: false, - } - adminUser := model.User{ - ID: 1002, - Username: "bob", - IsActive: true, - IsAdmin: true, - } - selfUser := model.User{ - ID: 1003, - Username: "charlie", - IsActive: true, - IsAdmin: false, - } - - if err := dbConn.Create(®ularUser).Error; err != nil { - t.Fatalf("failed to seed regular user: %v", err) - } - if err := dbConn.Create(&adminUser).Error; err != nil { - t.Fatalf("failed to seed admin user: %v", err) - } - if err := dbConn.Create(&selfUser).Error; err != nil { - t.Fatalf("failed to seed self user: %v", err) - } - if err := dbConn.Create(&model.AccessToken{ - UserID: regularUser.ID, - Name: "api", - TokenHash: "hash-for-delete-user-test", - MaskedToken: "at_****test", - }).Error; err != nil { - t.Fatalf("failed to seed access token: %v", err) - } - if err := dbConn.Create(&model.ExternalAccount{ - ID: 5001, - AuthSourceID: 1, - UserID: regularUser.ID, - ExternalID: "external-alice", - }).Error; err != nil { - t.Fatalf("failed to seed external account: %v", err) - } - - router := setupTestRouter(&selfUser) - - t.Run("delete regular user successfully", func(t *testing.T) { - req, _ := http.NewRequest("DELETE", "/api/v1/admin/users/1001", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - var count int64 - if err := dbConn.Model(&model.User{}).Where("id = ?", 1001).Count(&count).Error; err != nil { - t.Fatalf("failed to count deleted user: %v", err) - } - if count != 0 { - t.Errorf("expected deleted user count 0, got %d", count) - } - - if err := dbConn.Model(&model.AccessToken{}).Where("user_id = ?", 1001).Count(&count).Error; err != nil { - t.Fatalf("failed to count deleted access tokens: %v", err) - } - if count != 0 { - t.Errorf("expected deleted access token count 0, got %d", count) - } - - if err := dbConn.Model(&model.ExternalAccount{}).Where("user_id = ?", 1001).Count(&count).Error; err != nil { - t.Fatalf("failed to count deleted external accounts: %v", err) - } - if count != 0 { - t.Errorf("expected deleted external account count 0, got %d", count) - } - }) - - t.Run("cannot delete admin user", func(t *testing.T) { - req, _ := http.NewRequest("DELETE", "/api/v1/admin/users/1002", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusForbidden { - t.Fatalf("expected 403 Forbidden, got %d. Body: %s", w.Code, w.Body.String()) - } - }) - - t.Run("cannot delete current user", func(t *testing.T) { - req, _ := http.NewRequest("DELETE", "/api/v1/admin/users/1003", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusForbidden { - t.Fatalf("expected 403 Forbidden, got %d. Body: %s", w.Code, w.Body.String()) - } - }) - - t.Run("delete non-existent user", func(t *testing.T) { - req, _ := http.NewRequest("DELETE", "/api/v1/admin/users/9999", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusNotFound { - t.Errorf("expected 404 Not Found, got %d. Body: %s", w.Code, w.Body.String()) - } - }) -} diff --git a/internal/apps/cap/manager_test.go b/internal/apps/cap/manager_test.go deleted file mode 100644 index 4bf61c35..00000000 --- a/internal/apps/cap/manager_test.go +++ /dev/null @@ -1,162 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package cap - -import ( - "context" - "sync" - "sync/atomic" - "testing" - "time" - - pkgcap "github.com/Rain-kl/Wavelet/pkg/cap" -) - -func installTestManagerSettings(t *testing.T) func() { - t.Helper() - return InstallTestRuntimeSettings(RuntimeSettings{ - ChallengeCount: 3, - ChallengeSize: 32, - ChallengeDifficulty: 3, - ChallengeTTL: 5 * time.Second, - TokenTTL: 10 * time.Second, - }) -} - -func TestCapFullFlow(t *testing.T) { - cleanup := installTestManagerSettings(t) - defer cleanup() - - secret := []byte("a-very-long-secret-key-at-least-16-bytes") - store := pkgcap.NewMemoryStore(1 * time.Minute) - manager := NewManager(secret, store) - - scope := "test-scope" - ctx := context.Background() - resp, err := manager.Generate(ctx, scope) - if err != nil { - t.Fatalf("Generate() error = %v", err) - } - - if resp.Challenge.C != 3 { - t.Fatalf("Generate().Challenge.C = %d, want %d", resp.Challenge.C, 3) - } - - solutions := pkgcap.Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D) - - redeemResp, err := manager.Redeem(ctx, resp.Token, solutions, scope) - if err != nil { - t.Fatalf("Redeem() error = %v", err) - } - if !redeemResp.Success { - t.Fatalf("Redeem().Success = false, error = %s", redeemResp.Error) - } - if redeemResp.Token == "" { - t.Fatal("Redeem().Token is empty") - } - - valid, err := manager.VerifyToken(ctx, redeemResp.Token, scope) - if err != nil { - t.Fatalf("VerifyToken() error = %v", err) - } - if !valid { - t.Fatal("VerifyToken() = false, want true") - } - - validAgain, err := manager.VerifyToken(ctx, redeemResp.Token, scope) - if err != nil { - t.Fatalf("VerifyToken() second call error = %v", err) - } - if validAgain { - t.Fatal("VerifyToken() second call = true, want false") - } -} - -func TestRedeemConcurrentRace(t *testing.T) { - const goroutines = 50 - - cleanup := installTestManagerSettings(t) - defer cleanup() - - secret := []byte("race-test-secret-key-at-least-16-bytes") - store := pkgcap.NewMemoryStore(1 * time.Minute) - manager := NewManager(secret, store) - - ctx := context.Background() - resp, err := manager.Generate(ctx, "login") - if err != nil { - t.Fatalf("Generate() error = %v", err) - } - solutions := pkgcap.Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D) - - var ( - wg sync.WaitGroup - success atomic.Int32 - barrier = make(chan struct{}) - ) - - for range goroutines { - wg.Add(1) - go func() { - defer wg.Done() - <-barrier - r, _ := manager.Redeem(ctx, resp.Token, solutions, "login") - if r != nil && r.Success { - success.Add(1) - } - }() - } - close(barrier) - wg.Wait() - - if got := success.Load(); got != 1 { - t.Fatalf("successful Redeem count = %d, want %d", got, 1) - } -} - -func TestVerifyTokenConcurrentRace(t *testing.T) { - const goroutines = 50 - - cleanup := installTestManagerSettings(t) - defer cleanup() - - secret := []byte("race-test-secret-key-at-least-16-bytes") - store := pkgcap.NewMemoryStore(1 * time.Minute) - manager := NewManager(secret, store) - - ctx := context.Background() - resp, err := manager.Generate(ctx, "login") - if err != nil { - t.Fatalf("Generate() error = %v", err) - } - solutions := pkgcap.Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D) - redeemResp, err := manager.Redeem(ctx, resp.Token, solutions, "login") - if err != nil || !redeemResp.Success { - t.Fatalf("Redeem() error = %v, resp = %+v", err, redeemResp) - } - - var ( - wg sync.WaitGroup - success atomic.Int32 - barrier = make(chan struct{}) - ) - - for range goroutines { - wg.Add(1) - go func() { - defer wg.Done() - <-barrier - ok, _ := manager.VerifyToken(ctx, redeemResp.Token, "login") - if ok { - success.Add(1) - } - }() - } - close(barrier) - wg.Wait() - - if got := success.Load(); got != 1 { - t.Fatalf("successful VerifyToken count = %d, want %d", got, 1) - } -} diff --git a/internal/apps/cap/routers_test.go b/internal/apps/cap/routers_test.go deleted file mode 100644 index 3a0c66ae..00000000 --- a/internal/apps/cap/routers_test.go +++ /dev/null @@ -1,158 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package cap - -import ( - "bytes" - "context" - "encoding/json" - "net/http" - "net/http/httptest" - "sync" - "testing" - - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/infra/config" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/shared/response" - "github.com/Rain-kl/Wavelet/internal/testhelper" - pkgcap "github.com/Rain-kl/Wavelet/pkg/cap" -) - -func decodeAPIResponse[T any](t *testing.T, body []byte) T { - t.Helper() - var envelope struct { - ErrorMsg string `json:"error_msg"` - Data T `json:"data"` - } - if err := json.Unmarshal(body, &envelope); err != nil { - t.Fatalf("failed to unmarshal API envelope: %v", err) - } - if envelope.ErrorMsg != "" { - t.Fatalf("unexpected API error_msg: %s", envelope.ErrorMsg) - } - return envelope.Data -} - -func TestCapEndpointsAndMiddleware(t *testing.T) { - sqliteDB, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - oldSecret := config.Config.App.SessionSecret - config.Config.App.SessionSecret = "test-captcha-session-secret" - once = sync.Once{} - defaultManager = nil - t.Cleanup(func() { - config.Config.App.SessionSecret = oldSecret - once = sync.Once{} - defaultManager = nil - }) - - r := testhelper.NewTestGinEngine() - - // Mount CAPTCHA API endpoints - capGroup := r.Group("/api/cap") - { - capGroup.POST("/challenge", Challenge) - capGroup.POST("/redeem", Redeem) - } - - r.POST("/api/v1/user/login", VerifyMiddleware(GetDefaultManager(), "login"), func(c *gin.Context) { - c.JSON(http.StatusOK, response.OK("login success")) - }) - - // Ensure CAPTCHA is disabled initially for step 2 - if err := sqliteDB.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyCapLoginEnabled).Update("value", "false").Error; err != nil { - t.Fatalf("failed to disable cap_login_enabled in DB: %v", err) - } - if err := repository.InvalidateSystemConfigCache(context.Background(), model.ConfigKeyCapLoginEnabled); err != nil { - t.Fatalf("InvalidateSystemConfigCache() error = %v", err) - } - InvalidateRuntimeSettings() - - // 1. Test challenge generation - w := httptest.NewRecorder() - req, _ := http.NewRequest("POST", "/api/cap/challenge", nil) - r.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) - } - - challengeResp := decodeAPIResponse[pkgcap.ChallengeResponse](t, w.Body.Bytes()) - - if challengeResp.Token == "" { - t.Fatalf("expected token in challenge response") - } - - // 2. Test login with CAPTCHA disabled (should pass) - w = httptest.NewRecorder() - req, _ = http.NewRequest("POST", "/api/v1/user/login", nil) - r.ServeHTTP(w, req) - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK when CAPTCHA is disabled, got %d. Body: %s", w.Code, w.Body.String()) - } - - // 3. Enable CAPTCHA in DB and invalidate runtime snapshot - err := sqliteDB.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyCapLoginEnabled).Update("value", "true").Error - if err != nil { - t.Fatalf("failed to enable cap_login_enabled in DB: %v", err) - } - if err := repository.InvalidateSystemConfigCache(context.Background(), model.ConfigKeyCapLoginEnabled); err != nil { - t.Fatalf("InvalidateSystemConfigCache() error = %v", err) - } - InvalidateRuntimeSettings() - - // 4. Test login with CAPTCHA enabled but no header (should be blocked) - w = httptest.NewRecorder() - req, _ = http.NewRequest("POST", "/api/v1/user/login", nil) - r.ServeHTTP(w, req) - if w.Code != http.StatusUnauthorized { - t.Fatalf("expected 401 Unauthorized, got %d. Body: %s", w.Code, w.Body.String()) - } - - // 5. Solve the challenge - solutions := pkgcap.Solve(challengeResp.Token, challengeResp.Challenge.C, challengeResp.Challenge.S, challengeResp.Challenge.D) - - // 6. Redeem solutions - redeemReqPayload := redeemRequest{ - Token: challengeResp.Token, - Solutions: solutions, - } - bodyBytes, _ := json.Marshal(redeemReqPayload) - w = httptest.NewRecorder() - req, _ = http.NewRequest("POST", "/api/cap/redeem", bytes.NewBuffer(bodyBytes)) - req.Header.Set("Content-Type", "application/json") - r.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK for redeem, got %d. Body: %s", w.Code, w.Body.String()) - } - - redeemResp := decodeAPIResponse[RedeemResponse](t, w.Body.Bytes()) - - if !redeemResp.Success || redeemResp.Token == "" { - t.Fatalf("redeem failed or returned empty token: %+v", redeemResp) - } - - // 7. Login with valid redeem token (should pass) - w = httptest.NewRecorder() - req, _ = http.NewRequest("POST", "/api/v1/user/login", nil) - req.Header.Set("X-Cap-Token", redeemResp.Token) - r.ServeHTTP(w, req) - if w.Code != http.StatusOK { - t.Fatalf("expected 200 OK with valid cap token, got %d. Body: %s", w.Code, w.Body.String()) - } - - // 8. Replay attack: Login with the same redeem token again (should be blocked as it is single-use) - w = httptest.NewRecorder() - req, _ = http.NewRequest("POST", "/api/v1/user/login", nil) - req.Header.Set("X-Cap-Token", redeemResp.Token) - r.ServeHTTP(w, req) - if w.Code != http.StatusUnauthorized { - t.Fatalf("expected 401 Unauthorized on replayed token, got %d. Body: %s", w.Code, w.Body.String()) - } -} diff --git a/internal/apps/cap/runtime_settings_test.go b/internal/apps/cap/runtime_settings_test.go deleted file mode 100644 index 29f0177c..00000000 --- a/internal/apps/cap/runtime_settings_test.go +++ /dev/null @@ -1,119 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package cap - -import ( - "context" - "testing" - "time" - - "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/testhelper" -) - -func TestCurrentSettingsLoadsSnapshotOnce(t *testing.T) { - _, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - ctx := context.Background() - - ResetRuntimeSettingsForTest() - repository.ResetSystemConfigRAMCacheForTest() - - first, err := CurrentSettings(ctx) - if err != nil { - t.Fatalf("CurrentSettings() first error = %v", err) - } - if first.ChallengeCount != 1 { - t.Fatalf("CurrentSettings().ChallengeCount = %d, want %d", first.ChallengeCount, 1) - } - - if err := db.DB(ctx).Model(&model.SystemConfig{}). - Where("key = ?", model.ConfigKeyCapChallengeCount). - Update("value", "4").Error; err != nil { - t.Fatalf("Update(cap_challenge_count) error = %v", err) - } - if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeyCapChallengeCount); err != nil { - t.Fatalf("InvalidateSystemConfigCache() error = %v", err) - } - InvalidateRuntimeSettings() - - second, err := CurrentSettings(ctx) - if err != nil { - t.Fatalf("CurrentSettings() second error = %v", err) - } - if second.ChallengeCount != 4 { - t.Fatalf("CurrentSettings().ChallengeCount = %d, want %d", second.ChallengeCount, 4) - } -} - -func TestProtectionEnabledReflectsLoginSwitch(t *testing.T) { - _, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - ctx := context.Background() - - ResetRuntimeSettingsForTest() - - if ProtectionEnabled(ctx) { - t.Fatal("ProtectionEnabled() = true, want false from seed defaults") - } - - if err := db.DB(ctx).Model(&model.SystemConfig{}). - Where("key = ?", model.ConfigKeyCapLoginEnabled). - Update("value", "true").Error; err != nil { - t.Fatalf("Update(cap_login_enabled) error = %v", err) - } - if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeyCapLoginEnabled); err != nil { - t.Fatalf("InvalidateSystemConfigCache() error = %v", err) - } - InvalidateRuntimeSettings() - - if !ProtectionEnabled(ctx) { - t.Fatal("ProtectionEnabled() = false, want true after config update") - } -} - -func TestParseRuntimeSettingsUsesDefaultsForMissingKeys(t *testing.T) { - settings := parseRuntimeSettings(map[string]model.SystemConfig{}) - - if settings.ChallengeCount != defaultChallengeCount { - t.Fatalf("ChallengeCount = %d, want %d", settings.ChallengeCount, defaultChallengeCount) - } - if settings.ChallengeTTL != defaultChallengeTTL { - t.Fatalf("ChallengeTTL = %s, want %s", settings.ChallengeTTL, defaultChallengeTTL) - } - if settings.TokenTTL != defaultTokenTTL { - t.Fatalf("TokenTTL = %s, want %s", settings.TokenTTL, defaultTokenTTL) - } -} - -func TestIsRuntimeConfigKey(t *testing.T) { - if !IsRuntimeConfigKey(model.ConfigKeyCapChallengeCount) { - t.Fatalf("IsRuntimeConfigKey(%s) = false, want true", model.ConfigKeyCapChallengeCount) - } - if IsRuntimeConfigKey(model.ConfigKeySiteName) { - t.Fatalf("IsRuntimeConfigKey(%s) = true, want false", model.ConfigKeySiteName) - } -} - -func TestInstallTestRuntimeSettings(t *testing.T) { - cleanup := InstallTestRuntimeSettings(RuntimeSettings{ - LoginEnabled: true, - ChallengeCount: 2, - TokenTTL: 30 * time.Minute, - }) - defer cleanup() - - settings, err := CurrentSettings(context.Background()) - if err != nil { - t.Fatalf("CurrentSettings() error = %v", err) - } - if !settings.LoginEnabled { - t.Fatal("LoginEnabled = false, want true") - } - if settings.ChallengeCount != 2 { - t.Fatalf("ChallengeCount = %d, want %d", settings.ChallengeCount, 2) - } -} diff --git a/internal/apps/cap/testhelper_hook.go b/internal/apps/cap/testhelper_hook.go deleted file mode 100644 index 61eaaaf6..00000000 --- a/internal/apps/cap/testhelper_hook.go +++ /dev/null @@ -1,10 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package cap - -import "github.com/Rain-kl/Wavelet/internal/testhelper" - -func init() { - testhelper.RegisterCleanup(ResetRuntimeSettingsForTest) -} diff --git a/internal/apps/config/public_config_cache_test.go b/internal/apps/config/public_config_cache_test.go deleted file mode 100644 index b1770341..00000000 --- a/internal/apps/config/public_config_cache_test.go +++ /dev/null @@ -1,77 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package config - -import ( - "context" - "testing" - "time" - - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/testhelper" -) - -func TestListVisibleSystemConfigsUsesStoreCache(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - ctx := context.Background() - - repository.ResetSystemConfigRAMCacheForTest() - if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil { - t.Fatalf("InvalidateVisibleSystemConfigsCache() error = %v", err) - } - - // Warm cache - if _, err := repository.ListVisibleSystemConfigs(ctx); err != nil { - t.Fatalf("ListVisibleSystemConfigs() warm cache error = %v", err) - } - - // Directly insert a new config in DB (bypassing caching layer) - if err := dbConn.Create(&model.SystemConfig{ - Key: "cache_probe_public_key", - Value: "cache_probe_public_value", - Type: "system", - Visibility: model.ConfigVisibilityVisible, - Description: "cache probe", - }).Error; err != nil { - t.Fatalf("Create(cache_probe_public_key) error = %v", err) - } - - // Cached call: shouldn't return the new key yet - cached, err := repository.ListVisibleSystemConfigs(ctx) - if err != nil { - t.Fatalf("ListVisibleSystemConfigs() cached call error = %v", err) - } - for _, item := range cached { - if item.Key == "cache_probe_public_key" { - t.Fatal("cached visible config list should be stale before invalidation") - } - } - - // Invalidate: triggers reload - if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil { - t.Fatalf("InvalidateVisibleSystemConfigsCache() error = %v", err) - } - - // Wait for Pub/Sub delivery in test environment - time.Sleep(100 * time.Millisecond) - - // Refreshed call: should return the new key - refreshed, err := repository.ListVisibleSystemConfigs(ctx) - if err != nil { - t.Fatalf("ListVisibleSystemConfigs() refreshed call error = %v", err) - } - - var found bool - for _, item := range refreshed { - if item.Key == "cache_probe_public_key" { - found = true - break - } - } - if !found { - t.Fatal("refreshed visible config list should include newly created public config") - } -} diff --git a/internal/apps/config/routers.go b/internal/apps/config/routers.go deleted file mode 100644 index 031b0080..00000000 --- a/internal/apps/config/routers.go +++ /dev/null @@ -1,57 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package config 提供公开配置查询接口 -package config - -import ( - "net/http" - - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -// GetPublicConfig 获取公共配置 -// @Summary 获取公共配置 -// @Description 返回系统配置表中 visibility 为 1 的配置键值集合 -// @Tags config -// @Accept json -// @Produce json -// @Success 200 {object} response.Any -// @Router /api/v1/config/public [get] -func GetPublicConfig(c *gin.Context) { - ctx := c.Request.Context() - configs, err := repository.ListVisibleSystemConfigs(ctx) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - - resp := make(map[string]string, len(configs)) - for _, config := range configs { - resp[config.Key] = config.Value - } - - c.JSON(http.StatusOK, response.OK(resp)) -} - -// GetRobotsTXT 动态生成 robots.txt -// @Summary 获取 robots.txt -// @Description 根据系统配置决定是否允许搜索引擎检索,并返回相应的 robots.txt 文件内容 -// @Tags config -// @Produce text/plain -// @Success 200 {string} string "robots.txt 内容" -// @Router /robots.txt [get] -func GetRobotsTXT(c *gin.Context) { - ctx := c.Request.Context() - enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeySearchEngineIndexingEnabled) - content := "User-Agent: *\nDisallow: /\n" - if err == nil && enabled { - content = "User-Agent: *\nAllow: /\n" - } - c.Data(http.StatusOK, "text/plain; charset=utf-8", []byte(content)) -} diff --git a/internal/apps/config/routers_test.go b/internal/apps/config/routers_test.go deleted file mode 100644 index d05aaa25..00000000 --- a/internal/apps/config/routers_test.go +++ /dev/null @@ -1,70 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package config - -import ( - "encoding/json" - "net/http" - "net/http/httptest" - "testing" - - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -func TestGetPublicConfigUsesVisibility(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - if err := dbConn.Create(&model.SystemConfig{ - Key: "custom_public_key", - Value: "custom_public_value", - Type: "system", - Visibility: model.ConfigVisibilityVisible, - Description: "custom public config", - }).Error; err != nil { - t.Fatalf("Create(custom_public_key) error = %v", err) - } - if err := dbConn.Model(&model.SystemConfig{}). - Where("key = ?", model.ConfigKeySiteName). - Update("visibility", model.ConfigVisibilityHidden).Error; err != nil { - t.Fatalf("Update(%s.visibility) error = %v", model.ConfigKeySiteName, err) - } - - gin.SetMode(gin.TestMode) - router := gin.New() - router.GET("/api/v1/config/public", GetPublicConfig) - - req := httptest.NewRequest(http.MethodGet, "/api/v1/config/public", nil) - w := httptest.NewRecorder() - router.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Fatalf("GetPublicConfig() status = %d, want %d; body = %s", w.Code, http.StatusOK, w.Body.String()) - } - - var resp response.Any - if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil { - t.Fatalf("json.Unmarshal(GetPublicConfig()) error = %v", err) - } - dataBytes, err := json.Marshal(resp.Data) - if err != nil { - t.Fatalf("json.Marshal(GetPublicConfig().data) error = %v", err) - } - var configs map[string]string - if err := json.Unmarshal(dataBytes, &configs); err != nil { - t.Fatalf("json.Unmarshal(GetPublicConfig().data) error = %v", err) - } - - if got := configs["custom_public_key"]; got != "custom_public_value" { - t.Errorf("GetPublicConfig()[custom_public_key] = %q, want %q", got, "custom_public_value") - } - if _, ok := configs[model.ConfigKeySiteName]; ok { - t.Errorf("GetPublicConfig()[%s] is present, want hidden", model.ConfigKeySiteName) - } -} diff --git a/internal/apps/config/system_config_cache_test.go b/internal/apps/config/system_config_cache_test.go deleted file mode 100644 index 16719620..00000000 --- a/internal/apps/config/system_config_cache_test.go +++ /dev/null @@ -1,110 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package config - -import ( - "context" - "testing" - "time" - - "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/testhelper" -) - -func TestSystemConfigRAMCacheServesUntilInvalidated(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - ctx := context.Background() - - repository.ResetSystemConfigRAMCacheForTest() - if err := repository.InvalidateAllSystemConfigCaches(ctx); err != nil { - t.Fatalf("InvalidateAllSystemConfigCaches() error = %v", err) - } - - warm, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName) - if err != nil { - t.Fatalf("GetSystemConfigByKey(site_name) warm error = %v", err) - } - if warm.Value != "Wavelet" { - t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", warm.Value, "Wavelet") - } - - // Update DB directly (bypassing caching layer) - if err := dbConn.Model(&model.SystemConfig{}). - Where("key = ?", model.ConfigKeySiteName). - Update("value", "ram_probe_value").Error; err != nil { - t.Fatalf("Update(site_name) error = %v", err) - } - - // Should still return "Wavelet" since it's cached in RAM cache - cached, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName) - if err != nil { - t.Fatalf("GetSystemConfigByKey(site_name) cached error = %v", err) - } - if cached.Value != "Wavelet" { - t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want stale RAM value %q", cached.Value, "Wavelet") - } - - // Invalidate the cache (triggers refresh callback) - if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil { - t.Fatalf("InvalidateSystemConfigCache(site_name) error = %v", err) - } - - // Allow some time for broadcast listener in test environment - time.Sleep(100 * time.Millisecond) - - // Should return the updated value now - refreshed, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName) - if err != nil { - t.Fatalf("GetSystemConfigByKey(site_name) refreshed error = %v", err) - } - if refreshed.Value != "ram_probe_value" { - t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", refreshed.Value, "ram_probe_value") - } -} - -func TestInvalidateSystemConfigCacheBroadcast(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - ctx := context.Background() - - repository.ResetSystemConfigRAMCacheForTest() - - // Initially seed in cache - _, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName) - if err != nil { - t.Fatalf("GetSystemConfigByKey(site_name) error = %v", err) - } - - // Update DB directly - if err := dbConn.Model(&model.SystemConfig{}). - Where("key = ?", model.ConfigKeySiteName). - Update("value", "broadcast_value").Error; err != nil { - t.Fatalf("Update(site_name) error = %v", err) - } - - // Invalidate: publishes to Redis and refreshes locally/other nodes - if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil { - t.Fatalf("InvalidateSystemConfigCache(site_name) error = %v", err) - } - - // Wait for Redis Pub/Sub delivery in test - time.Sleep(100 * time.Millisecond) - - refreshed, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName) - if err != nil { - t.Fatalf("GetSystemConfigByKey(site_name) refreshed error = %v", err) - } - if refreshed.Value != "broadcast_value" { - t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", refreshed.Value, "broadcast_value") - } - - // Verify Redis pub/sub channel received message - if db.Redis != nil { - // Just a sanity check: we can publish a new manual update and verify subscription triggers - // which was done implicitly above. - } -} diff --git a/internal/apps/custom/routers.go b/internal/apps/custom/routers.go deleted file mode 100644 index e33c0839..00000000 --- a/internal/apps/custom/routers.go +++ /dev/null @@ -1,25 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package custom is a scaffold SAMPLE, not a product business home. -// Real domains live in internal/apps// (sibling of oauth, user, upload). -package custom - -import ( - "net/http" - - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -// Hello is a sample handler only — do not grow real product logic in this package. -// @Summary Sample Hello API -// @Description Scaffold demo API; product APIs use semantic paths under apps/ -// @Tags custom -// @Produce json -// @Success 200 {object} response.Any{data=string} "成功" -// @Router /api/v1/custom/hello [get] -func Hello(c *gin.Context) { - c.JSON(http.StatusOK, response.OK("Hello from custom business module!")) -} diff --git a/internal/apps/health/routers.go b/internal/apps/health/routers.go deleted file mode 100644 index d0d26bab..00000000 --- a/internal/apps/health/routers.go +++ /dev/null @@ -1,25 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package health 提供健康检查端点 -package health - -import ( - "net/http" - - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -// Health 健康检查 -// @Summary 健康检查 -// @Description 检查服务是否正常运行,可用于负载均衡存活探测 -// @Tags health -// @Produce json -// @Success 200 {object} response.Any{data=string} "服务正常" -// @Router /api/health [get] -func Health(c *gin.Context) { - c.JSON(http.StatusOK, response.OKNil()) -} diff --git a/internal/apps/message_gateway/errs.go b/internal/apps/message_gateway/errs.go deleted file mode 100644 index feb46b52..00000000 --- a/internal/apps/message_gateway/errs.go +++ /dev/null @@ -1,16 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package message_gateway - -import "errors" - -var ( - errCodeInvalid = errors.New("invalid or expired pairing code") - errChannelMismatch = errors.New("pairing code does not match channel") - errPlatformAlreadyBound = errors.New("this platform account is already bound") - errBindingNotFound = errors.New("binding not found") - errBindingForbidden = errors.New("cannot unbind another user's binding") - errChannelIDRequired = errors.New("channel_id is required") - errChannelDisabled = errors.New("channel is not enabled") -) diff --git a/internal/apps/message_gateway/handlers.go b/internal/apps/message_gateway/handlers.go deleted file mode 100644 index b30470ea..00000000 --- a/internal/apps/message_gateway/handlers.go +++ /dev/null @@ -1,136 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package message_gateway - -import ( - "errors" - "net/http" - "strconv" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/shared/response" - "github.com/gin-gonic/gin" -) - -func currentUser(c *gin.Context) (*model.User, bool) { - return oauth.GetFromContext[*model.User](c, oauth.UserObjKey) -} - -// 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=[]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, "login required") - return - } - rows, err := listEnabledPublicChannels(c.Request.Context()) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(rows)) -} - -// ListBindings lists the current user's bot bindings. -// @Summary List message gateway bindings -// @Description Returns the current user's bound messaging channels -// @Tags message-gateway -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=[]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, "login required") - return - } - rows, err := listUserBindings(c.Request.Context(), user.ID) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(rows)) -} - -// BindBinding consumes a pairing code and binds the platform identity. -// @Summary Bind a messaging channel -// @Description Binds the current user to a platform identity using a one-time pairing code -// @Tags message-gateway -// @Accept json -// @Produce json -// @Security SessionCookie -// @Param request body BindRequest true "bind body" -// @Success 200 {object} response.Any{data=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, "login required") - return - } - var req BindRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - dto, err := bindChannel(c.Request.Context(), user.ID, req) - if err != nil { - if errors.Is(err, errPlatformAlreadyBound) { - response.AbortConflict(c, err.Error()) - return - } - response.AbortBadRequest(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(dto)) -} - -// UnbindBinding removes the current user's binding. -// @Summary Unbind a messaging channel -// @Description Removes a binding owned by the current user -// @Tags message-gateway -// @Produce json -// @Security SessionCookie -// @Param id path int true "binding id" -// @Success 200 {object} response.Any -// @Failure 403 {object} response.Any -// @Failure 404 {object} response.Any -// @Router /api/v1/message-gateway/bindings/{id} [delete] -func UnbindBinding(c *gin.Context) { - user, ok := currentUser(c) - if !ok || user == nil { - response.AbortUnauthorized(c, "login required") - return - } - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, "invalid binding id") - return - } - if err := unbindChannel(c.Request.Context(), user.ID, id); err != nil { - if errors.Is(err, errBindingNotFound) { - response.AbortNotFound(c, err.Error()) - return - } - if errors.Is(err, errBindingForbidden) { - response.AbortForbidden(c, err.Error()) - return - } - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OKNil()) -} diff --git a/internal/apps/message_gateway/logics.go b/internal/apps/message_gateway/logics.go deleted file mode 100644 index 514ed867..00000000 --- a/internal/apps/message_gateway/logics.go +++ /dev/null @@ -1,158 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package message_gateway - -import ( - "context" - "errors" - "strconv" - "strings" - "time" - - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - pkgmg "github.com/Rain-kl/Wavelet/pkg/message_gateway" - "gorm.io/gorm" -) - -// 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"` -} - -func bindChannel(ctx context.Context, userID uint64, req BindRequest) (BindingDTO, error) { - channelID, err := strconv.ParseUint(strings.TrimSpace(req.ChannelID), 10, 64) - if err != nil || channelID == 0 { - return BindingDTO{}, errChannelIDRequired - } - code := pkgmg.NormalizeCode(req.Code) - if code == "" { - return BindingDTO{}, errCodeInvalid - } - pairing, err := repository.GetPairingCode(ctx, code) - if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return BindingDTO{}, errCodeInvalid - } - return BindingDTO{}, err - } - if !pairing.ExpiresAt.After(time.Now()) { - return BindingDTO{}, errCodeInvalid - } - if pairing.ChannelID != channelID { - return BindingDTO{}, errChannelMismatch - } - ch, err := repository.GetMessageChannel(ctx, channelID) - if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return BindingDTO{}, errCodeInvalid - } - return BindingDTO{}, err - } - if !ch.Enabled { - return BindingDTO{}, errChannelDisabled - } - - existing, err := repository.GetBindingByChannelPlatform(ctx, channelID, pairing.PlatformUserID) - if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { - return BindingDTO{}, err - } - if err == nil && existing != nil { - if existing.UserID != userID { - return BindingDTO{}, 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 BindingDTO{}, err - } - if err := repository.DeletePairingCode(ctx, pairing.Code); err != nil { - return BindingDTO{}, err - } - return toBindingDTO(row, ch), nil -} - -// 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"` -} - -func listEnabledPublicChannels(ctx context.Context) ([]PublicChannelDTO, error) { - rows, err := repository.ListEnabledMessageChannels(ctx) - if err != nil { - return nil, err - } - out := make([]PublicChannelDTO, 0, len(rows)) - for _, row := range rows { - out = append(out, PublicChannelDTO{ID: row.ID, Name: row.Name, Type: row.Type}) - } - return out, nil -} - -func listUserBindings(ctx context.Context, userID uint64) ([]BindingDTO, error) { - rows, err := repository.ListBindingsByUser(ctx, userID) - if err != nil { - return nil, err - } - out := make([]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 -} - -func unbindChannel(ctx context.Context, userID, bindingID uint64) error { - row, err := repository.GetMessageBinding(ctx, bindingID) - if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return errBindingNotFound - } - return err - } - if row.UserID != userID { - return errBindingForbidden - } - return repository.DeleteMessageBinding(ctx, bindingID) -} - -func toBindingDTO(row *model.MessageBinding, ch *model.MessageChannel) BindingDTO { - dto := 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 -} diff --git a/internal/apps/message_gateway/logics_test.go b/internal/apps/message_gateway/logics_test.go deleted file mode 100644 index a7d9185b..00000000 --- a/internal/apps/message_gateway/logics_test.go +++ /dev/null @@ -1,106 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package message_gateway - -import ( - "context" - "errors" - "fmt" - "testing" - "time" - - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/testhelper" - "gorm.io/gorm" -) - -func seedChannel(t *testing.T, ctx context.Context) *model.MessageChannel { - t.Helper() - ch := &model.MessageChannel{ - Name: "tg", - Type: model.MessageChannelTypeTelegram, - OwnerScope: model.MessageOwnerScopeSystem, - Enabled: true, - } - if err := repository.CreateMessageChannel(ctx, ch); err != nil { - t.Fatalf("CreateMessageChannel() error = %v", err) - } - return ch -} - -func TestBind_ExpiredCode(t *testing.T) { - _, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - ctx := context.Background() - ch := seedChannel(t, ctx) - if _, err := repository.UpsertPairingCode(ctx, ch.ID, "u1", "ABCD2345", time.Now().Add(-time.Minute)); err != nil { - t.Fatalf("UpsertPairingCode() error = %v", err) - } - _, err := bindChannel(ctx, 1, BindRequest{ChannelID: fmt.Sprint(ch.ID), Code: "ABCD-2345"}) - if err == nil { - t.Fatal("bindChannel() error = nil, want expired code") - } -} - -func TestBind_HappyPathDeletesCode(t *testing.T) { - _, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - ctx := context.Background() - ch := seedChannel(t, ctx) - if _, err := repository.UpsertPairingCode(ctx, ch.ID, "plat-1", "ABCD2345", time.Now().Add(15*time.Minute)); err != nil { - t.Fatalf("UpsertPairingCode() error = %v", err) - } - dto, err := bindChannel(ctx, 42, BindRequest{ChannelID: fmt.Sprint(ch.ID), Code: "abcd-2345"}) - if err != nil { - t.Fatalf("bindChannel() error = %v", err) - } - if dto.PlatformUserID != "plat-1" || dto.ChannelID != ch.ID || dto.UserID != 42 { - t.Fatalf("bindChannel() dto = %+v", dto) - } - _, err = repository.GetPairingCode(ctx, "ABCD2345") - if !errors.Is(err, gorm.ErrRecordNotFound) { - t.Fatalf("GetPairingCode() error = %v, want not found", err) - } -} - -func TestBind_ConflictOtherUser(t *testing.T) { - _, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - ctx := context.Background() - ch := seedChannel(t, ctx) - if err := repository.CreateMessageBinding(ctx, &model.MessageBinding{ - UserID: 7, ChannelID: ch.ID, PlatformUserID: "plat-1", - }); err != nil { - t.Fatalf("CreateMessageBinding() error = %v", err) - } - if _, err := repository.UpsertPairingCode(ctx, ch.ID, "plat-1", "ABCD2345", time.Now().Add(15*time.Minute)); err != nil { - t.Fatalf("UpsertPairingCode() error = %v", err) - } - _, err := bindChannel(ctx, 42, BindRequest{ChannelID: fmt.Sprint(ch.ID), Code: "ABCD2345"}) - if !errors.Is(err, errPlatformAlreadyBound) { - t.Fatalf("bindChannel() error = %v, want %v", err, errPlatformAlreadyBound) - } -} - -func TestUnbind_OnlyOwnBinding(t *testing.T) { - _, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - ctx := context.Background() - ch := seedChannel(t, ctx) - b := &model.MessageBinding{UserID: 7, ChannelID: ch.ID, PlatformUserID: "plat-1"} - if err := repository.CreateMessageBinding(ctx, b); err != nil { - t.Fatalf("CreateMessageBinding() error = %v", err) - } - if err := unbindChannel(ctx, 42, b.ID); err == nil { - t.Fatal("unbindChannel() error = nil, want forbidden") - } - if err := unbindChannel(ctx, 7, b.ID); err != nil { - t.Fatalf("unbindChannel() own binding error = %v", err) - } - _, err := repository.GetMessageBinding(ctx, b.ID) - if !errors.Is(err, gorm.ErrRecordNotFound) { - t.Fatalf("GetMessageBinding() error = %v, want deleted", err) - } -} diff --git a/internal/apps/message_gateway/routers.go b/internal/apps/message_gateway/routers.go deleted file mode 100644 index 00cc4b53..00000000 --- a/internal/apps/message_gateway/routers.go +++ /dev/null @@ -1,22 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package message_gateway provides user bind/unbind APIs and credential helpers. -package message_gateway - -import ( - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/gin-gonic/gin" -) - -// RegisterUserRoutes mounts login-required bind/unbind APIs. -func RegisterUserRoutes(apiV1Router *gin.RouterGroup) { - g := apiV1Router.Group("/message-gateway") - g.Use(oauth.LoginRequired()) - { - g.GET("/channels", ListChannels) - g.GET("/bindings", ListBindings) - g.POST("/bindings", BindBinding) - g.DELETE("/bindings/:id", UnbindBinding) - } -} diff --git a/internal/apps/message_gateway/runner/inbound.go b/internal/apps/message_gateway/runner/inbound.go deleted file mode 100644 index fccd799c..00000000 --- a/internal/apps/message_gateway/runner/inbound.go +++ /dev/null @@ -1,68 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package runner starts message-gateway adapters in the Worker process. -package runner - -import ( - "context" - "errors" - "fmt" - "time" - - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/pkg/message_gateway" - "gorm.io/gorm" -) - -const pairingTTL = 15 * time.Minute - -type inboundDeps struct { - LookupBinding func(ctx context.Context, channelID uint64, platformUserID string) (*model.MessageBinding, error) - UpsertCode func(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*model.MessagePairingCode, error) - GenerateCode func() (string, error) - Emit func(context.Context, message_gateway.InboundMessage) error - Send func(context.Context, message_gateway.Recipient, message_gateway.OutboundMessage) error -} - -// Handle pairs unbound senders or emits bound inbound messages. -func (d inboundDeps) Handle(ctx context.Context, msg message_gateway.InboundMessage) error { - binding, err := d.LookupBinding(ctx, msg.ChannelID, msg.PlatformUserID) - if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { - return err - } - to := message_gateway.Recipient{ChatID: msg.ChatID, PlatformUserID: msg.PlatformUserID} - if binding == nil || errors.Is(err, gorm.ErrRecordNotFound) { - gen := d.GenerateCode - if gen == nil { - gen = message_gateway.GenerateCode - } - code, err := gen() - if err != nil { - return err - } - row, err := d.UpsertCode(ctx, msg.ChannelID, msg.PlatformUserID, code, time.Now().Add(pairingTTL)) - if err != nil { - return err - } - display := message_gateway.FormatCode(row.Code) - return d.Send(ctx, to, message_gateway.OutboundMessage{ - Text: fmt.Sprintf("Your pairing code is %s. Open Settings → Profile → Bind a bot and enter this code. It expires in 15 minutes.", display), - ReplyToID: msg.MessageID, - }) - } - - uid := binding.UserID - msg.BindingUserID = &uid - if err := d.Emit(ctx, msg); err != nil { - _ = d.Send(ctx, to, message_gateway.OutboundMessage{ - Text: "could not save your message", - ReplyToID: msg.MessageID, - }) - return err - } - return d.Send(ctx, to, message_gateway.OutboundMessage{ - Text: "received", - ReplyToID: msg.MessageID, - }) -} diff --git a/internal/apps/message_gateway/runner/inbound_test.go b/internal/apps/message_gateway/runner/inbound_test.go deleted file mode 100644 index 68f00a1e..00000000 --- a/internal/apps/message_gateway/runner/inbound_test.go +++ /dev/null @@ -1,138 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package runner - -import ( - "context" - "errors" - "strings" - "testing" - "time" - - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/pkg/message_gateway" - "gorm.io/gorm" -) - -func TestHandle_UnboundMintsCodeAndDoesNotEmit(t *testing.T) { - var sent []message_gateway.OutboundMessage - var emitted int - var upserted string - d := inboundDeps{ - LookupBinding: func(ctx context.Context, channelID uint64, platformUserID string) (*model.MessageBinding, error) { - return nil, gorm.ErrRecordNotFound - }, - GenerateCode: func() (string, error) { return "ABCD2345", nil }, - UpsertCode: func(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*model.MessagePairingCode, error) { - upserted = code - return &model.MessagePairingCode{Code: code, ChannelID: channelID, PlatformUserID: platformUserID, ExpiresAt: expiresAt}, nil - }, - Emit: func(ctx context.Context, msg message_gateway.InboundMessage) error { - emitted++ - return nil - }, - Send: func(ctx context.Context, to message_gateway.Recipient, msg message_gateway.OutboundMessage) error { - sent = append(sent, msg) - return nil - }, - } - err := d.Handle(context.Background(), message_gateway.InboundMessage{ - ChannelID: 1, PlatformUserID: "u1", ChatID: "u1", Text: "hi", - }) - if err != nil { - t.Fatalf("Handle() error = %v", err) - } - if emitted != 0 { - t.Fatalf("Handle() emitted = %d, want 0", emitted) - } - if upserted != "ABCD2345" { - t.Fatalf("UpsertCode() code = %q, want %q", upserted, "ABCD2345") - } - if len(sent) != 1 { - t.Fatalf("Send() calls = %d, want 1", len(sent)) - } - if !strings.Contains(sent[0].Text, "ABCD-2345") { - t.Fatalf("Handle() send text = %q, want pairing code ABCD-2345", sent[0].Text) - } -} - -func TestHandle_UnboundReusesExistingCode(t *testing.T) { - var sent string - d := inboundDeps{ - LookupBinding: func(ctx context.Context, channelID uint64, platformUserID string) (*model.MessageBinding, error) { - return nil, gorm.ErrRecordNotFound - }, - GenerateCode: func() (string, error) { return "NEWCODE1", nil }, - UpsertCode: func(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*model.MessagePairingCode, error) { - return &model.MessagePairingCode{Code: "OLDCODE2"}, nil - }, - Emit: func(ctx context.Context, msg message_gateway.InboundMessage) error { - t.Fatal("must not emit") - return nil - }, - Send: func(ctx context.Context, to message_gateway.Recipient, msg message_gateway.OutboundMessage) error { - sent = msg.Text - return nil - }, - } - if err := d.Handle(context.Background(), message_gateway.InboundMessage{ChannelID: 1, PlatformUserID: "u1"}); err != nil { - t.Fatalf("Handle() error = %v", err) - } - if !strings.Contains(sent, "OLDC-ODE2") { - t.Fatalf("Handle() send text = %q, want reused code OLDC-ODE2", sent) - } -} - -func TestHandle_BoundEmitsAndAcks(t *testing.T) { - var got message_gateway.InboundMessage - var sent string - d := inboundDeps{ - LookupBinding: func(ctx context.Context, channelID uint64, platformUserID string) (*model.MessageBinding, error) { - return &model.MessageBinding{UserID: 9, ChannelID: 1, PlatformUserID: "u1"}, nil - }, - Emit: func(ctx context.Context, msg message_gateway.InboundMessage) error { - got = msg - return nil - }, - Send: func(ctx context.Context, to message_gateway.Recipient, msg message_gateway.OutboundMessage) error { - sent = msg.Text - return nil - }, - } - err := d.Handle(context.Background(), message_gateway.InboundMessage{ - ChannelID: 1, PlatformUserID: "u1", Text: "hello", - }) - if err != nil { - t.Fatalf("Handle() error = %v", err) - } - if got.Text != "hello" || got.BindingUserID == nil || *got.BindingUserID != 9 { - t.Fatalf("Handle() emit = %+v, want text=hello user=9", got) - } - if sent != "received" { - t.Fatalf("Handle() ack = %q, want %q", sent, "received") - } -} - -func TestHandle_BoundEmitError(t *testing.T) { - var sent string - d := inboundDeps{ - LookupBinding: func(ctx context.Context, channelID uint64, platformUserID string) (*model.MessageBinding, error) { - return &model.MessageBinding{UserID: 9, ChannelID: 1, PlatformUserID: "u1"}, nil - }, - Emit: func(ctx context.Context, msg message_gateway.InboundMessage) error { - return errors.New("listener failed") - }, - Send: func(ctx context.Context, to message_gateway.Recipient, msg message_gateway.OutboundMessage) error { - sent = msg.Text - return nil - }, - } - err := d.Handle(context.Background(), message_gateway.InboundMessage{ChannelID: 1, PlatformUserID: "u1", Text: "x"}) - if err == nil { - t.Fatal("Handle() error = nil, want listener error") - } - if sent != "could not save your message" { - t.Fatalf("Handle() send = %q, want could not save your message", sent) - } -} diff --git a/internal/apps/message_gateway/runner/runner.go b/internal/apps/message_gateway/runner/runner.go deleted file mode 100644 index 4be3fb99..00000000 --- a/internal/apps/message_gateway/runner/runner.go +++ /dev/null @@ -1,239 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package runner - -import ( - "context" - "fmt" - "os" - "sync" - "time" - - appgw "github.com/Rain-kl/Wavelet/internal/apps/message_gateway" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/listener" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/pkg/logger" - "github.com/Rain-kl/Wavelet/pkg/message_gateway" - "github.com/Rain-kl/Wavelet/pkg/message_gateway/channel/qq" - "github.com/Rain-kl/Wavelet/pkg/message_gateway/channel/telegram" -) - -const ( - reloadInterval = 5 * time.Second - lockTTL = 30 * time.Second - lockRefresh = 10 * time.Second -) - -var registerFactoriesOnce sync.Once - -func registerFactories() { - registerFactoriesOnce.Do(func() { - message_gateway.Register(message_gateway.ChannelTypeTelegram, telegram.New) - message_gateway.Register(message_gateway.ChannelTypeQQ, qq.New) - }) -} - -type runningChannel struct { - ch message_gateway.Channel - updatedAt time.Time - cancel context.CancelFunc -} - -// Start loads enabled channels, connects adapters, and reloads on change. -func Start(ctx context.Context) error { - registerFactories() - r := &gateway{ - node: nodeID(), - running: map[uint64]*runningChannel{}, - } - if err := r.sync(ctx); err != nil { - logger.ErrorF(ctx, "message-gateway initial sync: %v", err) - } - ticker := time.NewTicker(reloadInterval) - defer ticker.Stop() - for { - select { - case <-ctx.Done(): - r.stopAll(ctx) - return ctx.Err() - case <-ticker.C: - if err := repository.DeleteExpiredPairingCodes(ctx); err != nil { - logger.WarnF(ctx, "message-gateway expire pairing: %v", err) - } - if err := r.sync(ctx); err != nil { - logger.ErrorF(ctx, "message-gateway sync: %v", err) - } - } - } -} - -type gateway struct { - node string - mu sync.Mutex - running map[uint64]*runningChannel -} - -func (r *gateway) sync(ctx context.Context) error { - rows, err := repository.ListMessageChannels(ctx) - if err != nil { - return err - } - live := make(map[uint64]model.MessageChannel, len(rows)) - for _, row := range rows { - live[row.ID] = row - } - - r.mu.Lock() - defer r.mu.Unlock() - - for id, run := range r.running { - row, ok := live[id] - if !ok || !row.Enabled || !row.UpdatedAt.Equal(run.updatedAt) { - r.stopLocked(ctx, id) - } - } - - for id, row := range live { - if !row.Enabled { - continue - } - if _, ok := r.running[id]; ok { - continue - } - if err := r.startLocked(ctx, row); err != nil { - logger.ErrorF(ctx, "message-gateway start channel %d: %v", id, err) - } - } - return nil -} - -func (r *gateway) startLocked(ctx context.Context, row model.MessageChannel) error { - if !r.acquireLock(ctx, row.ID) { - logger.InfoF(ctx, "message-gateway skip channel %d: lock held", row.ID) - return nil - } - - creds, err := appgw.DecryptCredentials(row.Credentials) - if err != nil { - logger.ErrorF(ctx, "message-gateway decrypt channel %d: %v", row.ID, err) - return nil - } - factory, ok := message_gateway.Lookup(row.Type) - if !ok { - return fmt.Errorf("unknown channel type %q", row.Type) - } - - runCtx, cancel := context.WithCancel(ctx) - var live message_gateway.Channel - deps := inboundDeps{ - LookupBinding: repository.GetBindingByChannelPlatform, - UpsertCode: repository.UpsertPairingCode, - GenerateCode: message_gateway.GenerateCode, - Emit: func(ctx context.Context, msg message_gateway.InboundMessage) error { - listener.EmitMessageGatewayInbound(ctx, msg) - return nil - }, - Send: func(ctx context.Context, to message_gateway.Recipient, msg message_gateway.OutboundMessage) error { - if live == nil { - return fmt.Errorf("qq/telegram channel not ready") - } - return live.Send(ctx, to, msg) - }, - } - ch, err := factory(message_gateway.ChannelConfig{ - ID: row.ID, - Type: row.Type, - Name: row.Name, - Credentials: creds, - Extra: appgw.ParseExtra(row.Extra), - }, func(ctx context.Context, msg message_gateway.InboundMessage) error { - defer func() { - if rec := recover(); rec != nil { - logger.ErrorF(ctx, "message-gateway inbound panic channel %d: %v", row.ID, rec) - } - }() - return deps.Handle(ctx, msg) - }) - if err != nil { - cancel() - return err - } - live = ch - if err := ch.Connect(runCtx); err != nil { - cancel() - _ = ch.Disconnect(ctx) - return err - } - r.running[row.ID] = &runningChannel{ch: ch, updatedAt: row.UpdatedAt, cancel: cancel} - go r.keepLock(runCtx, row.ID) - logger.InfoF(ctx, "message-gateway connected channel %d type=%s", row.ID, row.Type) - return nil -} - -func (r *gateway) stopLocked(ctx context.Context, id uint64) { - run, ok := r.running[id] - if !ok { - return - } - run.cancel() - if err := run.ch.Disconnect(ctx); err != nil { - logger.WarnF(ctx, "message-gateway disconnect channel %d: %v", id, err) - } - delete(r.running, id) - r.releaseLock(ctx, id) -} - -func (r *gateway) stopAll(ctx context.Context) { - r.mu.Lock() - defer r.mu.Unlock() - for id := range r.running { - r.stopLocked(ctx, id) - } -} - -func (r *gateway) acquireLock(ctx context.Context, id uint64) bool { - if db.Redis == nil { - return true - } - ok, err := db.Redis.SetNX(ctx, lockKey(id), r.node, lockTTL).Result() - if err != nil { - logger.WarnF(ctx, "message-gateway lock channel %d: %v", id, err) - return true - } - return ok -} - -func (r *gateway) keepLock(ctx context.Context, id uint64) { - if db.Redis == nil { - return - } - ticker := time.NewTicker(lockRefresh) - defer ticker.Stop() - for { - select { - case <-ctx.Done(): - return - case <-ticker.C: - _ = db.Redis.Expire(ctx, lockKey(id), lockTTL).Err() - } - } -} - -func (r *gateway) releaseLock(ctx context.Context, id uint64) { - if db.Redis == nil { - return - } - _ = db.Redis.Del(ctx, lockKey(id)).Err() -} - -func lockKey(id uint64) string { - return db.PrefixedKey(fmt.Sprintf("wg:channel:%d", id)) -} - -func nodeID() string { - host, _ := os.Hostname() - return fmt.Sprintf("%s:%d", host, os.Getpid()) -} diff --git a/internal/apps/message_gateway/secret.go b/internal/apps/message_gateway/secret.go deleted file mode 100644 index 62f0490e..00000000 --- a/internal/apps/message_gateway/secret.go +++ /dev/null @@ -1,78 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package message_gateway - -import ( - "crypto/sha256" - "encoding/hex" - "encoding/json" - - "github.com/Rain-kl/Wavelet/internal/infra/config" - "github.com/Rain-kl/Wavelet/pkg/util" -) - -// CredentialKey is AES-256 hex derived from the session secret. -func CredentialKey() string { - secret := "" - if config.Config != nil { - secret = config.Config.App.SessionSecret - } - 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) -} diff --git a/internal/apps/oauth/audit.go b/internal/apps/oauth/audit.go deleted file mode 100644 index 5f5fcede..00000000 --- a/internal/apps/oauth/audit.go +++ /dev/null @@ -1,36 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package oauth 提供 OAuth/OIDC 认证与会话管理 -package oauth - -import ( - "context" - "encoding/json" - - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/pkg/logger" - "github.com/gin-gonic/gin" -) - -// LogForAudit 将登录鉴权审计日志写入 Logger -func LogForAudit(ctx context.Context, user *model.User, c *gin.Context) { - auditLog := loginRequiredAuditLog{ - UserID: user.ID, - Username: user.Username, - ClientIP: c.ClientIP(), - Method: c.Request.Method, - Path: c.Request.URL.Path, - RequestURI: c.Request.RequestURI, - UserAgent: c.Request.UserAgent(), - Referer: c.Request.Referer(), - } - auditJSON, err := json.Marshal(auditLog) - if err != nil { - logger.ErrorF(ctx, "[LoginRequiredAudit] marshal failed: %v", err) - logger.DebugF(ctx, "[LoginRequiredAudit] %s %d %s", c.ClientIP(), user.ID, user.Username) - } else { - logger.DebugF(ctx, "[LoginRequiredAudit] %s", auditJSON) - } -} diff --git a/internal/apps/oauth/auth_source_resolver.go b/internal/apps/oauth/auth_source_resolver.go deleted file mode 100644 index e996c4be..00000000 --- a/internal/apps/oauth/auth_source_resolver.go +++ /dev/null @@ -1,118 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -import ( - "context" - "errors" - "strings" - - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/coreos/go-oidc/v3/oidc" - "golang.org/x/oauth2" -) - -func isOIDCLoginEnabled(ctx context.Context) bool { - enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled) - if err != nil { - return true - } - return enabled -} - -func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSource, error) { - name := strings.TrimSpace(strings.ToLower(sourceName)) - if name == "" { - sources, err := repository.GetActiveAuthSourcesCached(ctx) - if err != nil { - return nil, err - } - if len(sources) == 0 { - return nil, errors.New(errNoActiveAuthSource) - } - return repository.GetAuthSourceByNameCached(ctx, sources[0].Name) - } - return repository.GetAuthSourceByNameCached(ctx, name) -} - -func activeLoginSources(ctx context.Context) []AuthSourceView { - enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled) - if err == nil && !enabled { - return nil - } - - dbSources, err := repository.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) { - sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress) - if err != nil || strings.TrimSpace(sc.Value) == "" { - return "", errors.New(errServerAddressMissing) - } - return strings.TrimRight(sc.Value, "/") + "/login", nil -} - -func buildOAuthConfig(ctx context.Context, source *model.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 (trim /.well-known/openid-configuration if configured by mistake) - 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,避免每次调用都向 issuer 发起 - // /.well-known/openid-configuration HTTP 请求。 - 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 -} diff --git a/internal/apps/oauth/cache.go b/internal/apps/oauth/cache.go deleted file mode 100644 index 4f91b5bb..00000000 --- a/internal/apps/oauth/cache.go +++ /dev/null @@ -1,252 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -import ( - "context" - "fmt" - "strconv" - "sync" - "time" - - "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/pkg/cache/ram" - "github.com/Rain-kl/Wavelet/pkg/util" -) - -const ( - tokenCacheTTL = 5 * time.Minute - userCacheTTL = 5 * time.Minute - - //nolint:gosec // This is a Redis Pub/Sub channel name, not a credential - oauthTokenInvalidationChannel = "oauth:token_invalidation" - oauthUserInvalidationChannel = "oauth:user_invalidation" -) - -var ( - tokenRAM = ram.MustNew[string, *model.AccessToken](ram.Options{MaximumSize: 2048}) - userRAM = ram.MustNew[uint64, *model.User](ram.Options{MaximumSize: 2048}) - - tokenListenerOnce sync.Once - tokenListenerCtx context.Context - tokenListenerCancel context.CancelFunc - tokenListenerDone chan struct{} - - userListenerOnce sync.Once - userListenerCtx context.Context - userListenerCancel context.CancelFunc - userListenerDone chan struct{} -) - -func tokenCacheKey(tokenHash string) string { - return "oauth:token:" + tokenHash -} - -func userCacheKey(userID uint64) string { - return fmt.Sprintf("oauth:user:%d", userID) -} - -func ensureTokenCacheListener() { - if db.Redis == nil { - return - } - tokenListenerOnce.Do(startTokenCacheInvalidationListener) -} - -func startTokenCacheInvalidationListener() { - tokenListenerCtx, tokenListenerCancel = context.WithCancel(context.Background()) - tokenListenerDone = make(chan struct{}) - - redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争 - util.Go(func() { - listenerCtx := tokenListenerCtx - defer close(tokenListenerDone) - - pubsub := redisClient.Subscribe(listenerCtx, oauthTokenInvalidationChannel) - defer func() { - _ = pubsub.Close() - }() - - util.Go(func() { - <-listenerCtx.Done() - _ = pubsub.Close() - }) - - for msg := range pubsub.Channel() { - tokenHash := msg.Payload - if tokenHash == "" || tokenHash == "*" || tokenHash == "reset" { - tokenRAM.InvalidateAll() - } else { - tokenRAM.Invalidate(tokenHash) - } - } - }) -} - -func publishTokenRAMInvalidation(ctx context.Context, tokenHash string) { - if db.Redis == nil { - return - } - _ = db.Redis.Publish(ctx, oauthTokenInvalidationChannel, tokenHash).Err() -} - -func ensureUserCacheListener() { - if db.Redis == nil { - return - } - userListenerOnce.Do(startUserCacheInvalidationListener) -} - -func startUserCacheInvalidationListener() { - userListenerCtx, userListenerCancel = context.WithCancel(context.Background()) - userListenerDone = make(chan struct{}) - - redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争 - util.Go(func() { - listenerCtx := userListenerCtx - defer close(userListenerDone) - - pubsub := redisClient.Subscribe(listenerCtx, oauthUserInvalidationChannel) - defer func() { - _ = pubsub.Close() - }() - - util.Go(func() { - <-listenerCtx.Done() - _ = pubsub.Close() - }) - - for msg := range pubsub.Channel() { - userIDStr := msg.Payload - if userIDStr == "" || userIDStr == "*" || userIDStr == "reset" { - userRAM.InvalidateAll() - } else if userID, err := strconv.ParseUint(userIDStr, 10, 64); err == nil { - userRAM.Invalidate(userID) - } - } - }) -} - -func publishUserRAMInvalidation(ctx context.Context, userID uint64) { - if db.Redis == nil { - return - } - _ = db.Redis.Publish(ctx, oauthUserInvalidationChannel, strconv.FormatUint(userID, 10)).Err() -} - -// GetCachedToken 获取缓存的 AccessToken -func GetCachedToken(ctx context.Context, tokenHash string) (*model.AccessToken, error) { - ensureTokenCacheListener() - - if val, ok := tokenRAM.GetIfPresent(tokenHash); ok { - return val, nil - } - - if db.Redis != nil { - var token model.AccessToken - key := tokenCacheKey(tokenHash) - if err := db.GetJSON(ctx, key, &token); err == nil { - // Write back to local cache - tokenRAM.Set(tokenHash, &token) - return &token, nil - } - } - return nil, fmt.Errorf("cache miss") -} - -// SetCachedToken 设置 AccessToken 缓存 -func SetCachedToken(ctx context.Context, tokenHash string, token *model.AccessToken) { - ensureTokenCacheListener() - - tokenRAM.Set(tokenHash, token) - if db.Redis != nil { - key := tokenCacheKey(tokenHash) - _ = db.SetJSON(ctx, key, token, tokenCacheTTL) - } -} - -// InvalidateCachedToken 吊销/删除 token 缓存 -func InvalidateCachedToken(ctx context.Context, tokenHash string) { - ensureTokenCacheListener() - - tokenRAM.Invalidate(tokenHash) - if db.Redis != nil { - key := tokenCacheKey(tokenHash) - _ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err() - publishTokenRAMInvalidation(ctx, tokenHash) - } -} - -// GetCachedUser 获取缓存的 User -func GetCachedUser(ctx context.Context, userID uint64) (*model.User, error) { - ensureUserCacheListener() - - if val, ok := userRAM.GetIfPresent(userID); ok { - return val, nil - } - - if db.Redis != nil { - var u model.User - key := userCacheKey(userID) - if err := db.GetJSON(ctx, key, &u); err == nil { - // Write back to local cache - userRAM.Set(userID, &u) - return &u, nil - } - } - return nil, fmt.Errorf("cache miss") -} - -// SetCachedUser 设置 User 缓存 -func SetCachedUser(ctx context.Context, userID uint64, u *model.User) { - ensureUserCacheListener() - - userRAM.Set(userID, u) - if db.Redis != nil { - key := userCacheKey(userID) - _ = db.SetJSON(ctx, key, u, userCacheTTL) - } -} - -// InvalidateCachedUser 吊销/失效 User 缓存 -func InvalidateCachedUser(ctx context.Context, userID uint64) { - ensureUserCacheListener() - - userRAM.Invalidate(userID) - if db.Redis != nil { - key := userCacheKey(userID) - _ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err() - publishUserRAMInvalidation(ctx, userID) - } -} - -// StopOauthCacheListener stops both token and user Redis Pub/Sub subscription listeners and resets the sync.Once guards. -func StopOauthCacheListener() { - if tokenListenerCancel != nil { - tokenListenerCancel() - if tokenListenerDone != nil { - <-tokenListenerDone - } - tokenListenerCancel = nil - tokenListenerDone = nil - } - tokenListenerOnce = sync.Once{} - - if userListenerCancel != nil { - userListenerCancel() - if userListenerDone != nil { - <-userListenerDone - } - userListenerCancel = nil - userListenerDone = nil - } - userListenerOnce = sync.Once{} -} - -// ResetOauthRAMCacheForTest clears only the process-local RAM cache. -func ResetOauthRAMCacheForTest() { - tokenRAM.InvalidateAll() - userRAM.InvalidateAll() -} diff --git a/internal/apps/oauth/cache_test.go b/internal/apps/oauth/cache_test.go deleted file mode 100644 index 966cb5e9..00000000 --- a/internal/apps/oauth/cache_test.go +++ /dev/null @@ -1,235 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -import ( - "context" - "testing" - "time" - - "github.com/alicebob/miniredis/v2" - "github.com/redis/go-redis/v9" - "github.com/redis/go-redis/v9/maintnotifications" - - "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/model" -) - -func setupOauthCacheTest(t *testing.T) (*miniredis.Miniredis, func()) { - t.Helper() - - miniRedis, err := miniredis.Run() - if err != nil { - t.Fatalf("failed to start miniredis: %v", err) - } - - db.Redis = redis.NewClient(&redis.Options{ - Addr: miniRedis.Addr(), - MaintNotificationsConfig: &maintnotifications.Config{ - Mode: maintnotifications.ModeDisabled, - }, - }) - - ResetOauthRAMCacheForTest() - - cleanup := func() { - StopOauthCacheListener() - ResetOauthRAMCacheForTest() - db.Redis.Close() - miniRedis.Close() - db.Redis = nil - } - return miniRedis, cleanup -} - -func TestTokenCache_GetSetInvalidate(t *testing.T) { - _, cleanup := setupOauthCacheTest(t) - defer cleanup() - ctx := context.Background() - - tokenHash := "test-token-hash" - token := &model.AccessToken{ - ID: 123, - UserID: 456, - TokenHash: tokenHash, - Name: "test-token", - } - - // 1. Get from empty cache -> miss - _, err := GetCachedToken(ctx, tokenHash) - if err == nil { - t.Fatal("expected cache miss for un-cached token") - } - - // 2. Set to cache - SetCachedToken(ctx, tokenHash, token) - - // 3. Get from cache -> hit - cached, err := GetCachedToken(ctx, tokenHash) - if err != nil { - t.Fatalf("GetCachedToken() failed: %v", err) - } - if cached.ID != token.ID || cached.UserID != token.UserID { - t.Fatalf("expected cached token %+v, got %+v", token, cached) - } - - // 4. Invalidate cache - InvalidateCachedToken(ctx, tokenHash) - - // 5. Get from cache -> miss - _, err = GetCachedToken(ctx, tokenHash) - if err == nil { - t.Fatal("expected cache miss after invalidation") - } -} - -func TestUserCache_GetSetInvalidate(t *testing.T) { - _, cleanup := setupOauthCacheTest(t) - defer cleanup() - ctx := context.Background() - - userID := uint64(789) - user := &model.User{ - ID: userID, - Username: "testuser", - Email: "test@example.com", - } - - // 1. Get from empty cache -> miss - _, err := GetCachedUser(ctx, userID) - if err == nil { - t.Fatal("expected cache miss for un-cached user") - } - - // 2. Set to cache - SetCachedUser(ctx, userID, user) - - // 3. Get from cache -> hit - cached, err := GetCachedUser(ctx, userID) - if err != nil { - t.Fatalf("GetCachedUser() failed: %v", err) - } - if cached.ID != user.ID || cached.Username != user.Username { - t.Fatalf("expected cached user %+v, got %+v", user, cached) - } - - // 4. Invalidate cache - InvalidateCachedUser(ctx, userID) - - // 5. Get from cache -> miss - _, err = GetCachedUser(ctx, userID) - if err == nil { - t.Fatal("expected cache miss after invalidation") - } -} - -func TestOauthCache_PubSubInvalidation(t *testing.T) { - _, cleanup := setupOauthCacheTest(t) - defer cleanup() - ctx := context.Background() - - tokenHash := "pubsub-token-hash" - token := &model.AccessToken{ - ID: 111, - UserID: 222, - TokenHash: tokenHash, - } - - userID := uint64(333) - user := &model.User{ - ID: userID, - Username: "pubsubuser", - } - - // Set caches so they are stored in RAM - SetCachedToken(ctx, tokenHash, token) - SetCachedUser(ctx, userID, user) - - // Verify they are cached - if _, ok := tokenRAM.GetIfPresent(tokenHash); !ok { - t.Fatal("expected token to be in RAM cache") - } - if _, ok := userRAM.GetIfPresent(userID); !ok { - t.Fatal("expected user to be in RAM cache") - } - - // Give Pub/Sub subscription time to establish - time.Sleep(100 * time.Millisecond) - - // Publish invalidation messages directly to simulate peer node invalidation - if err := db.Redis.Publish(ctx, oauthTokenInvalidationChannel, tokenHash).Err(); err != nil { - t.Fatalf("publish token invalidation error: %v", err) - } - if err := db.Redis.Publish(ctx, oauthUserInvalidationChannel, "333").Err(); err != nil { - t.Fatalf("publish user invalidation error: %v", err) - } - - // Wait for background pubsub handlers to process messages - deadline := time.Now().Add(500 * time.Millisecond) - for time.Now().Before(deadline) { - _, tokenOk := tokenRAM.GetIfPresent(tokenHash) - _, userOk := userRAM.GetIfPresent(userID) - if !tokenOk && !userOk { - break - } - time.Sleep(10 * time.Millisecond) - } - - if _, ok := tokenRAM.GetIfPresent(tokenHash); ok { - t.Fatal("expected token RAM cache to be invalidated by Pub/Sub") - } - if _, ok := userRAM.GetIfPresent(userID); ok { - t.Fatal("expected user RAM cache to be invalidated by Pub/Sub") - } -} - -func TestOauthCache_PubSubResetAll(t *testing.T) { - _, cleanup := setupOauthCacheTest(t) - defer cleanup() - ctx := context.Background() - - tokenHash := "reset-token-hash" - token := &model.AccessToken{ - ID: 444, - UserID: 555, - TokenHash: tokenHash, - } - - userID := uint64(666) - user := &model.User{ - ID: userID, - Username: "resetuser", - } - - SetCachedToken(ctx, tokenHash, token) - SetCachedUser(ctx, userID, user) - - // Give Pub/Sub subscription time to establish - time.Sleep(100 * time.Millisecond) - - // Publish dynamic reset/wildcard to clear all - if err := db.Redis.Publish(ctx, oauthTokenInvalidationChannel, "*").Err(); err != nil { - t.Fatalf("publish token reset error: %v", err) - } - if err := db.Redis.Publish(ctx, oauthUserInvalidationChannel, "reset").Err(); err != nil { - t.Fatalf("publish user reset error: %v", err) - } - - deadline := time.Now().Add(500 * time.Millisecond) - for time.Now().Before(deadline) { - _, tokenOk := tokenRAM.GetIfPresent(tokenHash) - _, userOk := userRAM.GetIfPresent(userID) - if !tokenOk && !userOk { - break - } - time.Sleep(10 * time.Millisecond) - } - - if _, ok := tokenRAM.GetIfPresent(tokenHash); ok { - t.Fatal("expected token RAM cache to be fully cleared by '*'") - } - if _, ok := userRAM.GetIfPresent(userID); ok { - t.Fatal("expected user RAM cache to be fully cleared by 'reset'") - } -} diff --git a/internal/apps/oauth/constants.go b/internal/apps/oauth/constants.go deleted file mode 100644 index 83d2d37f..00000000 --- a/internal/apps/oauth/constants.go +++ /dev/null @@ -1,58 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -import ( - "encoding/json" - "time" -) - -// Session 用户信息字段 Key -const ( - UserNameKey = "username" - UserIDKey = "user_id" - 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 - PasswordHashKey = "password_hash" -) - -// OAuth State 缓存 Key 格式与过期时间 -const ( - OAuthStateCacheKeyFormat = "oauth:state:%s" - OAuthStateCacheKeyExpiration = 10 * time.Minute - oauthStateLimitKeyFormat = "oauth:state:limit:%s" - oauthStateLimitMax = 10 -) - -// OAuth 授权用途常量 -const ( - OAuthPurposeLogin = "login" - OAuthPurposeBind = "bind" -) - -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 -} diff --git a/internal/apps/oauth/errs.go b/internal/apps/oauth/errs.go deleted file mode 100644 index 4ccd37cc..00000000 --- a/internal/apps/oauth/errs.go +++ /dev/null @@ -1,23 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -// OAuth 认证相关错误消息 -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 = "请求授权过于频繁,请稍后重试" -) diff --git a/internal/apps/oauth/gin_context.go b/internal/apps/oauth/gin_context.go deleted file mode 100644 index 646f50a0..00000000 --- a/internal/apps/oauth/gin_context.go +++ /dev/null @@ -1,22 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -import "github.com/gin-gonic/gin" - -// GetFromContext 从 Gin 请求上下文获取指定类型的值。 -func GetFromContext[T any](c *gin.Context, key string) (T, bool) { - value, exists := c.Get(key) - if !exists { - var zero T - return zero, false - } - typed, ok := value.(T) - return typed, ok -} - -// SetToContext 设置值到 Gin 请求上下文。 -func SetToContext[T any](c *gin.Context, key string, value T) { - c.Set(key, value) -} diff --git a/internal/apps/oauth/handler_authorize.go b/internal/apps/oauth/handler_authorize.go deleted file mode 100644 index 0935b6d7..00000000 --- a/internal/apps/oauth/handler_authorize.go +++ /dev/null @@ -1,200 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -import ( - "context" - "errors" - "fmt" - "net/http" - "strings" - - "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/shared" - "github.com/Rain-kl/Wavelet/internal/shared/response" - "github.com/coreos/go-oidc/v3/oidc" - "github.com/gin-contrib/sessions" - "github.com/gin-gonic/gin" - "github.com/google/uuid" -) - -// GetLoginURL 获取登录授权地址 -// @Summary 获取登录授权地址 -// @Description 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用第一个启用的认证源。 -// @Tags oauth -// @Produce json -// @Param source query string false "认证源名称,为空使用第一个启用的认证源" -// @Success 200 {object} response.Any{data=oauth.OAuthAuthorizeResponse} "授权 URL" -// @Failure 400 {object} response.Any "认证源不存在或未配置" -// @Failure 500 {object} response.Any "Redis 异常 or 构造 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 - } - if err := db.Redis.Set(ctx, db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); 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 *model.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 db.Redis == nil || sessionHash == "" { - return nil - } - key := db.PrefixedKey(fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash)) - n, err := db.Redis.Incr(ctx, key).Result() - if err != nil { - return err - } - if n == 1 { - _ = db.Redis.Expire(ctx, key, OAuthStateCacheKeyExpiration).Err() - } - if n > 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=oauth.OAuthAuthorizeResponse} "授权 URL" -// @Failure 400 {object} response.Any "认证源不存在或未启用" -// @Failure 500 {object} response.Any "Redis 异常或构造 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, shared.UnAuthorized) - 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 - } - if err := db.Redis.Set(ctx, db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); 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})) -} diff --git a/internal/apps/oauth/handler_callback.go b/internal/apps/oauth/handler_callback.go deleted file mode 100644 index 9f663d86..00000000 --- a/internal/apps/oauth/handler_callback.go +++ /dev/null @@ -1,231 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -import ( - "context" - "errors" - "fmt" - "net/http" - "time" - - "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/listener" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/shared" - "github.com/Rain-kl/Wavelet/internal/shared/response" - "github.com/Rain-kl/Wavelet/pkg/logger" - "github.com/gin-contrib/sessions" - "github.com/gin-gonic/gin" - "gorm.io/gorm" -) - -// Callback OAuth 回调处理 -// @Summary OAuth 回调处理 -// @Description 接收前端传回的 state 和 code,完成 OAuth/OIDC 认证并建立会话。支持登录(login)和账号绑定(bind)两种场景。 -// @Tags oauth -// @Accept json -// @Produce json -// @Param request body oauth.CallbackRequest true "回调请求参数" -// @Success 200 {object} response.Any{data=oauth.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 := db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State)) - payloadRaw, err := db.Redis.Get(ctx, stateKey).Result() - if err != nil { - response.AbortBadRequest(c, errInvalidState) - return - } - _ = db.Redis.Del(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, shared.UnAuthorized) - return - } - - token, ok := session.Get(SessionTokenKey).(string) - if !ok || token == "" { - response.AbortBadRequest(c, "invalid session context") - return - } - - if hashSessionToken(token) != payload.SessionHash { - response.AbortBadRequest(c, "session mismatch for oauth state") - return - } - - if payload.Purpose == OAuthPurposeBind && currentUserID != payload.UserID { - response.AbortBadRequest(c, "user context mismatch for oauth binding") - 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) -} - -// handleCallbackBind 处理 OAuth 回调中的帐号绑定流程 -func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) { - userID := GetUserIDFromContext(c) - if userID == 0 { - response.AbortUnauthorized(c, shared.UnAuthorized) - return - } - user, err := repository.GetUserByID(ctx, userID) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - if err := repository.BindExternalAccount(ctx, &model.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() - _ = repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt) - c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "bound"))) -} - -// handleCallbackLogin 处理 OAuth 回调中的登录流程(查找已有帐号或自动注册) -func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) { - var user model.User - - account, err := repository.FindExternalAccount(ctx, source.ID, userInfo.Sub) - switch { - case err == nil: - loaded, loadErr := repository.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() - _ = repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt) - if err := SetLoginSession(ctx, c, &user); err != nil { - response.AbortInternal(c, err.Error()) - return - } - - 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()) - - listener.EmitAdminLoggedIn(ctx, &user, c.ClientIP()) - - c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "logged_in"))) -} - -// handleCallbackRegister 处理 OAuth 回调中的自动注册流程 -// 若注册被禁用则保存 pending 信息并返回 false;若注册成功则返回新用户;若出错则返回 false -func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) (model.User, bool) { - registrationEnabled, regErr := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled) - if regErr != nil { - registrationEnabled = false - } - - if !registrationEnabled { - c.JSON(http.StatusOK, response.OK(buildCallbackResult(nil, "need_bind"))) - return model.User{}, false - } - - username, uniqueErr := uniqueUsername(ctx, userInfo.Username) - if uniqueErr != nil { - response.AbortInternal(c, uniqueErr.Error()) - return model.User{}, false - } - userInfo.Username = username - - var user model.User - if err := repository.CreateUserFromOAuth(ctx, &user, userInfo); err != nil { - response.AbortInternal(c, err.Error()) - return model.User{}, false - } - if err := repository.BindExternalAccount(ctx, &model.ExternalAccount{ - AuthSourceID: source.ID, - UserID: user.ID, - ExternalID: userInfo.Sub, - ExternalUsername: userInfo.Username, - Email: userInfo.Email, - }); err != nil { - response.AbortBadRequest(c, err.Error()) - return model.User{}, 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 -} diff --git a/internal/apps/oauth/handler_external_accounts.go b/internal/apps/oauth/handler_external_accounts.go deleted file mode 100644 index aa922c1d..00000000 --- a/internal/apps/oauth/handler_external_accounts.go +++ /dev/null @@ -1,66 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -import ( - "net/http" - "strconv" - "strings" - - "github.com/Rain-kl/Wavelet/internal/repository" - - "github.com/Rain-kl/Wavelet/internal/shared" - "github.com/Rain-kl/Wavelet/internal/shared/response" - "github.com/gin-gonic/gin" -) - -// ListExternalAccounts 获取当前用户的外部帐号绑定列表 -// @Summary 获取外部帐号列表 -// @Description 返回当前登录用户已绑定的所有外部 OAuth 帐号信息,需要登录 -// @Tags oauth -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=[]model.ExternalAccountView} "外部帐号列表" -// @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 := repository.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, shared.UnAuthorized) - 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 := repository.DeleteExternalAccountForUser(c.Request.Context(), id, userID); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OKNil()) -} diff --git a/internal/apps/oauth/handler_sources.go b/internal/apps/oauth/handler_sources.go deleted file mode 100644 index 2cefa790..00000000 --- a/internal/apps/oauth/handler_sources.go +++ /dev/null @@ -1,22 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -import ( - "net/http" - - "github.com/Rain-kl/Wavelet/internal/shared/response" - "github.com/gin-gonic/gin" -) - -// GetLoginSources 获取可用登录源列表 -// @Summary 获取可用登录源 -// @Description 返回当前系统已启用的所有 OAuth 登录源,前端展示登录按钮列表时调用 -// @Tags oauth -// @Produce json -// @Success 200 {object} response.Any{data=[]oauth.AuthSourceView} "登录源列表" -// @Router /api/v1/oauth/sources [get] -func GetLoginSources(c *gin.Context) { - c.JSON(http.StatusOK, response.OK(activeLoginSources(c.Request.Context()))) -} diff --git a/internal/apps/oauth/middlewares.go b/internal/apps/oauth/middlewares.go deleted file mode 100644 index 6a20fa8c..00000000 --- a/internal/apps/oauth/middlewares.go +++ /dev/null @@ -1,153 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -import ( - "context" - "errors" - - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/shared" - "github.com/Rain-kl/Wavelet/internal/shared/response" - - otel_trace "github.com/Rain-kl/Wavelet/pkg/trace" - "github.com/gin-contrib/sessions" - "github.com/gin-gonic/gin" -) - -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"` -} - -func getUserByToken(ctx context.Context, tokenStr string) (*model.User, *model.AccessToken, error) { - tokenHash := model.HashToken(tokenStr) - tokenRecord, err := GetCachedToken(ctx, tokenHash) - if err != nil { - dbToken, err := repository.GetAccessTokenByHash(ctx, tokenHash) - if err != nil { - return nil, nil, err - } - tokenRecord = &dbToken - SetCachedToken(ctx, tokenHash, tokenRecord) - } - - user, err := GetCachedUser(ctx, tokenRecord.UserID) - if err != nil || !user.IsActive { - dbUser, err := repository.GetActiveUserByID(ctx, tokenRecord.UserID) - if err != nil { - return nil, nil, err - } - user = &dbUser - SetCachedUser(ctx, tokenRecord.UserID, user) - } - return user, tokenRecord, nil -} - -// GetUserFromRequest 校验 Access Token 或 Session 并返回用户对象,如果未登录或用户失效则返回 error -func GetUserFromRequest(c *gin.Context) (*model.User, error) { - ctx := c.Request.Context() - - // check token in headers - tokenStr := c.GetHeader("X-Access-Token") - if tokenStr == "" { - 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 { - // 强行阻止 system 用户任何会话/Token 鉴权通过 - if user.Username == "system" { - return nil, errors.New("system user is not allowed to login") - } - SetToContext(c, TokenAuthKey, true) - SetToContext(c, TokenAdminKey, tokenRecord.IsAdmin) - return user, nil - } - } - - // 降级使用 Session 鉴权 - userID := GetUserIDFromContext(c) - if userID <= 0 { - return nil, errors.New("unauthorized") - } - - user, err := GetCachedUser(ctx, userID) - if err != nil || !user.IsActive { - // load user from db to make sure is active - dbUser, loadErr := repository.GetActiveUserByID(ctx, userID) - if loadErr != nil { - return nil, loadErr - } - user = &dbUser - SetCachedUser(ctx, userID, user) - } - - // 密码哈希校验:当用户存在本地密码时,要求 Session 中的密码哈希必须与当前数据库中一致 - if user.Password != "" { - session := sessions.Default(c) - sessionHash, _ := session.Get(PasswordHashKey).(string) - if sessionHash != user.Password { - return nil, errors.New("session expired due to password change") - } - } - - // set keys in context for session auth - SetToContext(c, TokenAuthKey, false) - SetToContext(c, TokenAdminKey, false) - - // 强行阻止 system 用户任何会话/Token 鉴权通过 - if user.Username == "system" { - return nil, errors.New("system user is not allowed to login") - } - - return user, nil -} - -// LoginRequired 返回登录鉴权中间件,校验 Access Token 或 Session -func LoginRequired() gin.HandlerFunc { - return func(c *gin.Context) { - // init trace - ctx, span := otel_trace.Start(c.Request.Context(), "LoginRequired") - defer span.End() - - user, err := GetUserFromRequest(c) - if err != nil { - response.AbortUnauthorized(c, shared.UnAuthorized) - return - } - - // log - LogForAudit(ctx, user, c) - - // set user info - SetToContext(c, UserObjKey, user) - - // next - c.Next() - } -} - -// DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点 -func DisallowTokenAuth() gin.HandlerFunc { - return func(c *gin.Context) { - if tokenAuth, _ := GetFromContext[bool](c, TokenAuthKey); tokenAuth { - response.AbortForbidden(c, ErrTokenAuthNotAllowed) - return - } - c.Next() - } -} diff --git a/internal/apps/oauth/oauth_test.go b/internal/apps/oauth/oauth_test.go deleted file mode 100644 index 87932e48..00000000 --- a/internal/apps/oauth/oauth_test.go +++ /dev/null @@ -1,1343 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -import ( - "bytes" - "context" - "crypto/rand" - "crypto/rsa" - "encoding/json" - "fmt" - "io" - "net/http" - "net/http/httptest" - "net/url" - "strconv" - "strings" - "testing" - "time" - - "github.com/coreos/go-oidc/v3/oidc" - "github.com/gin-contrib/sessions" - "github.com/gin-contrib/sessions/cookie" - "github.com/gin-gonic/gin" - "github.com/go-jose/go-jose/v4" - "github.com/redis/go-redis/v9" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - "golang.org/x/oauth2" - "gorm.io/driver/sqlite" - "gorm.io/gorm" - - "github.com/Rain-kl/Wavelet/internal/infra/config" - "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/Rain-kl/Wavelet/internal/util" -) - -// ----------------------------------------------------------------------------- -// Mocks Setup -// ----------------------------------------------------------------------------- - -type mockRedisClient struct { - redis.UniversalClient - store map[string]string -} - -func newMockRedisClient() *mockRedisClient { - return &mockRedisClient{ - store: make(map[string]string), - } -} - -func (m *mockRedisClient) Set(ctx context.Context, key string, value interface{}, expiration time.Duration) *redis.StatusCmd { - cmd := redis.NewStatusCmd(ctx) - var val string - switch v := value.(type) { - case []byte: - val = string(v) - case string: - val = v - default: - val = fmt.Sprintf("%v", v) - } - m.store[key] = val - cmd.SetVal("OK") - return cmd -} - -func (m *mockRedisClient) Get(ctx context.Context, key string) *redis.StringCmd { - cmd := redis.NewStringCmd(ctx) - val, ok := m.store[key] - if !ok { - cmd.SetErr(redis.Nil) - } else { - cmd.SetVal(val) - } - return cmd -} - -func (m *mockRedisClient) Del(ctx context.Context, keys ...string) *redis.IntCmd { - cmd := redis.NewIntCmd(ctx) - var count int64 - for _, key := range keys { - if _, ok := m.store[key]; ok { - delete(m.store, key) - count++ - } - } - cmd.SetVal(count) - return cmd -} - -func (m *mockRedisClient) Incr(ctx context.Context, key string) *redis.IntCmd { - cmd := redis.NewIntCmd(ctx) - n := int64(1) - if raw, ok := m.store[key]; ok { - fmt.Sscan(raw, &n) - n++ - } - m.store[key] = fmt.Sprintf("%d", n) - cmd.SetVal(n) - return cmd -} - -func (m *mockRedisClient) Expire(ctx context.Context, key string, expiration time.Duration) *redis.BoolCmd { - cmd := redis.NewBoolCmd(ctx) - _, ok := m.store[key] - cmd.SetVal(ok) - return cmd -} - -func (m *mockRedisClient) Scan(ctx context.Context, cursor uint64, match string, count int64) *redis.ScanCmd { - cmd := redis.NewScanCmd(ctx, nil, cursor, match, count) - var keys []string - for key := range m.store { - if redisMatchPattern(key, match) { - keys = append(keys, key) - } - } - cmd.SetVal(keys, 0) - return cmd -} - -func redisMatchPattern(key, pattern string) bool { - if pattern == "" || pattern == "*" { - return true - } - if strings.HasSuffix(pattern, "*") { - return strings.HasPrefix(key, strings.TrimSuffix(pattern, "*")) - } - return key == pattern -} - -func (m *mockRedisClient) HSet(ctx context.Context, key string, values ...interface{}) *redis.IntCmd { - cmd := redis.NewIntCmd(ctx) - if len(values) >= 2 { - field := fmt.Sprintf("%v", values[0]) - var val string - switch v := values[1].(type) { - case []byte: - val = string(v) - case string: - val = v - default: - val = fmt.Sprintf("%v", v) - } - compositeKey := key + ":" + field - m.store[compositeKey] = val - cmd.SetVal(1) - } else { - cmd.SetVal(0) - } - return cmd -} - -func (m *mockRedisClient) HGet(ctx context.Context, key string, field string) *redis.StringCmd { - cmd := redis.NewStringCmd(ctx) - compositeKey := key + ":" + field - val, ok := m.store[compositeKey] - if !ok { - cmd.SetErr(redis.Nil) - } else { - cmd.SetVal(val) - } - return cmd -} - -func (m *mockRedisClient) Publish(ctx context.Context, channel string, message interface{}) *redis.IntCmd { - cmd := redis.NewIntCmd(ctx) - cmd.SetVal(1) - return cmd -} - -func (m *mockRedisClient) Subscribe(ctx context.Context, channels ...string) *redis.PubSub { - return redis.NewClient(&redis.Options{ - Addr: "127.0.0.1:0", - }).Subscribe(ctx, channels...) -} - -type mockRoundTripper struct { - roundTripFunc func(req *http.Request) (*http.Response, error) -} - -func (m *mockRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { - return m.roundTripFunc(req) -} - -// Global cryptographic tools for custom OIDC mocking -var ( - testRSAPrivateKey *rsa.PrivateKey - testJWKS jose.JSONWebKeySet -) - -const ( - testIssuerURL = "https://connect.linux.do" - testAuthURL = "https://connect.linux.do/oauth2/authorize" - testTokenURL = "https://connect.linux.do/oauth2/token" - testJWKSURL = "https://connect.linux.do/oauth2/keys" - testClientID = "test_client_id" - testClientSecret = "test_client_secret" - testSourceName = "linuxdo" - testSourceDisplay = "LINUX DO" -) - -func init() { - var err error - testRSAPrivateKey, err = rsa.GenerateKey(rand.Reader, 2048) - if err != nil { - panic(fmt.Sprintf("failed to generate RSA key: %v", err)) - } - jwk := jose.JSONWebKey{ - Key: &testRSAPrivateKey.PublicKey, - KeyID: "test-key-id", - Algorithm: string(jose.RS256), - Use: "sig", - } - testJWKS = jose.JSONWebKeySet{ - Keys: []jose.JSONWebKey{jwk}, - } -} - -func normalizeIssuerURL(issuer string) string { - return strings.TrimRight(strings.TrimSpace(issuer), "/") -} - -func seedTestAuthSource(t *testing.T, dbConn *gorm.DB) { - t.Helper() - if err := dbConn.Create(&model.AuthSource{ - ID: 100, - Name: testSourceName, - Type: model.AuthSourceTypeOIDC, - DisplayName: testSourceDisplay, - IsActive: true, - ClientID: testClientID, - ClientSecret: testClientSecret, - OpenIDDiscoveryURL: testIssuerURL, - }).Error; err != nil { - t.Fatalf("failed to seed auth source: %v", err) - } -} - -func oidcDiscoveryResponse() *http.Response { - issuer := normalizeIssuerURL(testIssuerURL) - body := fmt.Sprintf(`{ - "issuer": %q, - "authorization_endpoint": %q, - "token_endpoint": %q, - "jwks_uri": %q, - "response_types_supported": ["code"], - "subject_types_supported": ["public"], - "id_token_signing_alg_values_supported": ["RS256"] - }`, issuer, issuer+"/oauth2/authorize", issuer+"/oauth2/token", issuer+"/oauth2/keys") - return &http.Response{ - StatusCode: http.StatusOK, - Body: io.NopCloser(strings.NewReader(body)), - Header: make(http.Header), - } -} - -type mockClaims struct { - ID uint64 `json:"id"` - Issuer string `json:"iss"` - Subject string `json:"sub"` - Audience string `json:"aud"` - Expiry int64 `json:"exp"` - IssuedAt int64 `json:"iat"` - Nonce string `json:"nonce"` - Username string `json:"preferred_username"` - Email string `json:"email"` - Name string `json:"name"` - Active bool `json:"active"` -} - -func generateMockIDToken(issuer, sub, aud, nonce, username, email, name string) string { - signer, err := jose.NewSigner(jose.SigningKey{Algorithm: jose.RS256, Key: testRSAPrivateKey}, (&jose.SignerOptions{}).WithType("JWT")) - if err != nil { - panic(err) - } - id, _ := strconv.ParseUint(sub, 10, 64) - claims := mockClaims{ - ID: id, - Issuer: issuer, - Subject: sub, - Audience: aud, - Expiry: time.Now().Add(time.Hour).Unix(), - IssuedAt: time.Now().Unix(), - Nonce: nonce, - Username: username, - Email: email, - Name: name, - Active: true, - } - payload, _ := json.Marshal(claims) - object, err := signer.Sign(payload) - if err != nil { - panic(err) - } - tokenStr, _ := object.CompactSerialize() - return tokenStr -} - -// ----------------------------------------------------------------------------- -// Test Helpers -func newMockOIDCClient(issuer, clientID string, expectedState *string, sub, username, email, name string) *http.Client { - cleanIssuer := normalizeIssuerURL(issuer) - return &http.Client{ - Transport: &mockRoundTripper{ - roundTripFunc: func(req *http.Request) (*http.Response, error) { - urlStr := req.URL.String() - if req.Method == http.MethodGet && strings.Contains(urlStr, "/.well-known/openid-configuration") { - body := fmt.Sprintf(`{ - "issuer": %q, - "authorization_endpoint": %q, - "token_endpoint": %q, - "jwks_uri": %q, - "response_types_supported": ["code"], - "subject_types_supported": ["public"], - "id_token_signing_alg_values_supported": ["RS256"] - }`, cleanIssuer, cleanIssuer+"/oauth2/authorize", cleanIssuer+"/oauth2/token", cleanIssuer+"/oauth2/keys") - return &http.Response{ - StatusCode: http.StatusOK, - Body: io.NopCloser(strings.NewReader(body)), - Header: make(http.Header), - }, nil - } - if req.Method == http.MethodGet && (strings.Contains(urlStr, "/keys") || strings.Contains(urlStr, "/jwks")) { - jwksJSON, _ := json.Marshal(testJWKS) - return &http.Response{ - StatusCode: http.StatusOK, - Body: io.NopCloser(bytes.NewReader(jwksJSON)), - Header: make(http.Header), - }, nil - } - if req.Method == http.MethodPost && (strings.Contains(urlStr, "/token") || strings.Contains(urlStr, "/access_token")) { - var stateVal string - if expectedState != nil { - stateVal = *expectedState - } - idToken := generateMockIDToken(cleanIssuer, sub, clientID, stateVal, username, email, name) - body := fmt.Sprintf(`{"access_token":"mock_access_token","token_type":"Bearer","expires_in":3600,"id_token":"%s"}`, idToken) - return &http.Response{ - StatusCode: http.StatusOK, - Body: io.NopCloser(strings.NewReader(body)), - Header: make(http.Header), - }, nil - } - return nil, fmt.Errorf("unexpected mock request: %s %s", req.Method, req.URL) - }, - }, - } -} - -func setupTestDB(t *testing.T) *gorm.DB { - repository.ResetSystemConfigRAMCacheForTest() - repository.ResetAuthSourceRAMCacheForTest() - - dbConn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) - if err != nil { - t.Fatalf("failed to open sqlite in memory: %v", err) - } - err = dbConn.AutoMigrate( - &model.User{}, - &model.AuthSource{}, - &model.ExternalAccount{}, - &model.SystemConfig{}, - ) - if err != nil { - t.Fatalf("failed to migrate schema: %v", err) - } - - // 注入测试所需的服务器地址配置 - if err := dbConn.Create(&model.SystemConfig{ - Key: model.ConfigKeyServerAddress, - Value: "http://localhost:3000", - }).Error; err != nil { - t.Fatalf("failed to seed server_address config: %v", err) - } - - return dbConn -} - -func mockContextMiddleware(mockClient *http.Client) gin.HandlerFunc { - return func(c *gin.Context) { - ctx := c.Request.Context() - ctx = context.WithValue(ctx, oauth2.HTTPClient, mockClient) - ctx = oidc.ClientContext(ctx, mockClient) - c.Request = c.Request.WithContext(ctx) - c.Next() - } -} - -func resetOIDCProviderCacheForTest() { - InvalidateOIDCProviderCache(normalizeIssuerURL(testIssuerURL)) - InvalidateOIDCProviderCache("https://github.com") -} - -func setupTestRouter(dbConn *gorm.DB, mockRedis *mockRedisClient, mockClient *http.Client) *gin.Engine { - resetOIDCProviderCacheForTest() - - r := testhelper.NewTestGinEngine(gin.Recovery()) - - // Inject context mock middleware - r.Use(mockContextMiddleware(mockClient)) - - store := cookie.NewStore([]byte(config.Config.App.SessionSecret)) - store.Options(GetSessionOptions(3600)) - r.Use(sessions.Sessions(config.Config.App.SessionCookieName, store)) - - db.SetDB(dbConn) - db.Redis = mockRedis - - api := r.Group("/api/v1") - { - api.GET("/oauth/sources", GetLoginSources) - api.GET("/oauth/login", GetLoginURL) - api.GET("/oauth/:source/authorize", Authorize) - api.GET("/oauth/logout", Logout) - api.POST("/oauth/callback", Callback) - api.GET("/oauth/user-info", LoginRequired(), UserInfo) - api.GET("/oauth/external-accounts", LoginRequired(), ListExternalAccounts) - api.POST("/oauth/external-accounts/:id/delete", LoginRequired(), DeleteExternalAccount) - } - - return r -} - -func performRequest(r http.Handler, method, path string, body []byte, headers map[string]string, cookies []*http.Cookie) *httptest.ResponseRecorder { - var bodyReader io.Reader - if body != nil { - bodyReader = bytes.NewReader(body) - } - req, _ := http.NewRequest(method, path, bodyReader) - for k, v := range headers { - req.Header.Set(k, v) - } - for _, cookie := range cookies { - if cookie != nil { - req.AddCookie(cookie) - } - } - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - return w -} - -func initializeTestConfig() { - config.Config.App.Env = "testing" - config.Config.App.SessionCookieName = "test_session_id" - config.Config.App.SessionSecret = "test_session_secret" - config.Config.App.APIPrefix = "/api" -} - -// ----------------------------------------------------------------------------- -// Tests -// ----------------------------------------------------------------------------- - -func TestGetLoginSources(t *testing.T) { - initializeTestConfig() - dbConn := setupTestDB(t) - mockRedis := newMockRedisClient() - - // Setup empty HTTP Mock - httpMock := &http.Client{ - Transport: &mockRoundTripper{ - roundTripFunc: func(req *http.Request) (*http.Response, error) { - return nil, fmt.Errorf("unexpected request") - }, - }, - } - util.SetHTTPClient(httpMock) - router := setupTestRouter(dbConn, mockRedis, httpMock) - - // Inject OIDC login enabled config - dbConn.Create(&model.SystemConfig{ - Key: model.ConfigKeyOIDCLoginEnabled, - Value: "true", - }) - - // Inject active DB auth source - dbConn.Create(&model.AuthSource{ - ID: 101, - Name: "github", - Type: model.AuthSourceTypeOIDC, - DisplayName: "GitHub OAuth", - IsActive: true, - ClientID: "gh_client", - ClientSecret: "gh_secret", - OpenIDDiscoveryURL: "https://github.com", - }) - - // Perform GET request - w := performRequest(router, http.MethodGet, "/api/v1/oauth/sources", nil, nil, nil) - if w.Code != http.StatusOK { - t.Fatalf("expected status 200, got %d", w.Code) - } - - var resp struct { - Data []AuthSourceView `json:"data"` - } - err := json.Unmarshal(w.Body.Bytes(), &resp) - if err != nil { - t.Fatalf("failed to unmarshal response: %v", err) - } - - if len(resp.Data) != 1 { - t.Fatalf("expected 1 active source, got %d", len(resp.Data)) - } - - if resp.Data[0].Name != "github" { - t.Errorf("expected github source, got %s", resp.Data[0].Name) - } - - // Test disabling OIDC - dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false") - repository.ResetSystemConfigRAMCacheForTest() - mockRedis.store = make(map[string]string) - - w2 := performRequest(router, http.MethodGet, "/api/v1/oauth/sources", nil, nil, nil) - var resp2 struct { - Data []AuthSourceView `json:"data"` - } - _ = json.Unmarshal(w2.Body.Bytes(), &resp2) - if len(resp2.Data) != 0 { - t.Errorf("expected 0 sources when OIDC is disabled, got %d", len(resp2.Data)) - } -} - -func TestGetLoginURL(t *testing.T) { - initializeTestConfig() - dbConn := setupTestDB(t) - mockRedis := newMockRedisClient() - seedTestAuthSource(t, dbConn) - - httpMock := &http.Client{ - Transport: &mockRoundTripper{ - roundTripFunc: func(req *http.Request) (*http.Response, error) { - if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/.well-known/openid-configuration") { - return oidcDiscoveryResponse(), nil - } - return nil, fmt.Errorf("unexpected request") - }, - }, - } - util.SetHTTPClient(httpMock) - router := setupTestRouter(dbConn, mockRedis, httpMock) - - // Case 1: Default Login URL - w := performRequest(router, http.MethodGet, "/api/v1/oauth/login", nil, nil, nil) - if w.Code != http.StatusOK { - t.Fatalf("expected 200, got %d", w.Code) - } - - var resp struct { - Data OAuthAuthorizeResponse `json:"data"` - } - err := json.Unmarshal(w.Body.Bytes(), &resp) - if err != nil { - t.Fatalf("failed to unmarshal response: %v", err) - } - - if !strings.Contains(resp.Data.AuthorizeURL, testAuthURL) { - t.Errorf("invalid authorize URL: %s", resp.Data.AuthorizeURL) - } - - parsedURL, _ := url.Parse(resp.Data.AuthorizeURL) - state := parsedURL.Query().Get("state") - if state == "" { - t.Error("missing state in URL") - } - - redisKey := db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)) - stateVal, err := mockRedis.Get(context.Background(), redisKey).Result() - if err != nil { - t.Fatalf("state not found in redis: %v", err) - } - - payload, err := decodeOAuthStatePayload(stateVal) - if err != nil { - t.Fatalf("failed to decode state payload: %v", err) - } - if payload.SourceName != testSourceName || payload.Purpose != OAuthPurposeLogin { - t.Errorf("unexpected payload: %+v", payload) - } - - // Case 2: Unknown Source Login URL - w2 := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source=nonexistent", nil, nil, nil) - if w2.Code != http.StatusBadRequest { - t.Errorf("expected 400 for unknown source, got %d", w2.Code) - } -} - -func TestAuthorize(t *testing.T) { - initializeTestConfig() - dbConn := setupTestDB(t) - mockRedis := newMockRedisClient() - - // Setup GitHub active source - dbConn.Create(&model.AuthSource{ - ID: 101, - Name: "github", - Type: model.AuthSourceTypeOIDC, - DisplayName: "GitHub OAuth", - IsActive: true, - ClientID: "gh_client", - ClientSecret: "gh_secret", - OpenIDDiscoveryURL: "https://github.com", - }) - - // Mock OIDC Discovery request - httpMock := &http.Client{ - Transport: &mockRoundTripper{ - roundTripFunc: func(req *http.Request) (*http.Response, error) { - if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/.well-known/openid-configuration") { - body := `{ - "issuer": "https://github.com", - "authorization_endpoint": "https://github.com/login/oauth/authorize", - "token_endpoint": "https://github.com/login/oauth/access_token", - "jwks_uri": "https://github.com/oauth/keys", - "response_types_supported": ["code"], - "subject_types_supported": ["public"], - "id_token_signing_alg_values_supported": ["RS256"] - }` - return &http.Response{ - StatusCode: http.StatusOK, - Body: io.NopCloser(strings.NewReader(body)), - Header: make(http.Header), - }, nil - } - return nil, fmt.Errorf("unexpected request: %s", req.URL) - }, - }, - } - util.SetHTTPClient(httpMock) - router := setupTestRouter(dbConn, mockRedis, httpMock) - - // Case 1a: Active Source Authorize with purpose=bind without login -> 401 - wUnauth := performRequest(router, http.MethodGet, "/api/v1/oauth/github/authorize?purpose=bind", nil, nil, nil) - if wUnauth.Code != http.StatusUnauthorized { - t.Errorf("expected 401 for unauthorized bind authorize, got %d", wUnauth.Code) - } - - // Case 1b: Active Source Authorize with purpose=bind (authenticated) - router.GET("/test-helper/login-777", func(c *gin.Context) { - session := sessions.Default(c) - session.Set(UserIDKey, uint64(777)) - _ = session.Save() - c.String(200, "ok") - }) - wLogin := performRequest(router, http.MethodGet, "/test-helper/login-777", nil, nil, nil) - var activeCookie *http.Cookie - for _, cookie := range wLogin.Result().Cookies() { - if cookie.Name == config.Config.App.SessionCookieName { - activeCookie = cookie - break - } - } - - w := performRequest(router, http.MethodGet, "/api/v1/oauth/github/authorize?purpose=bind", nil, nil, []*http.Cookie{activeCookie}) - if w.Code != http.StatusOK { - t.Fatalf("expected 200, got %d, body: %s", w.Code, w.Body.String()) - } - - var resp struct { - Data OAuthAuthorizeResponse `json:"data"` - } - - _ = json.Unmarshal(w.Body.Bytes(), &resp) - - parsedURL, _ := url.Parse(resp.Data.AuthorizeURL) - state := parsedURL.Query().Get("state") - redisKey := db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)) - stateVal, _ := mockRedis.Get(context.Background(), redisKey).Result() - payload, _ := decodeOAuthStatePayload(stateVal) - - if payload.SourceName != "github" || payload.Purpose != OAuthPurposeBind { - t.Errorf("expected source github with purpose bind, got %+v", payload) - } - - // Case 2: Inactive Source Authorize - dbConn.Model(&model.AuthSource{}).Where("id = ?", 101).Update("is_active", false) - _ = repository.InvalidateAuthSourceCache(context.Background()) - w2 := performRequest(router, http.MethodGet, "/api/v1/oauth/github/authorize", nil, nil, nil) - if w2.Code != http.StatusBadRequest { - t.Errorf("expected 400 for inactive source, got %d", w2.Code) - } -} - -func TestLogout(t *testing.T) { - initializeTestConfig() - dbConn := setupTestDB(t) - mockRedis := newMockRedisClient() - httpMock := &http.Client{} - router := setupTestRouter(dbConn, mockRedis, httpMock) - - w := performRequest(router, http.MethodGet, "/api/v1/oauth/logout", nil, nil, nil) - if w.Code != http.StatusOK { - t.Fatalf("expected 200, got %d", w.Code) - } -} - -func TestCallbackLoginAndUserInfo(t *testing.T) { - initializeTestConfig() - dbConn := setupTestDB(t) - mockRedis := newMockRedisClient() - seedTestAuthSource(t, dbConn) - dbConn.Create(&model.SystemConfig{ - Key: model.ConfigKeyRegistrationEnabled, - Value: "true", - Type: "system", - }) - - var state string - - // 1. Mock the outgoing HTTP client for token exchange and user info fetching - httpMock := newMockOIDCClient(testIssuerURL, testClientID, &state, "88888", "test_oauth_user", "oauth@linux.do", "Oauth Test User") - util.SetHTTPClient(httpMock) - router := setupTestRouter(dbConn, mockRedis, httpMock) - - // Get Login URL first to initialize the session and generate the state - wLogin := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) - if wLogin.Code != http.StatusOK { - t.Fatalf("failed to get login URL: %s", wLogin.Body.String()) - } - - var loginUrlResp struct { - Data OAuthAuthorizeResponse `json:"data"` - } - _ = json.Unmarshal(wLogin.Body.Bytes(), &loginUrlResp) - - parsedURL, _ := url.Parse(loginUrlResp.Data.AuthorizeURL) - state = parsedURL.Query().Get("state") - - var anonymousCookie *http.Cookie - for _, cookie := range wLogin.Result().Cookies() { - if cookie.Name == config.Config.App.SessionCookieName { - anonymousCookie = cookie - break - } - } - if anonymousCookie == nil { - t.Fatal("session cookie not found after login URL generation") - } - - // 3. Trigger Callback (Login flow - new user) - reqBody := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state) - w := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody), map[string]string{ - "Content-Type": "application/json", - }, []*http.Cookie{anonymousCookie}) - - if w.Code != http.StatusOK { - t.Fatalf("callback failed with status %d, body: %s", w.Code, w.Body.String()) - } - - var callbackResp struct { - Data OAuthCallbackResult `json:"data"` - } - _ = json.Unmarshal(w.Body.Bytes(), &callbackResp) - - if callbackResp.Data.Status != "logged_in" { - t.Errorf("expected logged_in status, got %s", callbackResp.Data.Status) - } - - if callbackResp.Data.User.Username != "test_oauth_user" || callbackResp.Data.User.ID != 88888 { - t.Errorf("unexpected user returned: %+v", callbackResp.Data.User) - } - - // Verify user is created in database - var user model.User - if err := dbConn.First(&user, "id = ?", 88888).Error; err != nil { - t.Fatalf("user was not created in DB: %v", err) - } - - // Extract session cookie - cookies := w.Result().Cookies() - var sessionCookie *http.Cookie - for _, cookie := range cookies { - if cookie.Name == config.Config.App.SessionCookieName { - sessionCookie = cookie - break - } - } - if sessionCookie == nil { - t.Fatal("session cookie not found in response") - } - - // 4. Test GET user-info - w2 := performRequest(router, http.MethodGet, "/api/v1/oauth/user-info", nil, nil, []*http.Cookie{sessionCookie}) - if w2.Code != http.StatusOK { - t.Fatalf("failed to fetch user info, status %d", w2.Code) - } - - // 5. Test Callback (Login flow - existing user, username collision check) - var state2 string - // Callback with same username but different external ID (99999) - httpMock2 := newMockOIDCClient(testIssuerURL, testClientID, &state2, "99999", "test_oauth_user", "another@linux.do", "Another User") - util.SetHTTPClient(httpMock2) - - // Create another router for this mock client - router2 := setupTestRouter(dbConn, mockRedis, httpMock2) - - // Call login to get state2 and new anonymous session - wLogin2 := performRequest(router2, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) - var loginUrlResp2 struct { - Data OAuthAuthorizeResponse `json:"data"` - } - _ = json.Unmarshal(wLogin2.Body.Bytes(), &loginUrlResp2) - parsedURL2, _ := url.Parse(loginUrlResp2.Data.AuthorizeURL) - state2 = parsedURL2.Query().Get("state") - - var anonymousCookie2 *http.Cookie - for _, cookie := range wLogin2.Result().Cookies() { - if cookie.Name == config.Config.App.SessionCookieName { - anonymousCookie2 = cookie - break - } - } - - reqBody2 := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state2) - w3 := performRequest(router2, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody2), map[string]string{ - "Content-Type": "application/json", - }, []*http.Cookie{anonymousCookie2}) - - if w3.Code != http.StatusOK { - t.Fatalf("callback for collision failed: %d, body: %s", w3.Code, w3.Body.String()) - } - - var collisionResp struct { - Data OAuthCallbackResult `json:"data"` - } - _ = json.Unmarshal(w3.Body.Bytes(), &collisionResp) - - if collisionResp.Data.User.Username != "test_oauth_user-1" { - t.Errorf("expected collision renamed username, got %s", collisionResp.Data.User.Username) - } - - t.Run("OIDC login when registration disabled - need bind", func(t *testing.T) { - dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyRegistrationEnabled).Update("value", "false") - repository.ResetSystemConfigRAMCacheForTest() - mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyRegistrationEnabled) - t.Cleanup(func() { - dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyRegistrationEnabled).Update("value", "true") - repository.ResetSystemConfigRAMCacheForTest() - mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyRegistrationEnabled) - }) - - var state4 string - httpMock4 := newMockOIDCClient(testIssuerURL, testClientID, &state4, "77777", "need_bind_user", "needbind@linux.do", "Need Bind User") - util.SetHTTPClient(httpMock4) - router4 := setupTestRouter(dbConn, mockRedis, httpMock4) - - wLogin4 := performRequest(router4, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) - var loginUrlResp4 struct { - Data OAuthAuthorizeResponse `json:"data"` - } - _ = json.Unmarshal(wLogin4.Body.Bytes(), &loginUrlResp4) - parsedURL4, _ := url.Parse(loginUrlResp4.Data.AuthorizeURL) - state4 = parsedURL4.Query().Get("state") - - var anonymousCookie4 *http.Cookie - for _, cookie := range wLogin4.Result().Cookies() { - if cookie.Name == config.Config.App.SessionCookieName { - anonymousCookie4 = cookie - break - } - } - - reqBody4 := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state4) - w4 := performRequest(router4, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody4), map[string]string{ - "Content-Type": "application/json", - }, []*http.Cookie{anonymousCookie4}) - - if w4.Code != http.StatusOK { - t.Fatalf("callback failed: %d, body: %s", w4.Code, w4.Body.String()) - } - - var needBindResp struct { - Data OAuthCallbackResult `json:"data"` - } - _ = json.Unmarshal(w4.Body.Bytes(), &needBindResp) - - if needBindResp.Data.Status != "need_bind" { - t.Errorf("expected status 'need_bind', got %s", needBindResp.Data.Status) - } - if needBindResp.Data.User != nil { - t.Errorf("expected User to be nil, got %+v", needBindResp.Data.User) - } - }) -} - -func TestCallbackBind(t *testing.T) { - initializeTestConfig() - dbConn := setupTestDB(t) - mockRedis := newMockRedisClient() - - // Create user - user := model.User{ - ID: 777, - Username: "existing_member", - Nickname: "Existing Member", - IsActive: true, - LastLoginAt: time.Now(), - } - dbConn.Create(&user) - - // Create auth source "github" - dbConn.Create(&model.AuthSource{ - ID: 2, - Name: "github", - Type: model.AuthSourceTypeOIDC, - DisplayName: "GitHub", - IsActive: true, - ClientID: "gh_client", - ClientSecret: "gh_secret", - OpenIDDiscoveryURL: "https://github.com", - }) - - var state string - // Mock OIDC discovery, JWKS, and Token exchange for custom source (GitHub) - httpMock := newMockOIDCClient("https://github.com", "gh_client", &state, "github_user_123", "github_tester", "tester@github.com", "GitHub Tester") - util.SetHTTPClient(httpMock) - router := setupTestRouter(dbConn, mockRedis, httpMock) - - // Set up login helper - router.GET("/test-helper/login-777", func(c *gin.Context) { - session := sessions.Default(c) - session.Set(UserIDKey, uint64(777)) - _ = session.Save() - c.String(200, "ok") - }) - - wLogin := performRequest(router, http.MethodGet, "/test-helper/login-777", nil, nil, nil) - var activeCookie *http.Cookie - for _, cookie := range wLogin.Result().Cookies() { - if cookie.Name == config.Config.App.SessionCookieName { - activeCookie = cookie - break - } - } - - // Generate OAuth authorize link (purpose=bind) to set state in Redis and Session - wAuth := performRequest(router, http.MethodGet, "/api/v1/oauth/github/authorize?purpose=bind", nil, nil, []*http.Cookie{activeCookie}) - if wAuth.Code != http.StatusOK { - t.Fatalf("authorize failed: %d, body: %s", wAuth.Code, wAuth.Body.String()) - } - // Extract the cookie from wAuth to get the session with the token! - for _, cookie := range wAuth.Result().Cookies() { - if cookie.Name == config.Config.App.SessionCookieName { - activeCookie = cookie - break - } - } - var authResp struct { - Data OAuthAuthorizeResponse `json:"data"` - } - _ = json.Unmarshal(wAuth.Body.Bytes(), &authResp) - parsedURL, _ := url.Parse(authResp.Data.AuthorizeURL) - state = parsedURL.Query().Get("state") - - // Case 1: Bind attempt without session -> 401 - reqBody := fmt.Sprintf(`{"state":"%s","code":"test_code"}`, state) - w1 := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody), map[string]string{ - "Content-Type": "application/json", - }, nil) - - if w1.Code != http.StatusUnauthorized { - t.Errorf("expected 401 for unauthenticated bind, got %d, body: %s", w1.Code, w1.Body.String()) - } - - // Case 2: Bind success (authenticated) - // Re-run authorize since state is consumed/deleted during Callback attempt - wAuth2 := performRequest(router, http.MethodGet, "/api/v1/oauth/github/authorize?purpose=bind", nil, nil, []*http.Cookie{activeCookie}) - // Extract updated cookie from wAuth2 - for _, cookie := range wAuth2.Result().Cookies() { - if cookie.Name == config.Config.App.SessionCookieName { - activeCookie = cookie - break - } - } - var authResp2 struct { - Data OAuthAuthorizeResponse `json:"data"` - } - _ = json.Unmarshal(wAuth2.Body.Bytes(), &authResp2) - parsedURL2, _ := url.Parse(authResp2.Data.AuthorizeURL) - state = parsedURL2.Query().Get("state") - - reqBody2 := fmt.Sprintf(`{"state":"%s","code":"test_code"}`, state) - w2 := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody2), map[string]string{ - "Content-Type": "application/json", - }, []*http.Cookie{activeCookie}) - - if w2.Code != http.StatusOK { - t.Fatalf("expected 200 for bind callback, got %d, body: %s", w2.Code, w2.Body.String()) - } - - var bindResult struct { - Data OAuthCallbackResult `json:"data"` - } - _ = json.Unmarshal(w2.Body.Bytes(), &bindResult) - if bindResult.Data.Status != "bound" { - t.Errorf("expected status bound, got %s", bindResult.Data.Status) - } - - // Verify DB binding - var binding model.ExternalAccount - if err := dbConn.First(&binding, "user_id = ? AND external_id = ?", 777, "github_user_123").Error; err != nil { - t.Fatalf("DB binding record not found: %v", err) - } - - // Case 3: Bind already bound account to another user - // Create another user - user2 := model.User{ - ID: 888, - Username: "another_member", - Nickname: "Another Member", - IsActive: true, - LastLoginAt: time.Now(), - } - dbConn.Create(&user2) - - router.GET("/test-helper/login-888", func(c *gin.Context) { - session := sessions.Default(c) - session.Set(UserIDKey, uint64(888)) - _ = session.Save() - c.String(200, "ok") - }) - - wLogin2 := performRequest(router, http.MethodGet, "/test-helper/login-888", nil, nil, nil) - var activeCookie2 *http.Cookie - for _, cookie := range wLogin2.Result().Cookies() { - if cookie.Name == config.Config.App.SessionCookieName { - activeCookie2 = cookie - break - } - } - - var state3 string - // Re-sign token for new state (since state serves as OIDC Nonce) - httpMock3 := newMockOIDCClient("https://github.com", "gh_client", &state3, "github_user_123", "github_tester", "tester@github.com", "GitHub Tester") - util.SetHTTPClient(httpMock3) - router3 := setupTestRouter(dbConn, mockRedis, httpMock3) - - // Generate state3 and SessionHash using activeCookie2 - wAuth3 := performRequest(router3, http.MethodGet, "/api/v1/oauth/github/authorize?purpose=bind", nil, nil, []*http.Cookie{activeCookie2}) - if wAuth3.Code != http.StatusOK { - t.Fatalf("authorize failed: %d, body: %s", wAuth3.Code, wAuth3.Body.String()) - } - // Extract the cookie to get the updated session token - for _, cookie := range wAuth3.Result().Cookies() { - if cookie.Name == config.Config.App.SessionCookieName { - activeCookie2 = cookie - break - } - } - var authResp3 struct { - Data OAuthAuthorizeResponse `json:"data"` - } - _ = json.Unmarshal(wAuth3.Body.Bytes(), &authResp3) - parsedURL3, _ := url.Parse(authResp3.Data.AuthorizeURL) - state3 = parsedURL3.Query().Get("state") - - reqBody3 := fmt.Sprintf(`{"state":"%s","code":"test_code"}`, state3) - w3 := performRequest(router3, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody3), map[string]string{ - "Content-Type": "application/json", - }, []*http.Cookie{activeCookie2}) - - if w3.Code != http.StatusBadRequest { - t.Errorf("expected 400 for already bound account, got %d, body: %s", w3.Code, w3.Body.String()) - } - -} - -func TestExternalAccountsListAndDelete(t *testing.T) { - initializeTestConfig() - dbConn := setupTestDB(t) - mockRedis := newMockRedisClient() - httpMock := &http.Client{} - router := setupTestRouter(dbConn, mockRedis, httpMock) - - // Create user and external accounts - dbConn.Create(&model.User{ - ID: 555, - Username: "account_holder", - IsActive: true, - }) - - dbConn.Create(&model.AuthSource{ - ID: 10, - Name: "gitlab", - Type: model.AuthSourceTypeOIDC, - IsActive: true, - }) - - dbConn.Create(&model.ExternalAccount{ - ID: 2001, - AuthSourceID: 10, - UserID: 555, - ExternalID: "gitlab_123", - ExternalUsername: "gitlab_user", - }) - - router.GET("/test-helper/login-555", func(c *gin.Context) { - session := sessions.Default(c) - session.Set(UserIDKey, uint64(555)) - _ = session.Save() - c.String(200, "ok") - }) - - wLogin := performRequest(router, http.MethodGet, "/test-helper/login-555", nil, nil, nil) - var activeCookie *http.Cookie - for _, cookie := range wLogin.Result().Cookies() { - if cookie.Name == config.Config.App.SessionCookieName { - activeCookie = cookie - break - } - } - - // 1. List accounts - wList := performRequest(router, http.MethodGet, "/api/v1/oauth/external-accounts", nil, nil, []*http.Cookie{activeCookie}) - if wList.Code != http.StatusOK { - t.Fatalf("failed to list external accounts: %d", wList.Code) - } - - var listResp struct { - Data []model.ExternalAccountView `json:"data"` - } - _ = json.Unmarshal(wList.Body.Bytes(), &listResp) - if len(listResp.Data) != 1 || listResp.Data[0].ExternalUsername != "gitlab_user" { - t.Errorf("unexpected list response: %+v", listResp.Data) - } - - // 2. Delete/Unbind account - wDelete := performRequest(router, http.MethodPost, "/api/v1/oauth/external-accounts/2001/delete", nil, nil, []*http.Cookie{activeCookie}) - if wDelete.Code != http.StatusOK { - t.Fatalf("failed to delete external account binding: %d, body: %s", wDelete.Code, wDelete.Body.String()) - } - - var count int64 - dbConn.Model(&model.ExternalAccount{}).Where("id = ?", 2001).Count(&count) - if count != 0 { - t.Error("binding record was not deleted from DB") - } -} - -func TestOIDCPolicyEnforcement(t *testing.T) { - initializeTestConfig() - dbConn := setupTestDB(t) - mockRedis := newMockRedisClient() - seedTestAuthSource(t, dbConn) // seeds testSourceName ("linuxdo") active=true - - // Set up mock client & router - var state string - httpMock := newMockOIDCClient(testIssuerURL, testClientID, &state, "88888", "test_oauth_user", "oauth@linux.do", "Oauth Test User") - util.SetHTTPClient(httpMock) - router := setupTestRouter(dbConn, mockRedis, httpMock) - - // --- 1. Test GetLoginURL enforcement --- - // Disable globally - dbConn.Create(&model.SystemConfig{ - Key: model.ConfigKeyOIDCLoginEnabled, - Value: "false", - }) - repository.ResetSystemConfigRAMCacheForTest() - mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) - wLoginDisabled := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) - if wLoginDisabled.Code != http.StatusBadRequest { - t.Errorf("expected 400 when OIDC globally disabled, got %d", wLoginDisabled.Code) - } - - // Re-enable globally, but deactivate source - dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true") - repository.ResetSystemConfigRAMCacheForTest() - mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) - dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false) - _ = repository.InvalidateAuthSourceCache(context.Background()) - - wSourceInactive := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) - if wSourceInactive.Code != http.StatusBadRequest { - t.Errorf("expected 400 when OIDC source is inactive, got %d", wSourceInactive.Code) - } - - // --- 2. Test Authorize enforcement --- - // Deactivate globally again - dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false") - repository.ResetSystemConfigRAMCacheForTest() - mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) - dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true) - _ = repository.InvalidateAuthSourceCache(context.Background()) - - wAuthDisabled := performRequest(router, http.MethodGet, "/api/v1/oauth/"+testSourceName+"/authorize", nil, nil, nil) - if wAuthDisabled.Code != http.StatusBadRequest { - t.Errorf("expected 400 when OIDC globally disabled in Authorize, got %d", wAuthDisabled.Code) - } - - // --- 3. Test Callback enforcement --- - // Set up a valid state beforehand (when enabled) - dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true") - repository.ResetSystemConfigRAMCacheForTest() - mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) - dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true) - _ = repository.InvalidateAuthSourceCache(context.Background()) - - wLogin := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) - - if wLogin.Code != http.StatusOK { - t.Fatalf("failed to setup login: %s", wLogin.Body.String()) - } - var loginUrlResp struct { - Data OAuthAuthorizeResponse `json:"data"` - } - _ = json.Unmarshal(wLogin.Body.Bytes(), &loginUrlResp) - parsedURL, _ := url.Parse(loginUrlResp.Data.AuthorizeURL) - state = parsedURL.Query().Get("state") - - var anonymousCookie *http.Cookie - for _, cookie := range wLogin.Result().Cookies() { - if cookie.Name == config.Config.App.SessionCookieName { - anonymousCookie = cookie - break - } - } - - // Now disable OIDC globally and attempt callback - dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false") - repository.ResetSystemConfigRAMCacheForTest() - mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) - reqBody := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state) - wCallbackDisabled := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody), map[string]string{ - "Content-Type": "application/json", - }, []*http.Cookie{anonymousCookie}) - if wCallbackDisabled.Code != http.StatusBadRequest { - t.Errorf("expected 400 for callback when OIDC globally disabled, got %d, body: %s", wCallbackDisabled.Code, wCallbackDisabled.Body.String()) - } - - // Enable globally but deactivate source and attempt callback - dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true") - repository.ResetSystemConfigRAMCacheForTest() - mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) - dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false) - - // Since callback deletes state, we need to generate state again - dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true) - _ = repository.InvalidateAuthSourceCache(context.Background()) - wLogin2 := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) - _ = json.Unmarshal(wLogin2.Body.Bytes(), &loginUrlResp) - parsedURL, _ = url.Parse(loginUrlResp.Data.AuthorizeURL) - state = parsedURL.Query().Get("state") - var anonymousCookie2 *http.Cookie - for _, cookie := range wLogin2.Result().Cookies() { - if cookie.Name == config.Config.App.SessionCookieName { - anonymousCookie2 = cookie - break - } - } - - // Deactivate source - dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false) - _ = repository.InvalidateAuthSourceCache(context.Background()) - reqBody2 := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state) - wCallbackSourceInactive := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody2), map[string]string{ - "Content-Type": "application/json", - }, []*http.Cookie{anonymousCookie2}) - if wCallbackSourceInactive.Code != http.StatusBadRequest { - t.Errorf("expected 400 for callback when OIDC source deactivated, got %d, body: %s", wCallbackSourceInactive.Code, wCallbackSourceInactive.Body.String()) - } -} - -func TestSystemUserBlockedByMiddleware(t *testing.T) { - initializeTestConfig() - dbConn := setupTestDB(t) - - // 1. 创建正常管理员 - adminUser := &model.User{ID: 1001, Username: "normal_admin", IsAdmin: true, IsActive: true} - err := dbConn.Create(adminUser).Error - require.NoError(t, err) - - // 2. 创建系统用户 (根据架构设计,系统用户 id = 999) - systemUser := &model.User{ID: 999, Username: "system", Nickname: "系统", Password: "*", IsActive: true} - err = dbConn.Create(systemUser).Error - require.NoError(t, err) - - // 3. 设置全局测试数据库连接并构建测试路由组 - db.SetDB(dbConn) - rProtected := testhelper.NewTestGinEngine() - store := cookie.NewStore([]byte("secret")) - rProtected.Use(sessions.Sessions("mysession", store)) - rProtected.Use(LoginRequired()) - rProtected.GET("/test-auth", func(c *gin.Context) { - c.JSON(http.StatusOK, gin.H{"status": "ok"}) - }) - - // 4. 测试未登录用户 (401) - w1 := httptest.NewRecorder() - req1, _ := http.NewRequest("GET", "/test-auth", nil) - rProtected.ServeHTTP(w1, req1) - assert.Equal(t, http.StatusUnauthorized, w1.Code) - - // 5. 测试正常用户登录并访问 (200) - rLogin := gin.New() - rLogin.Use(sessions.Sessions("mysession", store)) - rLogin.GET("/login-mock", func(c *gin.Context) { - session := sessions.Default(c) - session.Set("user_id", uint64(1001)) - _ = session.Save() - c.Status(200) - }) - - wLogin := httptest.NewRecorder() - reqLogin, _ := http.NewRequest("GET", "/login-mock", nil) - rLogin.ServeHTTP(wLogin, reqLogin) - cookieStr := wLogin.Header().Get("Set-Cookie") - - w2 := httptest.NewRecorder() - req2, _ := http.NewRequest("GET", "/test-auth", nil) - req2.Header.Set("Cookie", cookieStr) - rProtected.ServeHTTP(w2, req2) - assert.Equal(t, http.StatusOK, w2.Code) - - // 6. 测试 system 用户(ID: 999)登录并访问 (被中间件阻断返回 401) - rLoginSystem := gin.New() - rLoginSystem.Use(sessions.Sessions("mysession", store)) - rLoginSystem.GET("/login-system-mock", func(c *gin.Context) { - session := sessions.Default(c) - session.Set("user_id", uint64(999)) - _ = session.Save() - c.Status(200) - }) - - wLoginSystem := httptest.NewRecorder() - reqLoginSystem, _ := http.NewRequest("GET", "/login-system-mock", nil) - rLoginSystem.ServeHTTP(wLoginSystem, reqLoginSystem) - cookieSystemStr := wLoginSystem.Header().Get("Set-Cookie") - - w3 := httptest.NewRecorder() - req3, _ := http.NewRequest("GET", "/test-auth", nil) - req3.Header.Set("Cookie", cookieSystemStr) - rProtected.ServeHTTP(w3, req3) - assert.Equal(t, http.StatusUnauthorized, w3.Code) -} diff --git a/internal/apps/oauth/oauth_types.go b/internal/apps/oauth/oauth_types.go deleted file mode 100644 index 17ab267c..00000000 --- a/internal/apps/oauth/oauth_types.go +++ /dev/null @@ -1,36 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -// 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 响应 -// -//nolint:revive // OAuth 前缀保持包内语义清晰 -type OAuthAuthorizeResponse struct { - AuthorizeURL string `json:"authorize_url"` -} - -// OAuthCallbackResult 回调处理结果 -// -//nolint:revive // OAuth 前缀保持包内语义清晰 -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"` -} diff --git a/internal/apps/oauth/oauth_userinfo.go b/internal/apps/oauth/oauth_userinfo.go deleted file mode 100644 index c8b9c657..00000000 --- a/internal/apps/oauth/oauth_userinfo.go +++ /dev/null @@ -1,139 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -import ( - "context" - "errors" - "fmt" - "strings" - - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/coreos/go-oidc/v3/oidc" - "golang.org/x/oauth2" -) - -func uniqueUsername(ctx context.Context, base string) (string, error) { - base = strings.TrimSpace(base) - if base == "" { - base = "user" - } - - existingUsernames, err := repository.ListUsernamesMatchingBase(ctx, base) - if err != nil { - return "", err - } - - // 将现有的用户名放入 map 中,以便 O(1) 查找 - exists := make(map[string]bool, len(existingUsernames)) - for _, u := range existingUsernames { - exists[strings.ToLower(u)] = true - } - - // 检查 base 是否被占用 - 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 buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code string, nonce string, redirectURL string) (*model.OAuthUserInfo, 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 := &model.OAuthUserInfo{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 -} - -// verifyIDToken 验证 OIDC ID Token 并将 Claims 解析到 userInfo -func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *model.OAuthUserInfo) 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 *model.OAuthUserInfo) 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 *model.User, status string) OAuthCallbackResult { - result := OAuthCallbackResult{Status: status} - if user != nil { - info := BuildBasicUserInfo(user, false) - result.User = &info - } - return result -} diff --git a/internal/apps/oauth/provider_cache.go b/internal/apps/oauth/provider_cache.go deleted file mode 100644 index 814e05c2..00000000 --- a/internal/apps/oauth/provider_cache.go +++ /dev/null @@ -1,99 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -import ( - "context" - "net/http" - "sync" - - "github.com/coreos/go-oidc/v3/oidc" - "golang.org/x/oauth2" - "golang.org/x/sync/singleflight" -) - -// oidcProviderCache 进程级 OIDC provider 缓存。 -// -// oidc.NewProvider 每次调用都会向远端 issuer 的 -// /.well-known/openid-configuration 发起 HTTP 请求拉取元数据。 -// 由于 provider 元数据极少变动,将其缓存后可消除登录发起与回调时的 -// 重复外部 HTTP 往返。 -// -// 并发安全性: -// - mu + entries 防止并发读写 map。 -// - sfGroup 保证同一 issuer 同时只有一次在途的 NewProvider 调用 -// (singleflight),后续等待者复用同一结果,彻底消除 thundering herd。 -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), -} - -// discoveryContext 从请求 ctx 提取 HTTP 客户端,并绑定到不可取消的 Background ctx。 -// 这样既能在测试中注入 mock 客户端,又避免请求取消导致 provider 拉取失败。 -func discoveryContext(ctx context.Context) context.Context { - bg := context.Background() - if client, ok := ctx.Value(oauth2.HTTPClient).(*http.Client); ok && client != nil { - bg = oidc.ClientContext(bg, client) - } - return bg -} - -// get 返回缓存的 provider;若无则通过 oidc.NewProvider 获取并写入缓存。 -// 同一 issuer 并发调用时,singleflight 保证只有一次实际 HTTP 请求。 -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() - return p, nil - } - c.mu.RUnlock() - - // 慢路径:通过 singleflight 合并并发的首次请求。 - discCtx := discoveryContext(ctx) - v, err, _ := c.sfGroup.Do(issuer, func() (any, error) { - // 双检:singleflight 内再次检查,前一个并发组可能已写入缓存。 - c.mu.RLock() - if p, ok := c.entries[issuer]; ok { - c.mu.RUnlock() - return p, nil - } - c.mu.RUnlock() - - p, err := oidc.NewProvider(discCtx, issuer) - if err != nil { - return nil, err - } - - c.mu.Lock() - c.entries[issuer] = p - c.mu.Unlock() - return p, nil - }) - if err != nil { - return nil, err - } - return v.(*oidc.Provider), nil //nolint:forcetypeassert // singleflight value 由同函数写入,类型确定 -} - -// invalidate 从缓存中移除指定 issuer 对应的 provider。 -// 在认证源的 Discovery URL 被修改时调用,强制下次请求重新拉取元数据。 -func (c *oidcProviderCache) invalidate(issuer string) { - c.mu.Lock() - delete(c.entries, issuer) - c.mu.Unlock() -} - -// InvalidateOIDCProviderCache 从进程级缓存中清除指定 issuer 的 provider 条目。 -// 当管理员更新认证源的 Discovery URL 后调用,以确保下次登录时重新拉取最新元数据。 -// issuer 值应为去掉 /.well-known/openid-configuration 后缀的规范化 URL。 -func InvalidateOIDCProviderCache(issuer string) { - globalOIDCProviderCache.invalidate(issuer) -} diff --git a/internal/apps/oauth/routers.go b/internal/apps/oauth/routers.go deleted file mode 100644 index a5a44799..00000000 --- a/internal/apps/oauth/routers.go +++ /dev/null @@ -1,113 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -import ( - "net/http" - - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/pkg/logger" - "github.com/gin-contrib/sessions" - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -// BasicUserInfo 用户基本信息结构体 -type BasicUserInfo struct { - ID uint64 `json:"id"` - 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 将 User 模型转换为 BasicUserInfo -func BuildBasicUserInfo(user *model.User, needChange bool) BasicUserInfo { - return BasicUserInfo{ - ID: user.ID, - Username: user.Username, - Nickname: user.Nickname, - Email: user.Email, - AvatarURL: user.AvatarURL, - IsAdmin: user.IsAdmin, - NeedChangePassword: needChange, - Bio: user.Bio, - Phone: user.Phone, - Gender: user.Gender, - Website: user.Website, - Location: user.Location, - } -} - -// UserInfo 获取当前登录用户信息 -// @Summary 获取当前登录用户信息 -// @Description 返回当前登录用户的基本信息及余额数据,需要登录。包括用户 ID、用户名、信任等级、各类余额信息等。 -// @Tags oauth -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=oauth.BasicUserInfo} "用户信息" -// @Failure 401 {object} response.Any "未登录" -// @Router /api/v1/oauth/user-info [get] -// @Router /api/v1/user-info [get] -// @Router /api/v1/user/self [get] -func UserInfo(c *gin.Context) { - user, _ := GetFromContext[*model.User](c, UserObjKey) - session := sessions.Default(c) - needChange := session.Get("need_change_password") == true - - c.JSON( - http.StatusOK, - response.OK(BuildBasicUserInfo(user, needChange)), - ) -} - -// GetLoginURL 获取登录地址 -// @Summary 获取登录地址 -// @Description 生成 OAuth 登录 URL,前端跳转至该地址完成授权。返回的 URL 中包含 state 参数用于 CSRF 防护。 -// @Tags oauth -// @Produce json -// @Success 200 {object} response.Any{data=string} "OAuth 登录 URL" -// @Failure 500 {object} response.Any "Redis 异常或内部错误" -// @Router /api/v1/oauth/login [get] - -// 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, ok := userID.(uint64); ok { - InvalidateCachedUser(c.Request.Context(), id) - } else if idFloat, ok := userID.(float64); ok { - InvalidateCachedUser(c.Request.Context(), uint64(idFloat)) - } else if idInt, ok := userID.(int); ok && idInt >= 0 { - InvalidateCachedUser(c.Request.Context(), uint64(idInt)) - } - } - session.Options(GetSessionOptions(-1)) - session.Clear() - if err := session.Save(); err != nil { - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OKNil()) -} diff --git a/internal/apps/oauth/session.go b/internal/apps/oauth/session.go deleted file mode 100644 index 0aeb1486..00000000 --- a/internal/apps/oauth/session.go +++ /dev/null @@ -1,53 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package oauth provides authentication and OAuth integration. -package oauth - -import ( - "net/http" - "strings" - - "github.com/Rain-kl/Wavelet/internal/infra/config" - "github.com/gin-contrib/sessions" -) - -// GetSessionOptions 根据配置构建 Session 选项 -func GetSessionOptions(maxAge int) sessions.Options { - return sessions.Options{ - Path: "/", - Domain: config.Config.App.SessionDomain, - MaxAge: maxAge, - HttpOnly: config.Config.App.SessionHTTPOnly, - Secure: config.Config.App.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 -} diff --git a/internal/apps/oauth/session_context.go b/internal/apps/oauth/session_context.go deleted file mode 100644 index 4e73a0d6..00000000 --- a/internal/apps/oauth/session_context.go +++ /dev/null @@ -1,101 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package oauth - -import ( - "context" - "crypto/sha256" - "encoding/hex" - - "github.com/Rain-kl/Wavelet/internal/infra/config" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/gin-contrib/sessions" - "github.com/gin-gonic/gin" - "github.com/google/uuid" - gsessions "github.com/gorilla/sessions" -) - -// GetUserIDFromSession 从 Session 中提取用户 ID -func GetUserIDFromSession(s sessions.Session) uint64 { - userID, ok := s.Get(UserIDKey).(uint64) - if !ok { - return 0 - } - return userID -} - -// GetUserIDFromContext 从 Gin Context 的 Session 中提取用户 ID -func GetUserIDFromContext(c *gin.Context) uint64 { - 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 *model.User, extras ...map[string]any) error { - session := sessions.Default(c) - session.Clear() - rotateSessionID(session) - - session.Set(UserIDKey, user.ID) - session.Set(UserNameKey, user.Username) - session.Set(PasswordHashKey, user.Password) - if len(extras) > 0 { - for key, value := range extras[0] { - session.Set(key, value) - } - } - - // 根据系统配置动态设置 Session 过期时间 - maxAge := config.Config.App.SessionAge - isSessionCookie := false - - ttlHours, err := repository.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours) - if 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(), config.Config.App.SessionCookieName) - } - - return nil -} diff --git a/internal/apps/risk_control/logics.go b/internal/apps/risk_control/logics.go deleted file mode 100644 index e69aac6e..00000000 --- a/internal/apps/risk_control/logics.go +++ /dev/null @@ -1,145 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package risk_control - -import ( - "context" - "sync" - - "time" - - "github.com/Rain-kl/Wavelet/internal/infra/persistence/batchwriter" - "github.com/Rain-kl/Wavelet/internal/model/analytics" - "github.com/Rain-kl/Wavelet/internal/platform/lifecycle" - "github.com/Rain-kl/Wavelet/internal/repository/logstore" - "github.com/Rain-kl/Wavelet/pkg/logger" -) - -var ( - logWriterMu sync.RWMutex - logWriter *batchwriter.Writer[*analytics.UserAccessLog] -) - -// InitLogWriter initializes the access-log batch writer for the active log database. -func InitLogWriter(ctx context.Context) { - logWriterMu.Lock() - defer logWriterMu.Unlock() - if logWriter != nil { - return - } - - cfg := batchwriter.DefaultConfig() - writer, err := batchwriter.New[*analytics.UserAccessLog](cfg, func(ctx context.Context, items []*analytics.UserAccessLog) error { - rows := make([]analytics.UserAccessLog, 0, len(items)) - for _, item := range items { - if item == nil { - continue - } - rows = append(rows, *item) - } - store, err := logstore.Active(ctx) - if err != nil { - return err - } - return store.UserAccessLogs.BatchInsert(ctx, rows) - }, - batchwriter.WithDropHandler[*analytics.UserAccessLog](func(item *analytics.UserAccessLog) { - path := "" - if item != nil { - path = item.Path - } - logger.WarnF(context.Background(), "[RiskControl] Log queue full, dropping log item for path: %s", path) - }), - batchwriter.WithFlushErrorHandler[*analytics.UserAccessLog](func(ctx context.Context, items []*analytics.UserAccessLog, err error) { - logger.ErrorF(ctx, "[RiskControl] flush access-log batch failed (batch=%d): %v", len(items), err) - }), - ) - if err != nil { - logger.ErrorF(ctx, "[RiskControl] init log writer failed: %v", err) - return - } - - writer.Start(ctx) - logWriter = writer - lifecycle.OnShutdown("risk_control_log_writer", StopLogWriter) -} - -// StopLogWriter stops the ClickHouse access-log batch writer and drains pending logs. -func StopLogWriter(ctx context.Context) error { - writer := currentLogWriter() - if writer == nil { - return nil - } - return writer.Stop(ctx) -} - -// IsBufferFull reports whether the access-log queue has no remaining capacity. -func IsBufferFull() bool { - writer := currentLogWriter() - if writer == nil { - return false - } - return writer.IsFull() -} - -// QueueAccessLog enqueues an access log without blocking. -func QueueAccessLog(logItem *analytics.UserAccessLog) { - writer := currentLogWriter() - if writer == nil || logItem == nil { - return - } - writer.TryEnqueue(logItem) -} - -// SetLogWriterForTest swaps the access-log writer for unit tests. -func SetLogWriterForTest(writer *batchwriter.Writer[*analytics.UserAccessLog]) func() { - logWriterMu.Lock() - previous := logWriter - logWriter = writer - logWriterMu.Unlock() - return func() { - logWriterMu.Lock() - logWriter = previous - logWriterMu.Unlock() - } -} - -func currentLogWriter() *batchwriter.Writer[*analytics.UserAccessLog] { - logWriterMu.RLock() - defer logWriterMu.RUnlock() - return logWriter -} - -const drainPollInterval = 50 * time.Millisecond - -// Drain waits until the in-memory access-log queue has been empty for one flush interval. -func Drain(ctx context.Context) error { - writer := currentLogWriter() - if writer == nil { - return nil - } - quietPeriod := batchwriter.DefaultConfig().FlushInterval - if quietPeriod <= 0 { - quietPeriod = time.Second - } - ticker := time.NewTicker(drainPollInterval) - defer ticker.Stop() - var quietSince time.Time - for { - if writer.Len() == 0 { - if quietSince.IsZero() { - quietSince = time.Now() - } else if time.Since(quietSince) >= quietPeriod { - return nil - } - } else { - quietSince = time.Time{} - } - select { - case <-ctx.Done(): - return ctx.Err() - case <-ticker.C: - } - } -} diff --git a/internal/apps/risk_control/middleware.go b/internal/apps/risk_control/middleware.go deleted file mode 100644 index e330751d..00000000 --- a/internal/apps/risk_control/middleware.go +++ /dev/null @@ -1,88 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package risk_control 提供风险控制中间件 -package risk_control - -import ( - "encoding/json" - "net/http" - "time" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/infra/config" - "github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/model/analytics" - "github.com/Rain-kl/Wavelet/internal/shared/response" - "github.com/gin-gonic/gin" -) - -// RiskControlMiddleware 全局日志采集中间件 -func RiskControlMiddleware() gin.HandlerFunc { - return func(c *gin.Context) { - // 如果未启用 ClickHouse,直接放行 - if !config.Config.ClickHouse.Enabled { - c.Next() - return - } - - // 1. 限流背压检测(检测本地缓冲队列是否已满) - if IsBufferFull() { - response.AbortTooManyRequests(c, "系统繁忙,请稍后再试") - return - } - - start := time.Now() - - // 2. 执行后续请求(穿过业务处理和认证中间件) - c.Next() - - // 3. 后置身份检查:仅记录通过认证的请求 - userObj, exists := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) - if !exists || userObj == nil { - return - } - - // 4. 计算耗时并异步推送到缓冲队列 - latency := time.Since(start).Milliseconds() - - var headersStr string - if c.Request.Header != nil { - // 克隆 Header,避免污染原 HTTP 请求的 Header 对象 - clonedHeaders := make(http.Header) - for k, v := range c.Request.Header { - clonedHeaders[k] = v - } - clonedHeaders.Del("Cookie") - - if headersBytes, err := json.Marshal(clonedHeaders); err == nil { - headersStr = string(headersBytes) - } - } - - const maxHTTPStatus = 999 - status := c.Writer.Status() - if status < 0 { - status = 0 - } else if status > maxHTTPStatus { - status = maxHTTPStatus - } - - logItem := &analytics.UserAccessLog{ - ID: idgen.NextUint64ID(), - UserID: userObj.ID, // 直接从 Context 获取已登录用户ID,避免数据库查询 - Path: c.Request.URL.Path, - Method: c.Request.Method, - IP: c.ClientIP(), - UserAgent: c.Request.UserAgent(), - Headers: headersStr, - Status: int32(status), - Latency: latency, - CreatedAt: time.Now(), - } - - // 非阻塞地推入缓存队列 - QueueAccessLog(logItem) - } -} diff --git a/internal/apps/risk_control/middleware_test.go b/internal/apps/risk_control/middleware_test.go deleted file mode 100644 index 60f33b2e..00000000 --- a/internal/apps/risk_control/middleware_test.go +++ /dev/null @@ -1,197 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package risk_control - -import ( - "context" - "encoding/json" - "net/http" - "net/http/httptest" - "sync" - "testing" - "time" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/infra/config" - "github.com/Rain-kl/Wavelet/internal/infra/persistence/batchwriter" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/model/analytics" - "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/gin-gonic/gin" - "github.com/stretchr/testify/assert" -) - -func newTestAccessLogWriter(t *testing.T, cfg batchwriter.Config) (*batchwriter.Writer[*analytics.UserAccessLog], func() []*analytics.UserAccessLog) { - t.Helper() - - var ( - mu sync.Mutex - captured []*analytics.UserAccessLog - ) - writer, err := batchwriter.New(cfg, func(_ context.Context, items []*analytics.UserAccessLog) error { - mu.Lock() - captured = append(captured, items...) - mu.Unlock() - return nil - }) - if err != nil { - t.Fatalf("batchwriter.New() error = %v", err) - } - - writer.Start(context.Background()) - restore := SetLogWriterForTest(writer) - t.Cleanup(func() { - restore() - stopCtx, cancel := context.WithTimeout(context.Background(), time.Second) - defer cancel() - _ = writer.Stop(stopCtx) - }) - - return writer, func() []*analytics.UserAccessLog { - mu.Lock() - defer mu.Unlock() - return append([]*analytics.UserAccessLog(nil), captured...) - } -} - -func drainAccessLogWriter(t *testing.T, writer *batchwriter.Writer[*analytics.UserAccessLog]) { - t.Helper() - - stopCtx, cancel := context.WithTimeout(context.Background(), time.Second) - defer cancel() - if err := writer.Stop(stopCtx); err != nil { - t.Fatalf("writer.Stop() error = %v", err) - } -} - -func TestRiskControlMiddleware(t *testing.T) { - gin.SetMode(gin.TestMode) - - t.Run("ClickHouse disabled", func(t *testing.T) { - config.Config.ClickHouse.Enabled = false - defer func() { config.Config.ClickHouse.Enabled = false }() - - r := testhelper.NewTestGinEngine(RiskControlMiddleware()) - r.GET("/test", func(c *gin.Context) { - c.String(http.StatusOK, "ok") - }) - - w := httptest.NewRecorder() - req, _ := http.NewRequest(http.MethodGet, "/test", nil) - r.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - assert.Equal(t, "ok", w.Body.String()) - }) - - t.Run("ClickHouse enabled - Normal Authenticated Request", func(t *testing.T) { - config.Config.ClickHouse.Enabled = true - defer func() { config.Config.ClickHouse.Enabled = false }() - - cfg := batchwriter.DefaultConfig() - cfg.MaxBatchSize = 100 - cfg.FlushInterval = time.Hour - - writer, getCaptured := newTestAccessLogWriter(t, cfg) - - r := gin.New() - r.Use(func(c *gin.Context) { - user := &model.User{ID: 12345} - oauth.SetToContext(c, oauth.UserObjKey, user) - c.Next() - }) - r.Use(RiskControlMiddleware()) - r.GET("/test", func(c *gin.Context) { - c.String(http.StatusOK, "ok") - }) - - w := httptest.NewRecorder() - req, _ := http.NewRequest(http.MethodGet, "/test", nil) - req.Header.Set("X-Test-Header", "hello") - req.Header.Set("Cookie", "session_id=abcdef123456") - r.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - assert.Equal(t, "ok", w.Body.String()) - - drainAccessLogWriter(t, writer) - - captured := getCaptured() - if len(captured) != 1 { - t.Fatalf("captured access logs = %d, want 1", len(captured)) - } - logItem := captured[0] - assert.Equal(t, uint64(12345), logItem.UserID) - assert.Equal(t, "/test", logItem.Path) - assert.Equal(t, http.MethodGet, logItem.Method) - assert.Equal(t, int32(http.StatusOK), logItem.Status) - assert.NotEmpty(t, logItem.Headers) - assert.Contains(t, logItem.Headers, "X-Test-Header") - assert.NotContains(t, logItem.Headers, "Cookie") - }) - - t.Run("ClickHouse enabled - Unauthenticated Request", func(t *testing.T) { - config.Config.ClickHouse.Enabled = true - defer func() { config.Config.ClickHouse.Enabled = false }() - - cfg := batchwriter.DefaultConfig() - cfg.MaxBatchSize = 100 - cfg.FlushInterval = time.Hour - - writer, getCaptured := newTestAccessLogWriter(t, cfg) - - r := testhelper.NewTestGinEngine(RiskControlMiddleware()) - r.GET("/test", func(c *gin.Context) { - c.String(http.StatusOK, "ok") - }) - - w := httptest.NewRecorder() - req, _ := http.NewRequest(http.MethodGet, "/test", nil) - r.ServeHTTP(w, req) - - assert.Equal(t, http.StatusOK, w.Code) - assert.Equal(t, "ok", w.Body.String()) - - drainAccessLogWriter(t, writer) - - if len(getCaptured()) != 0 { - t.Fatal("expected no log item for unauthenticated request") - } - }) - - t.Run("ClickHouse enabled - Buffer Full Rate Limiting", func(t *testing.T) { - config.Config.ClickHouse.Enabled = true - defer func() { config.Config.ClickHouse.Enabled = false }() - - cfg := batchwriter.DefaultConfig() - cfg.QueueSize = 2 - cfg.MaxBatchSize = 100 - cfg.FlushInterval = time.Hour - - writer, _ := newTestAccessLogWriter(t, cfg) - - for range cfg.QueueSize { - writer.TryEnqueue(&analytics.UserAccessLog{}) - } - if !IsBufferFull() { - t.Fatal("IsBufferFull() = false, want true") - } - - r := testhelper.NewTestGinEngine(RiskControlMiddleware()) - r.GET("/test", func(c *gin.Context) { - c.String(http.StatusOK, "ok") - }) - - w := httptest.NewRecorder() - req, _ := http.NewRequest(http.MethodGet, "/test", nil) - r.ServeHTTP(w, req) - - assert.Equal(t, http.StatusTooManyRequests, w.Code) - - var resp map[string]interface{} - err := json.Unmarshal(w.Body.Bytes(), &resp) - assert.NoError(t, err) - assert.Contains(t, resp["error_msg"], "系统繁忙") - }) -} diff --git a/internal/apps/user/access_tokens.go b/internal/apps/user/access_tokens.go deleted file mode 100644 index 047ea420..00000000 --- a/internal/apps/user/access_tokens.go +++ /dev/null @@ -1,193 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package user 提供用户认证与帐户管理功能 -package user - -import ( - "net/http" - "strconv" - "strings" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -type createTokenRequest struct { - Name string `json:"name"` - IsAdmin bool `json:"is_admin"` -} - -type tokenResponse struct { - Token string `json:"token"` - Record model.AccessToken `json:"record"` -} - -// ListAccessTokens 获取当前用户的 AccessToken 列表 -// @Summary 获取当前用户的 AccessToken 列表 -// @Description 返回当前登录用户的所有 active access tokens(脱敏后) -// @Tags user -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=[]model.AccessToken} "令牌列表" -// @Failure 401 {object} response.Any "未登录" -// @Router /api/v1/user/access-tokens [get] -// ListAccessTokens 获取当前用户的 AccessToken 列表 -func ListAccessTokens(c *gin.Context) { - currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) - ctx := c.Request.Context() - - tokens, err := listAccessTokensLogic(ctx, currUser.ID) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - c.JSON(http.StatusOK, response.OK(tokens)) -} - -// CreateAccessToken 创建一个新的 AccessToken -// @Summary 创建一个新的 AccessToken -// @Description 为当前用户新建一个 API 访问令牌,仅在此接口返回一次明文令牌值,请妥善保存。可通过 is_admin 字段赋予令牌管理员权限(仅管理员用户可设置)。 -// @Tags user -// @Accept json -// @Produce json -// @Param request body user.createTokenRequest true "令牌名称" -// @Security SessionCookie -// @Success 200 {object} response.Any{data=user.tokenResponse} "新建令牌成功" -// @Failure 400 {object} response.Any "参数错误或超限" -// @Router /api/v1/user/access-tokens [post] -func CreateAccessToken(c *gin.Context) { - currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) - ctx := c.Request.Context() - - var req createTokenRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, errBindParamsFailed) - return - } - - req.Name = strings.TrimSpace(req.Name) - if req.Name == "" { - response.AbortBadRequest(c, errTokenNameRequired) - return - } - - // 只有管理员才能创建具有管理员权限的令牌 - if req.IsAdmin && !currUser.IsAdmin { - response.AbortBadRequest(c, errAdminTokenRequiresAdmin) - return - } - - // 检查最大限制(基于 ConfigKeyMaxAPIKeysPerUser 配置,默认值为 5) - maxLimit := 5 - if val, err := repository.GetIntByKey(ctx, model.ConfigKeyMaxAPIKeysPerUser); err == nil { - maxLimit = val - } - - count, err := countAccessTokensLogic(ctx, currUser.ID) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - if int(count) >= maxLimit { - response.AbortBadRequest(c, errAccessTokenLimitReached) - return - } - - // 生成 Token - tokenStr, err := model.GenerateTokenString() - if err != nil { - response.AbortBadRequest(c, errGenerateTokenFailed) - return - } - - tokenHash := model.HashToken(tokenStr) - maskedToken := model.MaskTokenString(tokenStr) - - tokenRecord := model.AccessToken{ - UserID: currUser.ID, - Name: req.Name, - TokenHash: tokenHash, - MaskedToken: maskedToken, - IsAdmin: req.IsAdmin, - } - - if err := createAccessTokenLogic(ctx, &tokenRecord); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - c.JSON(http.StatusOK, response.OK(tokenResponse{ - Token: tokenStr, - Record: tokenRecord, - })) -} - -// DeleteAccessToken 删除一个 AccessToken -// @Summary 删除一个 AccessToken -// @Description 撤销并删除一个属于当前用户的 API 访问令牌 -// @Tags user -// @Produce json -// @Param id path string true "令牌ID" -// @Security SessionCookie -// @Success 200 {object} response.Any{data=string} "删除成功" -// @Failure 400 {object} response.Any "参数错误" -// @Router /api/v1/user/access-tokens/{id} [delete] -func DeleteAccessToken(c *gin.Context) { - currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) - ctx := c.Request.Context() - - idStr := c.Param("id") - id, err := strconv.ParseUint(idStr, 10, 64) - if err != nil { - response.AbortBadRequest(c, errInvalidTokenID) - return - } - - if err := deleteAccessTokenLogic(ctx, id, currUser.ID); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - c.JSON(http.StatusOK, response.OK("删除成功")) -} - -// RotateAccessToken 轮换一个 AccessToken -// @Summary 轮换一个 AccessToken -// @Description 轮换(重新生成)一个属于当前用户的 API 访问令牌的密钥,旧令牌将立即失效 -// @Tags user -// @Produce json -// @Param id path string true "令牌ID" -// @Security SessionCookie -// @Success 200 {object} response.Any{data=user.tokenResponse} "令牌轮换成功" -// @Failure 400 {object} response.Any "参数错误" -// @Router /api/v1/user/access-tokens/{id}/rotate [post] -func RotateAccessToken(c *gin.Context) { - currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) - ctx := c.Request.Context() - - idStr := c.Param("id") - id, err := strconv.ParseUint(idStr, 10, 64) - if err != nil { - response.AbortBadRequest(c, errInvalidTokenID) - return - } - - newTokenStr, tokenRecord, err := rotateAccessTokenLogic(ctx, id, currUser.ID) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - c.JSON(http.StatusOK, response.OK(tokenResponse{ - Token: newTokenStr, - Record: *tokenRecord, - })) -} diff --git a/internal/apps/user/constants.go b/internal/apps/user/constants.go deleted file mode 100644 index b3851f5b..00000000 --- a/internal/apps/user/constants.go +++ /dev/null @@ -1,15 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package user - -import "time" - -const ( - verificationCodeRange = 900000 // 验证码随机范围 - verificationCodeOffset = 100000 // 验证码偏移量(保证 6 位) - emailCodeExpiry = 5 * time.Minute // 验证码有效期 - emailCodeCooldown = 60 * time.Second // 验证码发送冷却时间 - minPasswordLength = 8 // 密码最小长度 -) diff --git a/internal/apps/user/errs.go b/internal/apps/user/errs.go deleted file mode 100644 index 2884b5fb..00000000 --- a/internal/apps/user/errs.go +++ /dev/null @@ -1,48 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package user - -const ( - errBindParamsFailed = "参数绑定失败" - errInvalidParams = "无效的参数" - errPasswordLoginDisabled = "管理员关闭了密码登录" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errUsernameOrPasswordWrong = "用户名或密码错误" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errLoginEmailMissing = "该账号未绑定邮箱,请联系管理员绑定邮箱后再登录" - errNeedEmailCodePrefix = "need_email_code:" - errSMTPInvalidUseTempCodePrefix = "smtp_invalid:" - errSMTPInvalidUseTempCode = "smtp 配置无效,使用临时码登录" - errEmailCodeInvalidOrExpired = "验证码错误或已过期" - errSaveSessionFailed = "无法保存会话信息,请重试" - errRegistrationDisabled = "管理员关闭了注册" - errPasswordTooShort = "密码长度不能少于 8 位" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errEmailOrCodeRequired = "邮箱或验证码未填写" - errNewPasswordTooShort = "新密码长度不能少于 8 位" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errLoginRequired = "请先登录" - errUserNotFound = "未找到该用户" - errOldPasswordIncorrect = "原密码不正确" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errPasswordEncryptFailed = "密码加密失败,请重试" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errEmailRequired = "邮箱地址不能为空" - errUnsupportedEmailScene = "不支持的验证场景" - errEmailAlreadyRegistered = "该邮箱已被注册" - errEmailCodeCooldown = "验证码发送频繁,请稍后再试" - errEmailFormatInvalid = "邮箱格式不正确" - errEmailAlreadyBound = "该邮箱已被其他账号绑定" - errRenderEmailTemplateFailed = "渲染验证邮件模板失败:%w" - errGenerateEmailCodeFailed = "生成验证码失败,请重试" - errDispatchEmailTaskFailed = "投递验证邮件发送任务失败,请重试" - errTokenNameRequired = "令牌名称不能为空" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errAccessTokenLimitReached = "已达到访问令牌最大创建数量限制" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errGenerateTokenFailed = "生成令牌失败" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errInvalidTokenID = "无效的令牌ID" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errTokenNotFoundOrForbidden = "令牌不存在或无权操作" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errAdminTokenRequiresAdmin = "只有管理员才能创建具有管理员权限的令牌" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errTaskPayloadRequired = "任务参数不能为空" - errInvalidJSONFormat = "无效的 JSON 格式: %w" - errEmailTaskFieldsRequired = "to、subject、body 不能为空" - errParseEmailPayloadFailed = "解析邮件发送参数失败: %w" - errSMTPConfigIncomplete = "系统 SMTP 邮件服务配置不完整" - errSendMailFailed = "发送邮件失败: %w" - errLoginRateLimited = "请求过于频繁,请稍后再试" -) diff --git a/internal/apps/user/logics.go b/internal/apps/user/logics.go deleted file mode 100644 index 50c3b442..00000000 --- a/internal/apps/user/logics.go +++ /dev/null @@ -1,458 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package user - -import ( - "context" - "crypto/rand" - "crypto/sha256" - "crypto/subtle" - "encoding/json" - "errors" - "fmt" - "math/big" - "strings" - "time" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - pkgu "github.com/Rain-kl/Wavelet/pkg/util" -) - -// LoginEmailVerificationStatus 登录邮箱验证的处理结果。 -type LoginEmailVerificationStatus int - -const ( - // LoginEmailVerificationPassed 验证通过,可继续登录流程。 - LoginEmailVerificationPassed LoginEmailVerificationStatus = iota - // LoginEmailVerificationPending 需要用户输入邮箱验证码。 - LoginEmailVerificationPending - // LoginEmailVerificationRejected 验证被拒绝(验证码错误、临时码提示等)。 - LoginEmailVerificationRejected -) - -// LoginEmailVerificationResult 登录邮箱验证的业务结果。 -type LoginEmailVerificationResult struct { - Status LoginEmailVerificationStatus - Message string -} - -type updateProfileInput struct { - Nickname string - Email string - AvatarURL string - Bio string - Phone string - Gender string - Website string - Location string -} - -const ( - loginFailLimitKeyFormat = "login:fail:%s" - loginFailLimitMax = 20 - loginFailLimitWindow = 10 * time.Minute -) - -func loginFailLimitKey(ip string) string { - return fmt.Sprintf(loginFailLimitKeyFormat, strings.TrimSpace(ip)) -} - -func loginAttemptsBlocked(ctx context.Context, ip string) bool { - if db.Redis == nil { - return false - } - n, err := db.Redis.Get(ctx, db.PrefixedKey(loginFailLimitKey(ip))).Int() - if err != nil { - return false - } - return n >= loginFailLimitMax -} - -func recordFailedLogin(ctx context.Context, ip string) { - if db.Redis == nil { - return - } - key := db.PrefixedKey(loginFailLimitKey(ip)) - n, err := db.Redis.Incr(ctx, key).Result() - if err != nil { - return - } - if n == 1 { - _ = db.Redis.Expire(ctx, key, loginFailLimitWindow).Err() - } -} - -func clearFailedLogins(ctx context.Context, ip string) { - if db.Redis == nil { - return - } - _ = db.Redis.Del(ctx, db.PrefixedKey(loginFailLimitKey(ip))).Err() -} - -func isPasswordLoginEnabled(ctx context.Context) bool { - enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordLoginEnabled) - if err != nil { - return true - } - return enabled -} - -func isPasswordRegisterEnabled(ctx context.Context) bool { - enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordRegisterEnabled) - if err != nil { - return false - } - return enabled -} - -func isRegistrationEnabled(ctx context.Context) bool { - enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled) - if err != nil { - return false - } - return enabled -} - -func isEmailLoginVerificationEnabled(ctx context.Context) bool { - enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyEmailLoginVerificationEnabled) - if err != nil { - return false - } - return enabled -} - -func isEmailRegisterVerificationEnabled(ctx context.Context) bool { - enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyEmailRegisterVerificationEnabled) - if err != nil { - return false - } - return enabled -} - -func isSMTPConfigured(ctx context.Context) bool { - scHost, errHost := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPHost) - scPort, errPort := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPort) - scUser, errUser := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPUsername) - scPass, errPass := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPassword) - if errHost != nil || errPort != nil || errUser != nil || errPass != nil { - return false - } - return scHost.Value != "" && scPort.Value != "" && scUser.Value != "" && scPass.Value != "" -} - -func generateVerificationCode() (string, error) { - n, err := rand.Int(rand.Reader, big.NewInt(verificationCodeRange)) - if err != nil { - return "", err - } - return fmt.Sprintf("%06d", n.Int64()+verificationCodeOffset), nil -} - -func getEmailCodeKey(scene, email string) string { - return fmt.Sprintf("email_code:%s:%s", scene, email) -} - -func getEmailCooldownKey(scene, email string) string { - return fmt.Sprintf("email_code:cooldown:%s:%s", scene, email) -} - -func sendEmailVerificationCode(ctx context.Context, email, scene, templateName string) error { - if !isSMTPConfigured(ctx) { - return errors.New(errSMTPConfigIncomplete) - } - - code, err := generateVerificationCode() - if err != nil { - return errors.New(errGenerateEmailCodeFailed) - } - codeKey := getEmailCodeKey(scene, email) - cooldownKey := getEmailCooldownKey(scene, email) - - tmpl, err := repository.GetTemplateByKey(ctx, templateName) - if err != nil { - return fmt.Errorf("模板 %s 不存在或不可用: %w", templateName, err) - } - emailSubject, emailBody, err := tmpl.Render(map[string]any{"Code": code}) - if err != nil { - return fmt.Errorf(errRenderEmailTemplateFailed, err) - } - - if err := db.SetJSON(ctx, codeKey, code, emailCodeExpiry); err != nil { - return errors.New(errGenerateEmailCodeFailed) - } - _ = db.SetJSON(ctx, cooldownKey, "1", emailCodeCooldown) - - payload := SendEmailPayload{ - To: email, - Subject: emailSubject, - Body: emailBody, - } - payloadBytes, _ := json.Marshal(payload) - _, err = task.DispatchTask(ctx, TaskTypeSendEmail, payloadBytes, "system") - if err != nil { - return errors.New(errDispatchEmailTaskFailed) - } - return nil -} - -func verifyEmailCode(ctx context.Context, email, scene, code string) bool { - codeKey := getEmailCodeKey(scene, email) - var storedCode string - if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil { - return false - } - sumGot := sha256.Sum256([]byte(strings.TrimSpace(code))) - sumWant := sha256.Sum256([]byte(strings.TrimSpace(storedCode))) - if subtle.ConstantTimeCompare(sumGot[:], sumWant[:]) != 1 { - return false - } - _ = db.Redis.Del(ctx, db.PrefixedKey(codeKey)).Err() - return true -} - -func processLoginEmailVerification(ctx context.Context, code string, user *model.User) (LoginEmailVerificationResult, error) { - if code != "" { - if !verifyEmailCode(ctx, user.Email, "login", code) { - return LoginEmailVerificationResult{ - Status: LoginEmailVerificationRejected, - Message: errEmailCodeInvalidOrExpired, - }, nil - } - return LoginEmailVerificationResult{Status: LoginEmailVerificationPassed}, nil - } - - // 如果 SMTP 未配置,或者用户没有绑定邮箱(无法发送验证码),则使用临时码 888888 - if !isSMTPConfigured(ctx) || user.Email == "" { - codeKey := getEmailCodeKey("login", user.Email) - if err := db.SetJSON(ctx, codeKey, "888888", emailCodeExpiry); err != nil { - return LoginEmailVerificationResult{}, errors.New(errGenerateEmailCodeFailed) - } - var msg string - if !isSMTPConfigured(ctx) { - msg = errSMTPInvalidUseTempCodePrefix + errSMTPInvalidUseTempCode - } else { - msg = errSMTPInvalidUseTempCodePrefix + "该账号未绑定邮箱,使用临时码登录" - } - return LoginEmailVerificationResult{ - Status: LoginEmailVerificationRejected, - Message: msg, - }, nil - } - - cooldownKey := getEmailCooldownKey("login", user.Email) - var temp string - if err := db.GetJSON(ctx, cooldownKey, &temp); err != nil { - if err := sendEmailVerificationCode(ctx, user.Email, "login", "login_email"); err != nil { - return LoginEmailVerificationResult{}, err - } - } - - maskedEmail := pkgu.MaskEmail(user.Email) - return LoginEmailVerificationResult{ - Status: LoginEmailVerificationPending, - Message: errNeedEmailCodePrefix + maskedEmail, - }, nil -} - -func sendRegisterEmailCode(ctx context.Context, email string) error { - email = strings.TrimSpace(email) - if email == "" { - return errors.New(errEmailRequired) - } - - count, err := repository.CountUsersByEmail(ctx, email) - if err != nil { - return err - } - if count > 0 { - return errors.New(errEmailAlreadyRegistered) - } - - cooldownKey := getEmailCooldownKey("register", email) - var temp string - if err := db.GetJSON(ctx, cooldownKey, &temp); err == nil { - return errors.New(errEmailCodeCooldown) - } - - return sendEmailVerificationCode(ctx, email, "register", "register_email") -} - -func validateRegisterEmailVerification(ctx context.Context, email, code string) error { - if !isEmailRegisterVerificationEnabled(ctx) { - return nil - } - if email == "" || code == "" { - return errors.New(errEmailOrCodeRequired) - } - if !verifyEmailCode(ctx, email, "register", code) { - return errors.New(errEmailCodeInvalidOrExpired) - } - return nil -} - -func updateUserProfile(ctx context.Context, userID uint64, input updateProfileInput) (*model.User, error) { - dbUser, err := repository.GetUserByID(ctx, userID) - if err != nil { - return nil, errors.New(errUserNotFound) - } - - input.Email = strings.TrimSpace(input.Email) - if input.Email != "" && input.Email != dbUser.Email { - if !strings.Contains(input.Email, "@") || !strings.Contains(input.Email, ".") { - return nil, errors.New(errEmailFormatInvalid) - } - - count, err := repository.CountUsersByEmailExceptID(ctx, input.Email, dbUser.ID) - if err != nil { - return nil, err - } - if count > 0 { - return nil, errors.New(errEmailAlreadyBound) - } - } - - dbUser.Nickname = strings.TrimSpace(input.Nickname) - if dbUser.Nickname == "" { - dbUser.Nickname = dbUser.Username - } - dbUser.Email = input.Email - dbUser.AvatarURL = input.AvatarURL - dbUser.Bio = input.Bio - dbUser.Phone = strings.TrimSpace(input.Phone) - dbUser.Gender = strings.TrimSpace(input.Gender) - dbUser.Website = strings.TrimSpace(input.Website) - dbUser.Location = strings.TrimSpace(input.Location) - - if err := repository.UpdateUser(ctx, &dbUser); err != nil { - return nil, err - } - return &dbUser, nil -} - -func getUserByUsernameOrEmail(ctx context.Context, input string) (*model.User, error) { - user, err := repository.GetUserByUsernameOrEmail(ctx, input) - if err != nil { - return nil, err - } - return &user, nil -} - -func updateLastLogin(ctx context.Context, user *model.User) error { - return repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt) -} - -func registerUserLogic(ctx context.Context, u *model.User) error { - if err := repository.RegisterUserWithChecks(ctx, u); err != nil { - if strings.Contains(err.Error(), "duplicate key") || strings.Contains(err.Error(), "UNIQUE") { - return errors.New("用户名或邮箱已被占用") - } - return errors.New("注册失败,请稍后再试") - } - return nil -} - -func changePasswordLogic(ctx context.Context, userID uint64, oldPass, newPass string) error { - dbUser, err := repository.GetUserByID(ctx, userID) - if err != nil { - return errors.New(errUserNotFound) - } - - if !dbUser.CheckPassword(oldPass) { - return errors.New(errOldPasswordIncorrect) - } - - if err := dbUser.SetEncryptedPassword(newPass); err != nil { - return errors.New(errPasswordEncryptFailed) - } - - if err := repository.UpdateUserPassword(ctx, dbUser.ID, dbUser.Password); err != nil { - return errors.New("更新密码失败,请稍后再试") - } - - // 吊销该用户所有的 Access Token - if tokens, err := repository.ListAccessTokensByUserID(ctx, dbUser.ID); err == nil { - for _, token := range tokens { - oauth.InvalidateCachedToken(ctx, token.TokenHash) - } - } - if err := repository.DeleteAccessTokensByUserID(ctx, dbUser.ID); err != nil { - return errors.New("吊销 Access Token 失败,请稍后再试") - } - - oauth.InvalidateCachedUser(ctx, dbUser.ID) - return nil -} - -func listAccessTokensLogic(ctx context.Context, userID uint64) ([]model.AccessToken, error) { - tokens, err := repository.ListAccessTokensByUserID(ctx, userID) - if err != nil { - return nil, errors.New("获取令牌列表失败,请稍后再试") - } - return tokens, nil -} - -func countAccessTokensLogic(ctx context.Context, userID uint64) (int64, error) { - count, err := repository.CountAccessTokensByUserID(ctx, userID) - if err != nil { - return 0, errors.New("查询令牌数量失败,请稍后再试") - } - return count, nil -} - -func createAccessTokenLogic(ctx context.Context, record *model.AccessToken) error { - if err := repository.CreateAccessToken(ctx, record); err != nil { - return errors.New("创建令牌失败,请稍后再试") - } - oauth.SetCachedToken(ctx, record.TokenHash, record) - return nil -} - -func deleteAccessTokenLogic(ctx context.Context, id, userID uint64) error { - tokenRecord, err := repository.GetAccessTokenByIDAndUserID(ctx, id, userID) - if err != nil { - return errors.New(errTokenNotFoundOrForbidden) - } - oauth.InvalidateCachedToken(ctx, tokenRecord.TokenHash) - - rows, err := repository.DeleteAccessTokenForUser(ctx, id, userID) - if err != nil { - return errors.New("删除令牌失败,请稍后再试") - } - if rows == 0 { - return errors.New(errTokenNotFoundOrForbidden) - } - return nil -} - -func rotateAccessTokenLogic(ctx context.Context, id, userID uint64) (string, *model.AccessToken, error) { - tokenRecord, err := repository.GetAccessTokenByIDAndUserID(ctx, id, userID) - if err != nil { - return "", nil, errors.New(errTokenNotFoundOrForbidden) - } - - oauth.InvalidateCachedToken(ctx, tokenRecord.TokenHash) - - newTokenStr, err := model.GenerateTokenString() - if err != nil { - return "", nil, errors.New(errGenerateTokenFailed) - } - - newTokenHash := model.HashToken(newTokenStr) - newMaskedToken := model.MaskTokenString(newTokenStr) - - tokenRecord.TokenHash = newTokenHash - tokenRecord.MaskedToken = newMaskedToken - - if err := repository.SaveAccessToken(ctx, &tokenRecord); err != nil { - return "", nil, errors.New("轮换令牌失败,请稍后再试") - } - - oauth.SetCachedToken(ctx, tokenRecord.TokenHash, &tokenRecord) - - return newTokenStr, &tokenRecord, nil -} diff --git a/internal/apps/user/logics_test.go b/internal/apps/user/logics_test.go deleted file mode 100644 index f5ee503a..00000000 --- a/internal/apps/user/logics_test.go +++ /dev/null @@ -1,152 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package user - -import ( - "context" - "testing" - "time" - - "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/testhelper" -) - -func TestProcessLoginEmailVerificationSMTPFallback(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - const email = "smtpuser@example.com" - now := time.Now() - user := model.User{ - ID: 222, - Username: "smtpuser", - Nickname: "SMTP User", - Email: email, - IsActive: true, - LastLoginAt: now, - } - if err := user.SetEncryptedPassword("newpassword123"); err != nil { - t.Fatalf("set encrypted password failed: %v", err) - } - if err := dbConn.Create(&user).Error; err != nil { - t.Fatalf("create test user failed: %v", err) - } - - if err := dbConn.Model(&model.SystemConfig{}). - Where("key = ?", model.ConfigKeySMTPHost). - Update("value", "").Error; err != nil { - t.Fatalf("clear SMTP host failed: %v", err) - } - if err := db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Err(); err != nil { - t.Fatalf("invalidate system config cache failed: %v", err) - } - - ctx := context.Background() - result, err := processLoginEmailVerification(ctx, "", &user) - if err != nil { - t.Fatalf("processLoginEmailVerification() error = %v, want nil", err) - } - expected := errSMTPInvalidUseTempCodePrefix + errSMTPInvalidUseTempCode - if result.Status != LoginEmailVerificationRejected || result.Message != expected { - t.Fatalf("processLoginEmailVerification() = %+v, want rejected with %q", result, expected) - } - - codeKey := getEmailCodeKey("login", email) - var storedCode string - if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil { - t.Fatalf("get stored verification code failed: %v", err) - } - if storedCode != "888888" { - t.Errorf("stored verification code = %q, want %q", storedCode, "888888") - } - - passed, err := processLoginEmailVerification(ctx, "888888", &user) - if err != nil { - t.Fatalf("processLoginEmailVerification(valid code) error = %v, want nil", err) - } - if passed.Status != LoginEmailVerificationPassed { - t.Fatalf("processLoginEmailVerification(valid code) status = %v, want passed", passed.Status) - } -} - -func TestProcessLoginEmailVerificationEmptyEmailFallback(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - now := time.Now() - user := model.User{ - ID: 223, - Username: "emptyemailuser", - Nickname: "Empty Email User", - Email: "", - IsActive: true, - LastLoginAt: now, - } - if err := user.SetEncryptedPassword("newpassword123"); err != nil { - t.Fatalf("set encrypted password failed: %v", err) - } - if err := dbConn.Create(&user).Error; err != nil { - t.Fatalf("create test user failed: %v", err) - } - - for _, cfg := range []struct { - key string - value string - }{ - {model.ConfigKeySMTPHost, "smtp.example.com"}, - {model.ConfigKeySMTPPort, "587"}, - {model.ConfigKeySMTPUsername, "smtpuser"}, - {model.ConfigKeySMTPPassword, "smtppassword"}, - } { - if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", cfg.key).Update("value", cfg.value).Error; err != nil { - t.Fatalf("set %s failed: %v", cfg.key, err) - } - } - if err := db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Err(); err != nil { - t.Fatalf("invalidate system config cache failed: %v", err) - } - - ctx := context.Background() - result, err := processLoginEmailVerification(ctx, "", &user) - if err != nil { - t.Fatalf("processLoginEmailVerification() error = %v, want nil", err) - } - expected := errSMTPInvalidUseTempCodePrefix + "该账号未绑定邮箱,使用临时码登录" - if result.Status != LoginEmailVerificationRejected || result.Message != expected { - t.Fatalf("processLoginEmailVerification() = %+v, want rejected with %q", result, expected) - } -} - -func TestProcessLoginEmailVerificationInvalidCode(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - const email = "codeduser@example.com" - now := time.Now() - user := model.User{ - ID: 224, - Username: "codeduser", - Email: email, - IsActive: true, - LastLoginAt: now, - } - if err := dbConn.Create(&user).Error; err != nil { - t.Fatalf("create test user failed: %v", err) - } - - ctx := context.Background() - if err := db.SetJSON(ctx, getEmailCodeKey("login", email), "123456", emailCodeExpiry); err != nil { - t.Fatalf("seed verification code failed: %v", err) - } - - result, err := processLoginEmailVerification(ctx, "000000", &user) - if err != nil { - t.Fatalf("processLoginEmailVerification() error = %v, want nil", err) - } - if result.Status != LoginEmailVerificationRejected || result.Message != errEmailCodeInvalidOrExpired { - t.Fatalf("processLoginEmailVerification() = %+v, want rejected invalid code", result) - } -} diff --git a/internal/apps/user/routers.go b/internal/apps/user/routers.go deleted file mode 100644 index 445176ba..00000000 --- a/internal/apps/user/routers.go +++ /dev/null @@ -1,392 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package user - -import ( - "net/http" - "strings" - "time" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen" - "github.com/Rain-kl/Wavelet/internal/listener" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/shared/response" - "github.com/Rain-kl/Wavelet/pkg/logger" - pkgu "github.com/Rain-kl/Wavelet/pkg/util" - "github.com/gin-contrib/sessions" - "github.com/gin-gonic/gin" -) - -type loginRequest struct { - Username string `json:"username"` - Password string `json:"password"` - Code string `json:"code"` -} - -type registerRequest struct { - Username string `json:"username"` - Password string `json:"password"` - Nickname string `json:"nickname"` - DisplayName string `json:"display_name"` - Email string `json:"email"` - Code string `json:"code"` -} - -type sendEmailCodeRequest struct { - Email string `json:"email" binding:"required,email"` - Scene string `json:"scene" binding:"required"` -} - -type updateProfileRequest struct { - Nickname string `json:"nickname"` - Email string `json:"email"` - AvatarURL string `json:"avatar_url"` - Bio string `json:"bio"` - Phone string `json:"phone"` - Gender string `json:"gender"` - Website string `json:"website"` - Location string `json:"location"` -} - -// Login 用户密码登录 -// @Summary 用户密码登录 -// @Description 使用用户名和密码登录,登录成功后建立 Session。若管理员已关闭密码登录功能则返回错误。 -// @Tags user -// @Accept json -// @Produce json -// @Param request body user.loginRequest true "登录请求参数" -// @Success 200 {object} response.Any{data=oauth.BasicUserInfo} "登录成功,返回用户信息" -// @Failure 400 {object} response.Any "用户名或密码错误" -// @Failure 500 {object} response.Any "服务内部错误" -// @Router /api/v1/user/login [post] -func Login(c *gin.Context) { - ctx := c.Request.Context() - if !isPasswordLoginEnabled(ctx) { - response.AbortBadRequest(c, errPasswordLoginDisabled) - return - } - if loginAttemptsBlocked(ctx, c.ClientIP()) { - response.AbortBadRequest(c, errLoginRateLimited) - return - } - var req loginRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - req.Username = strings.TrimSpace(req.Username) - if req.Username == "" || req.Password == "" { - response.AbortBadRequest(c, errInvalidParams) - return - } - - user, err := getUserByUsernameOrEmail(ctx, req.Username) - if err != nil { - pkgu.DummyCheckPassword(req.Password) - recordFailedLogin(ctx, c.ClientIP()) - logger.WarnF(ctx, "[LoginAudit] failed login attempt (username not found) for input: %s, IP: %s", req.Username, c.ClientIP()) - response.AbortBadRequest(c, errUsernameOrPasswordWrong) - return - } - if !user.IsActive { - recordFailedLogin(ctx, c.ClientIP()) - logger.WarnF(ctx, "[LoginAudit] banned user login attempt for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP()) - response.AbortBadRequest(c, errUsernameOrPasswordWrong) - return - } - - // 判定是否是明文密码存储 - isPlaintext := !user.IsPasswordEncrypted() - - if !user.CheckPassword(req.Password) { - recordFailedLogin(ctx, c.ClientIP()) - logger.WarnF(ctx, "[LoginAudit] failed login attempt (incorrect password) for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP()) - response.AbortBadRequest(c, errUsernameOrPasswordWrong) - return - } - - if isEmailLoginVerificationEnabled(ctx) { - result, err := processLoginEmailVerification(ctx, req.Code, user) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - if result.Status != LoginEmailVerificationPassed { - response.AbortBadRequest(c, result.Message) - return - } - } - - needChangePassword := isPlaintext - - user.LastLoginAt = time.Now() - if err := updateLastLogin(ctx, user); err != nil { - response.AbortBadRequest(c, "更新登录时间失败,请稍后再试") - return - } - extras := map[string]any{} - if isPlaintext { - extras["need_change_password"] = true - } - clearFailedLogins(ctx, c.ClientIP()) - if err := oauth.SetLoginSession(ctx, c, user, extras); err != nil { - response.AbortBadRequest(c, errSaveSessionFailed) - return - } - - oauth.SetCachedUser(ctx, user.ID, user) - - logger.InfoF(ctx, "[LoginAudit] successful login for user: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP()) - - listener.EmitAdminLoggedIn(ctx, user, c.ClientIP()) - - c.JSON(http.StatusOK, response.OK(oauth.BuildBasicUserInfo(user, needChangePassword))) -} - -// Register 用户注册 -// @Summary 用户注册 -// @Description 使用用户名和密码注册新账号,注册成功后自动登录并建立 Session。密码长度不能少于 8 位。 -// @Tags user -// @Accept json -// @Produce json -// @Param request body user.registerRequest true "注册请求参数" -// @Success 200 {object} response.Any{data=oauth.BasicUserInfo} "注册并登录成功,返回用户信息" -// @Failure 400 {object} response.Any "参数错误、用户名已存在或注册已关闭" -// @Failure 500 {object} response.Any "服务内部错误" -// @Router /api/v1/user/register [post] -func Register(c *gin.Context) { - ctx := c.Request.Context() - if !isRegistrationEnabled(ctx) || !isPasswordRegisterEnabled(ctx) { - response.AbortBadRequest(c, errRegistrationDisabled) - return - } - - var req registerRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - req.Username = strings.TrimSpace(req.Username) - req.Password = strings.TrimSpace(req.Password) - req.Nickname = strings.TrimSpace(req.Nickname) - req.DisplayName = strings.TrimSpace(req.DisplayName) - req.Email = strings.TrimSpace(req.Email) - req.Code = strings.TrimSpace(req.Code) - - if req.Username == "" || req.Password == "" { - response.AbortBadRequest(c, errInvalidParams) - return - } - if req.Email == "" { - response.AbortBadRequest(c, errEmailRequired) - return - } - if len(req.Password) < minPasswordLength { - response.AbortBadRequest(c, errPasswordTooShort) - return - } - - // 邮箱注册验证校验 - if err := validateRegisterEmailVerification(ctx, req.Email, req.Code); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - user := model.User{ - ID: idgen.NextUint64ID(), - Username: req.Username, - Nickname: req.Nickname, - Email: req.Email, - AvatarURL: "", - IsActive: true, - IsAdmin: false, - LastLoginAt: time.Now(), - } - if user.Nickname == "" { - user.Nickname = req.DisplayName - } - if user.Nickname == "" { - user.Nickname = req.Username - } - if err := user.SetEncryptedPassword(req.Password); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - if err := registerUserLogic(ctx, &user); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - if err := oauth.SetLoginSession(ctx, c, &user); err != nil { - response.AbortBadRequest(c, errSaveSessionFailed) - return - } - - c.JSON(http.StatusOK, response.OK(oauth.BuildBasicUserInfo(&user, false))) -} - -// Logout 用户退出登录 -// @Summary 用户退出登录 -// @Description 清除用户登录 Session,完成退出 -// @Tags user -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=string} "退出成功" -// @Failure 500 {object} response.Any "Session 清除失败" -// @Router /api/v1/user/logout [get] -func Logout(c *gin.Context) { - session := sessions.Default(c) - userID := session.Get(oauth.UserIDKey) - username := session.Get(oauth.UserNameKey) - if userID != nil { - logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP()) - if id, ok := userID.(uint64); ok { - oauth.InvalidateCachedUser(c.Request.Context(), id) - } else if idFloat, ok := userID.(float64); ok { - oauth.InvalidateCachedUser(c.Request.Context(), uint64(idFloat)) - } else if idInt, ok := userID.(int); ok && idInt >= 0 { - oauth.InvalidateCachedUser(c.Request.Context(), uint64(idInt)) - } - } - session.Options(oauth.GetSessionOptions(-1)) - session.Clear() - if err := session.Save(); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK("")) -} - -type changePasswordRequest struct { - OldPassword string `json:"old_password"` - NewPassword string `json:"new_password"` -} - -// ChangePassword 修改用户密码 -// @Summary 修改用户密码 -// @Description 修改当前登录用户的密码。修改成功后,如果是首次明文登录的升级提示,则清除修改密码的提示状态。 -// @Tags user -// @Accept json -// @Produce json -// @Param request body user.changePasswordRequest true "修改密码请求参数" -// @Success 200 {object} response.Any{data=string} "修改密码成功" -// @Failure 400 {object} response.Any "原密码错误或新密码不符合要求" -// @Failure 401 {object} response.Any "请先登录" -// @Router /api/v1/user/change-password [post] -func ChangePassword(c *gin.Context) { - var req changePasswordRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - req.OldPassword = strings.TrimSpace(req.OldPassword) - req.NewPassword = strings.TrimSpace(req.NewPassword) - - if req.OldPassword == "" || req.NewPassword == "" { - response.AbortBadRequest(c, errInvalidParams) - return - } - if len(req.NewPassword) < minPasswordLength { - response.AbortBadRequest(c, errNewPasswordTooShort) - return - } - - userObj, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) - if userObj == nil { - response.AbortUnauthorized(c, errLoginRequired) - return - } - - ctx := c.Request.Context() - if err := changePasswordLogic(ctx, userObj.ID, req.OldPassword, req.NewPassword); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - // 销毁当前活跃会话以强制重新登录 - session := sessions.Default(c) - session.Clear() - _ = session.Save() - - c.JSON(http.StatusOK, response.OK("密码修改成功")) -} - -// SendEmailCode 发送邮箱验证码 -// @Summary 发送邮箱验证码 -// @Description 向指定邮箱发送验证码(用于注册场景) -// @Tags user -// @Accept json -// @Produce json -// @Param request body user.sendEmailCodeRequest true "发送验证码请求参数" -// @Success 200 {object} response.Any "发送成功" -// @Failure 400 {object} response.Any "参数错误" -// @Router /api/v1/user/send-email-code [post] -func SendEmailCode(c *gin.Context) { - var req sendEmailCodeRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - req.Email = strings.TrimSpace(req.Email) - if req.Email == "" { - response.AbortBadRequest(c, errEmailRequired) - return - } - - if req.Scene != "register" { - response.AbortBadRequest(c, errUnsupportedEmailScene) - return - } - - ctx := c.Request.Context() - if err := sendRegisterEmailCode(ctx, req.Email); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - c.JSON(http.StatusOK, response.OKNil()) -} - -// UpdateProfile 修改当前登录用户的个人资料 -// @Summary 修改当前登录用户的个人资料 -// @Description 修改当前登录用户的昵称、邮箱、头像、简介、电话、性别、个人网站和所在地。 -// @Tags user -// @Accept json -// @Produce json -// @Param request body user.updateProfileRequest true "更新请求参数" -// @Success 200 {object} response.Any{data=oauth.BasicUserInfo} "修改成功,返回更新后的用户信息" -// @Failure 400 {object} response.Any "邮箱已被占用或参数错误" -// @Failure 401 {object} response.Any "未登录" -// @Router /api/v1/user/profile [put] -func UpdateProfile(c *gin.Context) { - var req updateProfileRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - userObj, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) - if userObj == nil { - response.AbortUnauthorized(c, errLoginRequired) - return - } - - ctx := c.Request.Context() - dbUser, err := updateUserProfile(ctx, userObj.ID, updateProfileInput(req)) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - oauth.InvalidateCachedUser(ctx, userObj.ID) - - session := sessions.Default(c) - needChange := session.Get("need_change_password") == true - - c.JSON(http.StatusOK, response.OK(oauth.BuildBasicUserInfo(dbUser, needChange))) -} diff --git a/internal/apps/user/routers_test.go b/internal/apps/user/routers_test.go deleted file mode 100644 index 873ac5e4..00000000 --- a/internal/apps/user/routers_test.go +++ /dev/null @@ -1,684 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package user - -import ( - "bytes" - "context" - "encoding/json" - "net/http" - "net/http/httptest" - "strings" - "testing" - "time" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/infra/config" - "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/testhelper" - "github.com/gin-contrib/sessions" - "github.com/gin-contrib/sessions/cookie" - "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/shared/response" -) - -func setupUserTestRouter(t *testing.T) *gin.Engine { - t.Helper() - - oldCookieName := config.Config.App.SessionCookieName - oldSecret := config.Config.App.SessionSecret - oldDomain := config.Config.App.SessionDomain - oldSecure := config.Config.App.SessionSecure - oldHTTPOnly := config.Config.App.SessionHTTPOnly - t.Cleanup(func() { - config.Config.App.SessionCookieName = oldCookieName - config.Config.App.SessionSecret = oldSecret - config.Config.App.SessionDomain = oldDomain - config.Config.App.SessionSecure = oldSecure - config.Config.App.SessionHTTPOnly = oldHTTPOnly - }) - - config.Config.App.SessionCookieName = "test_session_id" - config.Config.App.SessionSecret = "test_session_secret" - config.Config.App.SessionDomain = "" - config.Config.App.SessionSecure = false - config.Config.App.SessionHTTPOnly = true - - store := cookie.NewStore([]byte(config.Config.App.SessionSecret)) - store.Options(oauth.GetSessionOptions(3600)) - r := testhelper.NewTestGinEngine(sessions.Sessions(config.Config.App.SessionCookieName, store)) - - api := r.Group("/api/v1") - api.POST("/user/register", Register) - api.POST("/user/login", Login) - api.GET("/user-info", oauth.LoginRequired(), oauth.UserInfo) - return r -} - -func performUserRequest(r http.Handler, method, path string, body []byte, cookies []*http.Cookie) *httptest.ResponseRecorder { - var reader *bytes.Reader - if body != nil { - reader = bytes.NewReader(body) - } else { - reader = bytes.NewReader(nil) - } - - req, _ := http.NewRequest(method, path, reader) - req.Header.Set("Content-Type", "application/json") - for _, c := range cookies { - req.AddCookie(c) - } - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - return w -} - -func sessionCookieFromResponse(t *testing.T, w *httptest.ResponseRecorder) *http.Cookie { - t.Helper() - - for _, c := range w.Result().Cookies() { - if c.Name == config.Config.App.SessionCookieName { - return c - } - } - t.Fatalf("sessionCookieFromResponse() did not find %q cookie", config.Config.App.SessionCookieName) - return nil -} - -func basicUserInfoFromResponse(t *testing.T, w *httptest.ResponseRecorder) oauth.BasicUserInfo { - t.Helper() - - var resp response.Any - if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil { - t.Fatalf("basicUserInfoFromResponse() decode response failed: %v", err) - } - if resp.ErrorMsg != "" { - t.Fatalf("basicUserInfoFromResponse() error_msg = %q, want empty", resp.ErrorMsg) - } - data, _ := json.Marshal(resp.Data) - var info oauth.BasicUserInfo - if err := json.Unmarshal(data, &info); err != nil { - t.Fatalf("basicUserInfoFromResponse() decode data failed: %v", err) - } - return info -} - -func TestEmailCooldownKeyIncludesScene(t *testing.T) { - email := "user@example.com" - - loginKey := getEmailCooldownKey("login", email) - registerKey := getEmailCooldownKey("register", email) - if loginKey == registerKey { - t.Errorf("getEmailCooldownKey(%q, %q) = %q, want different key from register scene", "login", email, loginKey) - } - if want := "email_code:cooldown:login:user@example.com"; loginKey != want { - t.Errorf("getEmailCooldownKey(%q, %q) = %q, want %q", "login", email, loginKey, want) - } -} - -func TestGenerateVerificationCode(t *testing.T) { - code, err := generateVerificationCode() - if err != nil { - t.Fatalf("generateVerificationCode() error = %v, want nil", err) - } - if len(code) != 6 { - t.Fatalf("generateVerificationCode() length = %d, want 6. Code: %q", len(code), code) - } - for _, r := range code { - if r < '0' || r > '9' { - t.Fatalf("generateVerificationCode() = %q, want only digits", code) - } - } -} - -func TestRegisterCreatesAuthenticatedEncryptedUser(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - router := setupUserTestRouter(t) - payload := registerRequest{ - Username: "newuser", - Password: "newpassword123", - Nickname: "New User", - Email: "newuser@example.com", - } - body, _ := json.Marshal(payload) - - w := performUserRequest(router, http.MethodPost, "/api/v1/user/register", body, nil) - if w.Code != http.StatusOK { - t.Fatalf("Register(%q) status = %d, want %d. Body: %s", payload.Username, w.Code, http.StatusOK, w.Body.String()) - } - info := basicUserInfoFromResponse(t, w) - if info.NeedChangePassword { - t.Errorf("Register(%q) need_change_password = true, want false", payload.Username) - } - - var dbUser model.User - if err := dbConn.Where("username = ?", payload.Username).First(&dbUser).Error; err != nil { - t.Fatalf("Register(%q) db lookup failed: %v", payload.Username, err) - } - if dbUser.ID < 1000 { - t.Errorf("Register(%q) user ID = %d, want generated snowflake ID", payload.Username, dbUser.ID) - } - if !dbUser.IsPasswordEncrypted() { - t.Errorf("Register(%q) stored plaintext password, want encrypted password", payload.Username) - } - if !dbUser.CheckPassword(payload.Password) { - t.Errorf("Register(%q) stored password does not match original password", payload.Username) - } - - sessionCookie := sessionCookieFromResponse(t, w) - w = performUserRequest(router, http.MethodGet, "/api/v1/user-info", nil, []*http.Cookie{sessionCookie}) - if w.Code != http.StatusOK { - t.Fatalf("UserInfo() after Register(%q) status = %d, want %d. Body: %s", payload.Username, w.Code, http.StatusOK, w.Body.String()) - } -} - -func TestLoginRequiresPasswordChangeForInitialPlaintextAdmin(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - const ( - adminID = uint64(1) - adminUsername = "admin" - adminPassword = "12345678" - ) - now := time.Now() - if err := dbConn.Exec( - `INSERT INTO w_users (id, username, password, nickname, is_active, is_admin, last_login_at, created_at, updated_at) -VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`, - adminID, - adminUsername, - adminPassword, - "Administrator", - true, - true, - now, - now, - now, - ).Error; err != nil { - t.Fatalf("seed initial admin failed: %v", err) - } - - router := setupUserTestRouter(t) - payload := loginRequest{ - Username: adminUsername, - Password: adminPassword, - } - body, _ := json.Marshal(payload) - - w := performUserRequest(router, http.MethodPost, "/api/v1/user/login", body, nil) - if w.Code != http.StatusOK { - t.Fatalf("Login(%q) status = %d, want %d. Body: %s", adminUsername, w.Code, http.StatusOK, w.Body.String()) - } - info := basicUserInfoFromResponse(t, w) - if !info.NeedChangePassword { - t.Errorf("Login(%q) need_change_password = false, want true", adminUsername) - } - - var dbUser model.User - if err := dbConn.Where("username = ?", adminUsername).First(&dbUser).Error; err != nil { - t.Fatalf("Login(%q) db lookup failed: %v", adminUsername, err) - } - if dbUser.ID != adminID { - t.Errorf("Login(%q) user ID = %d, want %d", adminUsername, dbUser.ID, adminID) - } - if dbUser.IsPasswordEncrypted() { - t.Errorf("Login(%q) encrypted password during login, want plaintext until password change", adminUsername) - } - if !dbUser.CheckPassword(adminPassword) { - t.Errorf("Login(%q) stored password does not match original password", adminUsername) - } - - sessionCookie := sessionCookieFromResponse(t, w) - w = performUserRequest(router, http.MethodGet, "/api/v1/user-info", nil, []*http.Cookie{sessionCookie}) - if w.Code != http.StatusOK { - t.Fatalf("UserInfo() after Login(%q) status = %d, want %d. Body: %s", adminUsername, w.Code, http.StatusOK, w.Body.String()) - } - info = basicUserInfoFromResponse(t, w) - if !info.NeedChangePassword { - t.Errorf("UserInfo() after Login(%q) need_change_password = false, want true", adminUsername) - } -} - -func TestLoginEmailVerificationFallbackWhenSMTPUnconfigured(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - const ( - userID = uint64(222) - username = "smtpuser" - password = "newpassword123" - email = "smtpuser@example.com" - ) - now := time.Now() - user := model.User{ - ID: userID, - Username: username, - Nickname: "SMTP User", - Email: email, - IsActive: true, - IsAdmin: false, - LastLoginAt: now, - } - if err := user.SetEncryptedPassword(password); err != nil { - t.Fatalf("set encrypted password failed: %v", err) - } - if err := dbConn.Create(&user).Error; err != nil { - t.Fatalf("create test user failed: %v", err) - } - - // 1. Enable email login verification - if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyEmailLoginVerificationEnabled).Update("value", "true").Error; err != nil { - t.Fatalf("enable email login verification failed: %v", err) - } - // 2. Clear SMTP host to simulate unconfigured SMTP - if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPHost).Update("value", "").Error; err != nil { - t.Fatalf("clear SMTP host failed: %v", err) - } - - // 2.5 Invalidate the system config cache in Redis - if err := db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Err(); err != nil { - t.Fatalf("invalidate system config cache failed: %v", err) - } - - router := setupUserTestRouter(t) - - // 3. Perform login request without verification code - payload := loginRequest{ - Username: username, - Password: password, - } - body, _ := json.Marshal(payload) - - w := performUserRequest(router, http.MethodPost, "/api/v1/user/login", body, nil) - if w.Code != http.StatusBadRequest { - t.Fatalf("Login(%q) status = %d, want %d. Body: %s", username, w.Code, http.StatusBadRequest, w.Body.String()) - } - - // Check response error msg - var resp struct { - ErrorMsg string `json:"error_msg"` - } - if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil { - t.Fatalf("unmarshal response failed: %v", err) - } - expectedError := errSMTPInvalidUseTempCodePrefix + errSMTPInvalidUseTempCode - if resp.ErrorMsg != expectedError { - t.Errorf("expected error %q, got %q", expectedError, resp.ErrorMsg) - } - - // 4. Check that verification code stored in Redis is "888888" - ctx := context.Background() - codeKey := getEmailCodeKey("login", email) - var storedCode string - if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil { - t.Fatalf("get stored verification code failed: %v", err) - } - if storedCode != "888888" { - t.Errorf("expected verification code '888888', got %q", storedCode) - } - - // 5. Retry login with code "888888" - payload.Code = "888888" - bodyWithCode, _ := json.Marshal(payload) - w = performUserRequest(router, http.MethodPost, "/api/v1/user/login", bodyWithCode, nil) - if w.Code != http.StatusOK { - t.Fatalf("Login with code status = %d, want %d. Body: %s", w.Code, http.StatusOK, w.Body.String()) - } - - var successResp struct { - ErrorMsg string `json:"error_msg"` - } - if err := json.Unmarshal(w.Body.Bytes(), &successResp); err != nil { - t.Fatalf("unmarshal success response failed: %v", err) - } - if successResp.ErrorMsg != "" { - t.Errorf("expected login success, got error %q", successResp.ErrorMsg) - } -} - -func TestLoginEmailVerificationFallbackForEmptyEmail(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - const ( - userID = uint64(223) - username = "emptyemailuser" - password = "newpassword123" - email = "" - ) - now := time.Now() - user := model.User{ - ID: userID, - Username: username, - Nickname: "Empty Email User", - Email: email, - IsActive: true, - IsAdmin: true, - LastLoginAt: now, - } - if err := user.SetEncryptedPassword(password); err != nil { - t.Fatalf("set encrypted password failed: %v", err) - } - if err := dbConn.Create(&user).Error; err != nil { - t.Fatalf("create test user failed: %v", err) - } - - // 1. Enable email login verification - if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyEmailLoginVerificationEnabled).Update("value", "true").Error; err != nil { - t.Fatalf("enable email login verification failed: %v", err) - } - // 2. Make sure SMTP is configured so we only trigger empty email check - if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPHost).Update("value", "smtp.example.com").Error; err != nil { - t.Fatalf("set SMTP host failed: %v", err) - } - if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPPort).Update("value", "587").Error; err != nil { - t.Fatalf("set SMTP port failed: %v", err) - } - if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPUsername).Update("value", "smtpuser").Error; err != nil { - t.Fatalf("set SMTP username failed: %v", err) - } - if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySMTPPassword).Update("value", "smtppassword").Error; err != nil { - t.Fatalf("set SMTP password failed: %v", err) - } - - // Invalidate the system config cache in Redis - if err := db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Err(); err != nil { - t.Fatalf("invalidate system config cache failed: %v", err) - } - - router := setupUserTestRouter(t) - - // 3. Perform login request without verification code - payload := loginRequest{ - Username: username, - Password: password, - } - body, _ := json.Marshal(payload) - - w := performUserRequest(router, http.MethodPost, "/api/v1/user/login", body, nil) - if w.Code != http.StatusBadRequest { - t.Fatalf("Login(%q) status = %d, want %d. Body: %s", username, w.Code, http.StatusBadRequest, w.Body.String()) - } - - // Check response error msg - var resp struct { - ErrorMsg string `json:"error_msg"` - } - if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil { - t.Fatalf("unmarshal response failed: %v", err) - } - expectedError := errSMTPInvalidUseTempCodePrefix + "该账号未绑定邮箱,使用临时码登录" - if resp.ErrorMsg != expectedError { - t.Errorf("expected error %q, got %q", expectedError, resp.ErrorMsg) - } - - // 4. Check that verification code stored in Redis is "888888" - ctx := context.Background() - codeKey := getEmailCodeKey("login", email) - var storedCode string - if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil { - t.Fatalf("get stored verification code failed: %v", err) - } - if storedCode != "888888" { - t.Errorf("expected verification code '888888', got %q", storedCode) - } - - // 5. Retry login with code "888888" - payload.Code = "888888" - bodyWithCode, _ := json.Marshal(payload) - w = performUserRequest(router, http.MethodPost, "/api/v1/user/login", bodyWithCode, nil) - if w.Code != http.StatusOK { - t.Fatalf("Login with code status = %d, want %d. Body: %s", w.Code, http.StatusOK, w.Body.String()) - } - - var successResp struct { - ErrorMsg string `json:"error_msg"` - } - if err := json.Unmarshal(w.Body.Bytes(), &successResp); err != nil { - t.Fatalf("unmarshal success response failed: %v", err) - } - if successResp.ErrorMsg != "" { - t.Errorf("expected login success, got error %q", successResp.ErrorMsg) - } -} - -func TestAccessTokenEndpointsDisallowTokenAuth(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - // 1. Seed a user - const ( - userID = uint64(500) - username = "tokenuser" - password = "tokenpassword123" - ) - now := time.Now() - userRecord := model.User{ - ID: userID, - Username: username, - Nickname: "Token User", - Email: "tokenuser@example.com", - IsActive: true, - IsAdmin: true, // Make them an admin so we can test with is_admin=true token requests if needed - LastLoginAt: now, - } - if err := userRecord.SetEncryptedPassword(password); err != nil { - t.Fatalf("set encrypted password failed: %v", err) - } - if err := dbConn.Create(&userRecord).Error; err != nil { - t.Fatalf("create test user failed: %v", err) - } - - // Seed an active AccessToken for this user - tokenStr, err := model.GenerateTokenString() - if err != nil { - t.Fatalf("generate token string failed: %v", err) - } - tokenHash := model.HashToken(tokenStr) - tokenRecord := model.AccessToken{ - UserID: userID, - Name: "Test Token", - TokenHash: tokenHash, - MaskedToken: model.MaskTokenString(tokenStr), - IsAdmin: false, - } - if err := dbConn.Create(&tokenRecord).Error; err != nil { - t.Fatalf("create test access token failed: %v", err) - } - - // 2. Set up router with access-token routes and oauth middlewares - store := cookie.NewStore([]byte("test_session_secret")) - r := testhelper.NewTestGinEngine(sessions.Sessions("test_session_id", store)) - - apiV1Router := r.Group("/api/v1") - userRouter := apiV1Router.Group("/user") - tokenRouter := userRouter.Group("/access-tokens") - tokenRouter.Use(oauth.LoginRequired(), oauth.DisallowTokenAuth()) - { - tokenRouter.GET("", ListAccessTokens) - tokenRouter.POST("", CreateAccessToken) - tokenRouter.DELETE("/:id", DeleteAccessToken) - tokenRouter.POST("/:id/rotate", RotateAccessToken) - } - - // 3. Test that accessing using an Access Token fails with 403 Forbidden - req, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil) - req.Header.Set("X-Access-Token", tokenStr) - w := httptest.NewRecorder() - r.ServeHTTP(w, req) - - if w.Code != http.StatusForbidden { - t.Errorf("expected status 403 Forbidden when accessing with Access Token, got %d. Body: %s", w.Code, w.Body.String()) - } - - var resp struct { - ErrorMsg string `json:"error_msg"` - } - if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil { - t.Fatalf("decode response failed: %v", err) - } - if resp.ErrorMsg != oauth.ErrTokenAuthNotAllowed { - t.Errorf("expected error message %q, got %q", oauth.ErrTokenAuthNotAllowed, resp.ErrorMsg) - } - - // 4. Test that accessing using a Session succeeds - sessionCookieStore := cookie.NewStore([]byte("test_session_secret")) - rSession := testhelper.NewTestGinEngine(sessions.Sessions("test_session_id", sessionCookieStore)) - rSession.GET("/api/v1/user/access-tokens", oauth.LoginRequired(), oauth.DisallowTokenAuth(), ListAccessTokens) - - // We can login/register or just mock the session handler to set user ID - rSession.GET("/mock-login", func(c *gin.Context) { - session := sessions.Default(c) - session.Set(oauth.UserIDKey, userID) - session.Set(oauth.UserNameKey, username) - session.Set(oauth.PasswordHashKey, userRecord.Password) - _ = session.Save() - c.String(http.StatusOK, "ok") - }) - - wMock := httptest.NewRecorder() - reqMock, _ := http.NewRequest(http.MethodGet, "/mock-login", nil) - rSession.ServeHTTP(wMock, reqMock) - cookieVal := wMock.Header().Get("Set-Cookie") - - reqSession, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil) - reqSession.Header.Set("Cookie", cookieVal) - wSession := httptest.NewRecorder() - rSession.ServeHTTP(wSession, reqSession) - - if wSession.Code != http.StatusOK { - t.Errorf("expected status 200 OK when accessing with Session, got %d. Body: %s", wSession.Code, wSession.Body.String()) - } -} - -func TestChangePasswordRevocation(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - // 1. Seed a user with a password - user := model.User{ - ID: uint64(888), - Username: "revoketest", - Nickname: "Revoke Test", - IsActive: true, - } - if err := user.SetEncryptedPassword("oldpassword123"); err != nil { - t.Fatalf("set encrypted password failed: %v", err) - } - if err := dbConn.Create(&user).Error; err != nil { - t.Fatalf("create test user failed: %v", err) - } - - // 2. Seed an active AccessToken for this user - tokenStr, err := model.GenerateTokenString() - if err != nil { - t.Fatalf("generate token string failed: %v", err) - } - tokenHash := model.HashToken(tokenStr) - tokenRecord := model.AccessToken{ - UserID: user.ID, - Name: "Test Token", - TokenHash: tokenHash, - MaskedToken: model.MaskTokenString(tokenStr), - } - if err := dbConn.Create(&tokenRecord).Error; err != nil { - t.Fatalf("create test access token failed: %v", err) - } - - // 3. Set up router - store := cookie.NewStore([]byte("test_session_secret")) - r := testhelper.NewTestGinEngine(sessions.Sessions("test_session_id", store)) - - r.GET("/mock-login", func(c *gin.Context) { - session := sessions.Default(c) - session.Set(oauth.UserIDKey, user.ID) - session.Set(oauth.UserNameKey, user.Username) - session.Set(oauth.PasswordHashKey, user.Password) - _ = session.Save() - c.String(http.StatusOK, "ok") - }) - - r.GET("/mock-old-session-login", func(c *gin.Context) { - session := sessions.Default(c) - session.Set(oauth.UserIDKey, user.ID) - session.Set(oauth.UserNameKey, user.Username) - session.Set(oauth.PasswordHashKey, "invalid_old_password_hash") - _ = session.Save() - c.String(http.StatusOK, "ok") - }) - - r.POST("/api/v1/user/change-password", oauth.LoginRequired(), ChangePassword) - r.GET("/api/v1/user/access-tokens", oauth.LoginRequired(), ListAccessTokens) - - // 4. Perform mock login to get cookie - wMock := httptest.NewRecorder() - reqMock, _ := http.NewRequest(http.MethodGet, "/mock-login", nil) - r.ServeHTTP(wMock, reqMock) - cookieVal := wMock.Header().Get("Set-Cookie") - - // 5. Test that session and token work initially - reqSession, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil) - reqSession.Header.Set("Cookie", cookieVal) - wSession := httptest.NewRecorder() - r.ServeHTTP(wSession, reqSession) - if wSession.Code != http.StatusOK { - t.Errorf("expected 200, got %d", wSession.Code) - } - - reqToken, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil) - reqToken.Header.Set("X-Access-Token", tokenStr) - wToken := httptest.NewRecorder() - r.ServeHTTP(wToken, reqToken) - if wToken.Code != http.StatusOK { - t.Errorf("expected 200 for token, got %d", wToken.Code) - } - - // 6. Change password using the active session - reqBody := `{"old_password": "oldpassword123", "new_password": "newpassword12345"}` - reqChange, _ := http.NewRequest(http.MethodPost, "/api/v1/user/change-password", strings.NewReader(reqBody)) - reqChange.Header.Set("Content-Type", "application/json") - reqChange.Header.Set("Cookie", cookieVal) - wChange := httptest.NewRecorder() - r.ServeHTTP(wChange, reqChange) - if wChange.Code != http.StatusOK { - t.Fatalf("expected change password to return 200, got %d. Body: %s", wChange.Code, wChange.Body.String()) - } - - // 7. Verification: The active session that performed change-password is now cleared (401) - reqSessionAfter, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil) - reqSessionAfter.Header.Set("Cookie", cookieVal) - wSessionAfter := httptest.NewRecorder() - r.ServeHTTP(wSessionAfter, reqSessionAfter) - if wSessionAfter.Code != http.StatusUnauthorized { - t.Errorf("expected session to be revoked (401), got %d", wSessionAfter.Code) - } - - // 8. Verification: An old session (holding an outdated password hash) should be rejected (401) - wMockOld := httptest.NewRecorder() - reqMockOld, _ := http.NewRequest(http.MethodGet, "/mock-old-session-login", nil) - r.ServeHTTP(wMockOld, reqMockOld) - oldCookieVal := wMockOld.Header().Get("Set-Cookie") - - reqOldSession, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil) - reqOldSession.Header.Set("Cookie", oldCookieVal) - wOldSession := httptest.NewRecorder() - r.ServeHTTP(wOldSession, reqOldSession) - if wOldSession.Code != http.StatusUnauthorized { - t.Errorf("expected old session with invalid hash to return 401, got %d", wOldSession.Code) - } - - // 9. Verification: The Access Token should be deleted from DB and rejected (401) - reqTokenAfter, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil) - reqTokenAfter.Header.Set("X-Access-Token", tokenStr) - wTokenAfter := httptest.NewRecorder() - r.ServeHTTP(wTokenAfter, reqTokenAfter) - if wTokenAfter.Code != http.StatusUnauthorized { - t.Errorf("expected access token to be revoked (401), got %d", wTokenAfter.Code) - } -} diff --git a/internal/apps/user/tasks.go b/internal/apps/user/tasks.go deleted file mode 100644 index 167854a6..00000000 --- a/internal/apps/user/tasks.go +++ /dev/null @@ -1,161 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package user - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "strconv" - "strings" - - "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/pkg/mail" -) - -// 异步任务名称与管理类型定义 -const ( - // SendEmailTask 发送邮件任务标识 - SendEmailTask = "mail:send" - // TaskTypeSendEmail 发送邮件管理类型 - TaskTypeSendEmail = "send_email" -) - -// SendEmailMeta represents the task metadata. -var SendEmailMeta = task.TaskMeta{ - Type: TaskTypeSendEmail, - AsynqTask: SendEmailTask, - Name: "发送邮件", - Description: "异步发送系统邮件", - SupportsTime: false, - MaxRetry: task.DefaultMaxRetry, - Queue: task.QueueDefault, - Retryable: true, - Params: []task.TaskParam{ - { - Name: "to", - Label: "接收邮箱 (To)", - Type: "string", - Required: true, - Placeholder: "receiver@example.com", - Description: "接收邮件的目标邮箱地址", - }, - { - Name: "subject", - Label: "邮件主题 (Subject)", - Type: "string", - Required: true, - Placeholder: "请输入邮件主题", - Description: "发送邮件的主题标题", - }, - { - Name: "body", - Label: "邮件内容 (Body)", - Type: "text", - Required: true, - Placeholder: "请输入邮件内容(支持 HTML 格式)", - Description: "发送邮件的内容主体", - }, - }, -} - -// SendEmailPayload 邮件发送任务载荷 -type SendEmailPayload struct { - To string `json:"to"` - Subject string `json:"subject"` - Body string `json:"body"` -} - -// SendEmailHandler 发送验证码邮件的异步任务处理器 -type SendEmailHandler struct{} - -// ValidatePayload 实现 task.PayloadValidator 接口 -// 校验并标准化邮件发送参数,框架在 Admin 下发时自动调用 -func (h *SendEmailHandler) ValidatePayload(payload []byte) ([]byte, error) { - if len(payload) == 0 { - return nil, errors.New(errTaskPayloadRequired) - } - - var req SendEmailPayload - if err := json.Unmarshal(payload, &req); err != nil { - return nil, fmt.Errorf(errInvalidJSONFormat, err) - } - - req.To = strings.TrimSpace(req.To) - req.Subject = strings.TrimSpace(req.Subject) - req.Body = strings.TrimSpace(req.Body) - - if req.To == "" || req.Subject == "" || req.Body == "" { - return nil, errors.New(errEmailTaskFieldsRequired) - } - - return json.Marshal(req) -} - -// Execute 执行邮件异步发送逻辑 -func (h *SendEmailHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) { - var req SendEmailPayload - if err := json.Unmarshal(payload, &req); err != nil { - task.AppendLog(ctx, "解析邮件发送参数失败: %v", err) - return nil, fmt.Errorf(errParseEmailPayloadFailed, err) - } - - task.AppendLog(ctx, "开始准备发送邮件到: %s, 主题: %s", req.To, req.Subject) - - // 从数据库读取最新的 SMTP 系统配置 - var smtpHost string - var smtpPortVal string - var smtpUsername string - var smtpPassword string - - if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPHost); err == nil { - smtpHost = sc.Value - } - if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPort); err == nil { - smtpPortVal = sc.Value - } - if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPUsername); err == nil { - smtpUsername = sc.Value - } - if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPassword); err == nil { - smtpPassword = sc.Value - } - - if smtpHost == "" || smtpPortVal == "" || smtpUsername == "" { - err := errors.New(errSMTPConfigIncomplete) - task.AppendLog(ctx, "发送失败: %v", err) - return nil, err - } - - smtpPort, err := strconv.Atoi(smtpPortVal) - if err != nil { - smtpPort = 587 - } - - cfg := mail.Config{ - Host: smtpHost, - Port: smtpPort, - Username: smtpUsername, - Password: smtpPassword, - } - - task.AppendLog(ctx, "连接 SMTP 服务器: %s:%d, 用户名: %s", smtpHost, smtpPort, smtpUsername) - - // 调用 SendMailHTML 执行邮件发送,这里会有 5s 拨号超时和 10s 读写限制 - err = mail.SendMailHTML(ctx, cfg, req.To, req.Subject, req.Body) - if err != nil { - task.AppendLog(ctx, "邮件发送失败: %v", err) - return nil, fmt.Errorf(errSendMailFailed, err) - } - - msg := fmt.Sprintf("邮件成功发送至: %s", req.To) - task.AppendLog(ctx, "%s", msg) - - return &task.TaskResult{ - Message: msg, - }, nil -} diff --git a/internal/cmd/app.go b/internal/cmd/app.go index 32c2f9d1..daa0dc4a 100644 --- a/internal/cmd/app.go +++ b/internal/cmd/app.go @@ -5,33 +5,26 @@ package cmd import ( "context" - "errors" - "fmt" - "log" - "net" - "net/http" - "sync" "time" "github.com/Rain-kl/Wavelet/core" "github.com/Rain-kl/Wavelet/core/extpoints" - gwrunner "github.com/Rain-kl/Wavelet/internal/apps/message_gateway/runner" "github.com/Rain-kl/Wavelet/internal/infra/config" - "github.com/Rain-kl/Wavelet/internal/infra/task/scheduler" - "github.com/Rain-kl/Wavelet/internal/infra/task/worker" - "github.com/Rain-kl/Wavelet/internal/platform/bootstrap" - "github.com/Rain-kl/Wavelet/internal/router" - "github.com/Rain-kl/Wavelet/pkg/util" "github.com/Rain-kl/Wavelet/plugins/domain/admin" "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/Rain-kl/Wavelet/plugins/domain/cap" "github.com/Rain-kl/Wavelet/plugins/domain/message_gateway" "github.com/Rain-kl/Wavelet/plugins/domain/risk_control" + "github.com/Rain-kl/Wavelet/plugins/domain/system" + "github.com/Rain-kl/Wavelet/plugins/domain/upload" "github.com/Rain-kl/Wavelet/plugins/domain/user" + "github.com/Rain-kl/Wavelet/plugins/drivers/driver_asynq_cron" + "github.com/Rain-kl/Wavelet/plugins/drivers/driver_asynq_worker" + "github.com/Rain-kl/Wavelet/plugins/drivers/driver_http" "github.com/Rain-kl/Wavelet/plugins/infra/cache" "github.com/Rain-kl/Wavelet/plugins/infra/database" "github.com/Rain-kl/Wavelet/plugins/infra/logger" "github.com/Rain-kl/Wavelet/plugins/infra/storage" - "github.com/hibiken/asynq" ) // newWaveletApp creates a core.App wired with Wavelet platform infrastructure, domain plugins, and profile drivers. @@ -43,7 +36,7 @@ func newWaveletApp(profile core.Profile) *core.App { core.WithShutdownTimeout(time.Duration(config.Config.App.GracefulShutdownTimeout)*time.Second), ) - // Register standard infrastructure plugins + // 1. Register standard infrastructure plugins app.Use( database.New(), cache.New(), @@ -51,272 +44,30 @@ func newWaveletApp(profile core.Profile) *core.App { storage.New(), ) - // Register domain plugins + // 2. Register all 8 domain business plugins app.Use( auth.New(), user.New(), message_gateway.New(), risk_control.New(), admin.New(), + upload.New(), + cap.New(), + system.New(), ) - // Bind Goose migration runner + // 3. Bind Goose migration runner app.SetMigrationRunner(func(_ context.Context, _ []extpoints.MigrationEntry) error { runMigrations() return nil }) - // Mount drivers for each aspect + // 4. Mount runtime drivers for each aspect app.Use( - newWaveletHTTPDriver(profile), - newWaveletWorkerDriver(profile), - newWaveletSchedulerDriver(profile), + driver_http.New(driver_http.WithAddr(config.Config.App.Addr)), + driver_asynq_worker.New(), + driver_asynq_cron.New(), ) return app } - -type waveletHTTPDriver struct { - profile core.Profile - server *http.Server - mu sync.Mutex - running bool -} - -func newWaveletHTTPDriver(profile core.Profile) *waveletHTTPDriver { - return &waveletHTTPDriver{profile: profile} -} - -func (d *waveletHTTPDriver) Name() string { - return "driver_wavelet_http" -} - -func (d *waveletHTTPDriver) Apply(ctx *core.Context) error { - return ctx.RegisterDriver(d) -} - -func (d *waveletHTTPDriver) Type() core.DriverType { - return core.DriverTypeHTTP -} - -//nolint:contextcheck -func (d *waveletHTTPDriver) Start(ctx context.Context) error { - d.mu.Lock() - defer d.mu.Unlock() - - if d.running { - return nil - } - - bootstrap.RegisterAPI() - runBootstrap(bootstrap.Options{API: true}) - - engine, err := router.BuildEngine() - if err != nil { - return fmt.Errorf("[API] build router engine failed: %w", err) - } - - srv := &http.Server{ - Addr: config.Config.App.Addr, - Handler: engine, - ReadHeaderTimeout: 10 * time.Second, - } - - listener, err := (&net.ListenConfig{}).Listen(ctx, "tcp", config.Config.App.Addr) - if err != nil { - return fmt.Errorf("[API] listen on %s failed: %w", config.Config.App.Addr, err) - } - - mode := "API" - if d.profile == core.ProfileAll { - mode = "API + Worker + Scheduler" - } - printStartupBanner(startupState{ - mode: mode, - relationalDB: latestMigrationState.relationalDB, - clickHouseDB: latestMigrationState.clickHouseDB, - listensForHTTP: true, - }) - - d.server = srv - d.running = true - - util.Go(func() { - log.Printf("[API] server listening on %s\n", config.Config.App.Addr) - if serveErr := srv.Serve(listener); serveErr != nil && !errors.Is(serveErr, http.ErrServerClosed) { - log.Fatalf("[API] server failed: %v\n", serveErr) - } - }) - - return nil -} - -func (d *waveletHTTPDriver) Stop(ctx context.Context) error { - d.mu.Lock() - defer d.mu.Unlock() - - if !d.running { - return nil - } - d.running = false - - var err error - if d.server != nil { - err = d.server.Shutdown(ctx) - d.server = nil - } - bootstrap.Stop(ctx) - log.Println("[API] server exited") - return err -} - -type waveletWorkerDriver struct { - profile core.Profile - server *asynq.Server - mu sync.Mutex - running bool -} - -func newWaveletWorkerDriver(profile core.Profile) *waveletWorkerDriver { - return &waveletWorkerDriver{profile: profile} -} - -func (d *waveletWorkerDriver) Name() string { - return "driver_wavelet_worker" -} - -func (d *waveletWorkerDriver) Apply(ctx *core.Context) error { - return ctx.RegisterDriver(d) -} - -func (d *waveletWorkerDriver) Type() core.DriverType { - return core.DriverTypeWorker -} - -//nolint:contextcheck -func (d *waveletWorkerDriver) Start(_ context.Context) error { - d.mu.Lock() - defer d.mu.Unlock() - - if d.running { - return nil - } - - if d.profile == core.ProfileAll { - bootstrap.RegisterAll() - } else { - bootstrap.RegisterWorker() - } - runBootstrap(bootstrap.Options{}) - - if d.profile == core.ProfileWorker { - printStartupBanner(startupState{ - mode: "Worker", - relationalDB: latestMigrationState.relationalDB, - clickHouseDB: latestMigrationState.clickHouseDB, - }) - } - - util.Go(func() { - if err := gwrunner.Start(context.Background()); err != nil { - log.Printf("[Worker] message gateway stopped: %v", err) - } - }) - - log.Println("[Worker] 启动任务处理服务") - srv, err := worker.StartWorkerServer() - if err != nil { - return fmt.Errorf("[Worker] 启动失败: %w", err) - } - - d.server = srv - d.running = true - return nil -} - -func (d *waveletWorkerDriver) Stop(_ context.Context) error { - d.mu.Lock() - defer d.mu.Unlock() - - if !d.running { - return nil - } - d.running = false - - if d.server != nil { - d.server.Stop() - d.server.Shutdown() - d.server = nil - } - log.Println("[Worker] 任务处理服务已退出") - return nil -} - -type waveletSchedulerDriver struct { - profile core.Profile - mu sync.Mutex - running bool -} - -func newWaveletSchedulerDriver(profile core.Profile) *waveletSchedulerDriver { - return &waveletSchedulerDriver{profile: profile} -} - -func (d *waveletSchedulerDriver) Name() string { - return "driver_wavelet_scheduler" -} - -func (d *waveletSchedulerDriver) Apply(ctx *core.Context) error { - return ctx.RegisterDriver(d) -} - -func (d *waveletSchedulerDriver) Type() core.DriverType { - return core.DriverTypeScheduler -} - -//nolint:contextcheck -func (d *waveletSchedulerDriver) Start(_ context.Context) error { - d.mu.Lock() - defer d.mu.Unlock() - - if d.running { - return nil - } - - if d.profile == core.ProfileAll { - bootstrap.RegisterAll() - } else { - bootstrap.RegisterScheduler() - } - runBootstrap(bootstrap.Options{}) - - if d.profile == core.ProfileSchedule { - printStartupBanner(startupState{ - mode: "Scheduler", - relationalDB: latestMigrationState.relationalDB, - clickHouseDB: latestMigrationState.clickHouseDB, - }) - } - - log.Println("[Scheduler] 启动定时任务调度服务") - if err := scheduler.ReloadScheduler(); err != nil { - return fmt.Errorf("[Scheduler] 启动失败: %w", err) - } - - d.running = true - return nil -} - -func (d *waveletSchedulerDriver) Stop(_ context.Context) error { - d.mu.Lock() - defer d.mu.Unlock() - - if !d.running { - return nil - } - d.running = false - - scheduler.StopScheduler() - log.Println("[Scheduler] 定时任务调度服务已退出") - return nil -} diff --git a/internal/cmd/app_test.go b/internal/cmd/app_test.go index 3c3b95d1..d653d256 100644 --- a/internal/cmd/app_test.go +++ b/internal/cmd/app_test.go @@ -26,9 +26,9 @@ func TestNewWaveletAppProfiles(t *testing.T) { require.NotNil(t, app) assert.Equal(t, prof, app.Profile()) - // Verify 4 infra plugins + 5 domain plugins + 3 driver plugins registered + // Verify 4 infra plugins + 8 domain plugins + 3 driver plugins registered plugins := app.Plugins() - assert.Len(t, plugins, 12) + assert.Len(t, plugins, 15) // Verify each standard infra plugin is registered _, ok := app.Plugin("database") @@ -59,14 +59,23 @@ func TestNewWaveletAppProfiles(t *testing.T) { _, ok = app.Plugin("admin") assert.True(t, ok, "admin plugin missing") + _, ok = app.Plugin("upload") + assert.True(t, ok, "upload plugin missing") + + _, ok = app.Plugin("cap") + assert.True(t, ok, "cap plugin missing") + + _, ok = app.Plugin("system") + assert.True(t, ok, "system plugin missing") + // Verify driver plugins - _, ok = app.Plugin("driver_wavelet_http") + _, ok = app.Plugin("driver_http") assert.True(t, ok, "http driver missing") - _, ok = app.Plugin("driver_wavelet_worker") + _, ok = app.Plugin("driver_asynq_worker") assert.True(t, ok, "worker driver missing") - _, ok = app.Plugin("driver_wavelet_scheduler") + _, ok = app.Plugin("driver_asynq_cron") assert.True(t, ok, "scheduler driver missing") }) } diff --git a/internal/cmd/banner.go b/internal/cmd/banner.go index 15d8f990..11892eb5 100644 --- a/internal/cmd/banner.go +++ b/internal/cmd/banner.go @@ -1,6 +1,6 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - +// Package cmd provides CLI command entry points. +// +//nolint:unused package cmd import ( @@ -14,6 +14,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/infra/persistence/migrator" ) +//nolint:unused // startup banner formatting utilities type startupState struct { mode string relationalDB migrator.Report diff --git a/internal/cmd/bootstrap.go b/internal/cmd/bootstrap.go deleted file mode 100644 index 2460cdaf..00000000 --- a/internal/cmd/bootstrap.go +++ /dev/null @@ -1,17 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package cmd - -import ( - "context" - - "github.com/Rain-kl/Wavelet/internal/platform/bootstrap" - "github.com/Rain-kl/Wavelet/pkg/trace" -) - -func runBootstrap(opts bootstrap.Options) { - ctx, span := trace.Start(context.Background(), "bootstrap.Init") - defer span.End() - bootstrap.Init(ctx, opts) -} diff --git a/internal/cmd/reset_passwd.go b/internal/cmd/reset_passwd.go index 61dc9903..6f69e16b 100644 --- a/internal/cmd/reset_passwd.go +++ b/internal/cmd/reset_passwd.go @@ -13,12 +13,11 @@ import ( "os" "strings" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/persistence/migrator" "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/platform/bootstrap" "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" "github.com/spf13/cobra" "gorm.io/gorm" ) @@ -52,7 +51,6 @@ var resetPasswdCmd = &cobra.Command{ }, Run: func(_ *cobra.Command, _ []string) { ctx := context.Background() - runBootstrap(bootstrap.Options{}) var username string if usernameFlag != "" { @@ -101,7 +99,7 @@ var resetPasswdCmd = &cobra.Command{ var tokens []model.AccessToken if err := tx.Where("user_id = ?", user.ID).Find(&tokens).Error; err == nil { for _, token := range tokens { - oauth.InvalidateCachedToken(ctx, token.TokenHash) + auth.InvalidateCachedToken(ctx, token.TokenHash) } } @@ -111,7 +109,7 @@ var resetPasswdCmd = &cobra.Command{ log.Fatalf("重置密码失败: %v\n", err) } - oauth.InvalidateCachedUser(ctx, user.ID) + auth.InvalidateCachedUser(ctx, user.ID) fmt.Println("成功重置密码!") fmt.Printf("用户名: %s\n", user.Username) diff --git a/internal/infra/task/handlers/register.go b/internal/infra/task/handlers/register.go deleted file mode 100644 index ddb8f73f..00000000 --- a/internal/infra/task/handlers/register.go +++ /dev/null @@ -1,43 +0,0 @@ -// Copyright 2025 linux.do -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package handlers 注册异步任务处理器 -package handlers - -import ( - "github.com/Rain-kl/Wavelet/internal/apps/admin/logs" - "github.com/Rain-kl/Wavelet/internal/apps/admin/push" - "github.com/Rain-kl/Wavelet/internal/apps/upload" - "github.com/Rain-kl/Wavelet/internal/apps/user" - "github.com/Rain-kl/Wavelet/internal/infra/task" -) - -// Register registers all built-in task handlers and their metadata. -func Register() { - task.RegisterHandler(upload.StorageMigrationTask, &upload.MigrationHandler{}) - task.RegisterTaskMeta(upload.StorageMigrationMeta) - - // system cleanup - task.RegisterHandler(upload.SystemCleanupTask, &upload.SystemCleanupHandler{}) - task.RegisterTaskMeta(upload.SystemCleanupMeta) - - // upload - task.RegisterHandler(upload.WarmImageCacheTask, &upload.WarmImageCacheHandler{}) - task.RegisterTaskMeta(upload.WarmImageCacheMeta) - - task.RegisterHandler(upload.RebuildUploadStatsTask, &upload.RebuildUploadStatsHandler{}) - task.RegisterTaskMeta(upload.RebuildUploadStatsMeta) - - // user - task.RegisterHandler(user.SendEmailTask, &user.SendEmailHandler{}) - task.RegisterTaskMeta(user.SendEmailMeta) - - // push - task.RegisterHandler(push.SendNotificationTask, &push.PushHandler{}) - task.RegisterTaskMeta(push.SendNotificationMeta) - - // logs - task.RegisterHandler(logs.LogDBSwitchTask, &logs.LogDBSwitchHandler{}) - task.RegisterTaskMeta(logs.LogDBSwitchMeta) -} diff --git a/internal/infra/task/meta_test.go b/internal/infra/task/meta_test.go index 40f5f14c..a5465546 100644 --- a/internal/infra/task/meta_test.go +++ b/internal/infra/task/meta_test.go @@ -7,13 +7,17 @@ import ( "testing" "github.com/Rain-kl/Wavelet/internal/infra/task" - taskhandlers "github.com/Rain-kl/Wavelet/internal/infra/task/handlers" ) func TestDuplicateTaskMeta(t *testing.T) { - // Call Register twice to simulate being imported by multiple packages (routers, worker, etc.) - taskhandlers.Register() - taskhandlers.Register() + dummyMeta := task.TaskMeta{ + Type: "test_duplicate_task", + AsynqTask: "test:duplicate_task", + Name: "Test Duplicate Task", + } + + task.RegisterTaskMeta(dummyMeta) + task.RegisterTaskMeta(dummyMeta) metas := task.GetDispatchableTasks() diff --git a/internal/infra/task/scheduler/scheduler.go b/internal/infra/task/scheduler/scheduler.go index 230a5231..0bc107e8 100644 --- a/internal/infra/task/scheduler/scheduler.go +++ b/internal/infra/task/scheduler/scheduler.go @@ -12,7 +12,7 @@ import ( "time" "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/internal/platform/bootstrap" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/logger" @@ -33,7 +33,6 @@ func GetAsynqClient() *asynq.Client { // StartScheduler 启动调度器 (该函数阻塞,直到调度器退出) func StartScheduler() error { - bootstrap.RegisterScheduler() var err error schedulerOnce.Do(func() { diff --git a/internal/infra/task/worker/worker.go b/internal/infra/task/worker/worker.go index 854c3655..3244e786 100644 --- a/internal/infra/task/worker/worker.go +++ b/internal/infra/task/worker/worker.go @@ -9,7 +9,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/infra/config" "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/internal/platform/bootstrap" + "github.com/hibiken/asynq" ) @@ -18,7 +18,7 @@ const workerShutdownTimeout = 3 * time.Minute // StartWorker 启动任务处理服务器 func StartWorker() error { - bootstrap.RegisterWorker() + asynqServer := asynq.NewServer( task.RedisOpt, asynq.Config{ @@ -46,7 +46,7 @@ func StartWorker() error { // StartWorkerServer 异步启动 Asynq 工作器服务并返回 Server 实例以支持平滑停机 func StartWorkerServer() (*asynq.Server, error) { - bootstrap.RegisterWorker() + asynqServer := asynq.NewServer( task.RedisOpt, asynq.Config{ diff --git a/internal/platform/bootstrap/bootstrap.go b/internal/platform/bootstrap/bootstrap.go deleted file mode 100644 index 858a7c1a..00000000 --- a/internal/platform/bootstrap/bootstrap.go +++ /dev/null @@ -1,240 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package bootstrap wires cross-module integrations and process-level subsystem initialization. -// All registrations use sync.Once so entry points can call them safely without import-order side effects. -package bootstrap - -import ( - "context" - "errors" - "fmt" - "log" - "sync" - - admin_push "github.com/Rain-kl/Wavelet/internal/apps/admin/push" - "github.com/Rain-kl/Wavelet/internal/apps/admin/push/custom_events" - "github.com/Rain-kl/Wavelet/internal/apps/risk_control" - "github.com/Rain-kl/Wavelet/internal/infra/config" - taskhandlers "github.com/Rain-kl/Wavelet/internal/infra/task/handlers" - "github.com/Rain-kl/Wavelet/internal/listener" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/platform/lifecycle" - "github.com/Rain-kl/Wavelet/internal/repository" - "github.com/Rain-kl/Wavelet/internal/repository/logstore" - "github.com/Rain-kl/Wavelet/pkg/cache/ram" - "github.com/Rain-kl/Wavelet/pkg/logger" - "gorm.io/gorm" -) - -// Options selects role-specific runtime bootstrap steps for the current process. -type Options struct { - // API enables HTTP-only subsystems such as the ClickHouse access-log writer. - API bool -} - -// CacheRegistry holds settings for a registered cache type. -type CacheRegistry struct { - Loader ram.Loader -} - -var ( - registerTasksOnce sync.Once - registerPushDomainEventsOnce sync.Once - registerTaskListenersOnce sync.Once - registerMessageGatewayListenersOnce sync.Once - initRuntimeOnce sync.Once - - cacheRegistries = make(map[string]CacheRegistry) - cacheRegistriesMu sync.RWMutex - - refreshLocks = make(map[string]*sync.Mutex) - refreshLocksMu sync.Mutex -) - -// RegisterCache registers a cache type with its Loader for unified preheating and refreshing. -func RegisterCache(configType string, reg CacheRegistry) { - cacheRegistriesMu.Lock() - defer cacheRegistriesMu.Unlock() - cacheRegistries[configType] = reg -} - -func getRefreshLock(configType string) *sync.Mutex { - refreshLocksMu.Lock() - defer refreshLocksMu.Unlock() - lock, found := refreshLocks[configType] - if !found { - lock = &sync.Mutex{} - refreshLocks[configType] = lock - } - return lock -} - -// PreheatAllCaches preheats all registered caches. -func PreheatAllCaches(ctx context.Context) error { - cacheRegistriesMu.RLock() - defer cacheRegistriesMu.RUnlock() - - for configType, reg := range cacheRegistries { - lock := getRefreshLock(configType) - lock.Lock() - err := ram.Refresh(ctx, configType, "", reg.Loader) - lock.Unlock() - if err != nil { - logger.ErrorF(ctx, "[Bootstrap] preheating cache type %s failed: %v", configType, err) - } - } - return nil -} - -// RegisterTasks registers all built-in task handlers and metadata. -func RegisterTasks() { - registerTasksOnce.Do(func() { - taskhandlers.Register() - }) -} - -// RegisterPushDomainEvents wires push notification handlers for domain events. -func RegisterPushDomainEvents() { - registerPushDomainEventsOnce.Do(func() { - custom_events.Register() - }) -} - -// RegisterMessageGatewayListeners registers the default log-only inbound handler. -func RegisterMessageGatewayListeners() { - registerMessageGatewayListenersOnce.Do(func() { - listener.OnMessageGatewayInbound(func(ctx context.Context, event listener.MessageGatewayInbound) { - userID := uint64(0) - if event.Msg.BindingUserID != nil { - userID = *event.Msg.BindingUserID - } - logger.InfoF(ctx, "[%s] channel=%d type=%s user=%d platform_user=%s", - listener.EventMessageGatewayInbound, - event.Msg.ChannelID, - event.Msg.ChannelType, - userID, - event.Msg.PlatformUserID, - ) - }) - }) -} - -// RegisterTaskListeners wires operational listeners to task framework hooks. -func RegisterTaskListeners() { - registerTaskListenersOnce.Do(func() { - admin_push.RegisterTaskListeners() - }) -} - -// RegisterAPI wires integrations required by the HTTP API process. -func RegisterAPI() { - RegisterTasks() - RegisterPushDomainEvents() - RegisterMessageGatewayListeners() -} - -// RegisterWorker wires integrations required by the task worker process. -func RegisterWorker() { - RegisterTasks() - RegisterTaskListeners() - RegisterMessageGatewayListeners() -} - -// RegisterScheduler wires integrations required by the task scheduler process. -func RegisterScheduler() { - RegisterTasks() -} - -// RegisterAll wires integrations for fused mode (API + Worker + Scheduler). -func RegisterAll() { - RegisterTasks() - RegisterPushDomainEvents() - RegisterTaskListeners() - RegisterMessageGatewayListeners() -} - -// Init runs shared runtime bootstrap exactly once per process. -// Call from cmd entry points after wiring registration and database migration, not from router. -func Init(ctx context.Context, opts Options) { - initRuntimeOnce.Do(func() { - if err := validateAndSeedLogDatabase(ctx); err != nil { - logger.ErrorF(ctx, "[Bootstrap] 日志主库配置校验失败: %v", err) - log.Fatalf("[Bootstrap] 日志主库配置校验失败: %v", err) - } - - logstore.SetConfigReader(func(ctx context.Context, key string) (string, error) { - cfg, err := repository.GetSystemConfigByKey(ctx, key) - if err != nil { - return "", err - } - return cfg.Value, nil - }) - logstore.Init(ctx) - - // Register config cache loader - RegisterCache(repository.ConfigCacheType, CacheRegistry{ - Loader: repository.ConfigLoader{}, - }) - - // Preheat config cache initially (using PreheatAllCaches) - if err := PreheatAllCaches(ctx); err != nil { - logger.ErrorF(ctx, "[Bootstrap] preheating all caches failed: %v", err) - } - - if err := admin_push.SyncEvents(ctx); err != nil { - logger.ErrorF(ctx, "[Bootstrap] sync push events failed: %v", err) - } - if opts.API { - risk_control.InitLogWriter(ctx) - } - }) -} - -func validateAndSeedLogDatabase(ctx context.Context) error { - cfg, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyLogDatabase) - if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { - return fmt.Errorf("读取日志主库配置失败: %w", err) - } - current := cfg.Value - if current == "" { - current = "sqlite" - if config.Config.Database.Enabled { - current = "postgres" - } - if config.Config.ClickHouse.Enabled { - current = "clickhouse" - } - if err := repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyLogDatabase, current); err != nil { - return fmt.Errorf("初始化日志主库配置失败: %w", err) - } - return nil - } - switch current { - case "clickhouse": - if !config.Config.ClickHouse.Enabled { - return errors.New("当前日志主库为 ClickHouse 但 ClickHouse 未启用。请先重新启用 ClickHouse 配置并启动,在任务管理运行「切换日志数据库」迁移到 PostgreSQL/SQLite 后再禁用 ClickHouse") - } - case "postgres": - if !config.Config.Database.Enabled { - return errors.New("当前日志主库为 PostgreSQL 但 PostgreSQL 未启用(当前为 SQLite 主库)。请运行「切换日志数据库」迁回 SQLite 或启用 PostgreSQL") - } - case "sqlite": - if config.Config.Database.Enabled { - return errors.New("当前日志主库为 SQLite 但当前主库为 PostgreSQL。请运行「切换日志数据库」迁移到 PostgreSQL") - } - default: - return fmt.Errorf("未知的日志主库配置: %s", current) - } - return nil -} - -// Stop stops all batch writers and background resources. -func Stop(ctx context.Context) { - lifecycle.Stop(ctx) -} - -// ResetInitRuntimeOnceForTest clears initRuntimeOnce so Init can run again in unit tests. -func ResetInitRuntimeOnceForTest() { - initRuntimeOnce = sync.Once{} -} diff --git a/internal/platform/bootstrap/bootstrap_test.go b/internal/platform/bootstrap/bootstrap_test.go deleted file mode 100644 index f9a28fc5..00000000 --- a/internal/platform/bootstrap/bootstrap_test.go +++ /dev/null @@ -1,52 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package bootstrap - -import ( - "context" - "testing" - - admin_push "github.com/Rain-kl/Wavelet/internal/apps/admin/push" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/testhelper" -) - -func TestInitSyncsPushEventsOnce(t *testing.T) { - ResetInitRuntimeOnceForTest() - t.Cleanup(ResetInitRuntimeOnceForTest) - - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - if err := dbConn.AutoMigrate(&model.PushEvent{}); err != nil { - t.Fatalf("auto migrate push events failed: %v", err) - } - - RegisterPushDomainEvents() - - wantCount := len(admin_push.BuiltInEvents) - if wantCount < 1 { - t.Fatalf("built-in push events = %d, want at least 1", wantCount) - } - - ctx := context.Background() - Init(ctx, Options{API: true}) - Init(ctx, Options{}) // second Init must not duplicate events (initRuntimeOnce) - - var count int64 - if err := dbConn.Model(&model.PushEvent{}).Count(&count).Error; err != nil { - t.Fatalf("count push events failed: %v", err) - } - if count != int64(wantCount) { - t.Fatalf("push event count = %d, want %d", count, wantCount) - } - - var adminLogin model.PushEvent - if err := dbConn.Where("event_key = ?", "admin_login").First(&adminLogin).Error; err != nil { - t.Fatalf("admin_login event not found after Init: %v", err) - } - if adminLogin.Name != "管理员登录" { - t.Fatalf("admin_login name = %q, want %q", adminLogin.Name, "管理员登录") - } -} diff --git a/internal/router/root/custom.go b/internal/router/root/custom.go deleted file mode 100644 index 15c5be1c..00000000 --- a/internal/router/root/custom.go +++ /dev/null @@ -1,16 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package root registers custom business routes and frontend serving. -package root - -import ( - "github.com/gin-gonic/gin" -) - -// RegisterCustomRootRoutes is a scaffold SAMPLE placeholder for root-path routes -// (webhooks, short links). Prefer semantic paths and/or a dedicated root file; -// do not treat this as the only place for all product APIs. See skill new-api. -func RegisterCustomRootRoutes(_ *gin.Engine) { - // Sample only — add root-path demos here if needed -} diff --git a/internal/router/root/default.go b/internal/router/root/default.go deleted file mode 100644 index d19baddc..00000000 --- a/internal/router/root/default.go +++ /dev/null @@ -1,32 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package root - -import ( - _ "github.com/Rain-kl/Wavelet/docs" // Swagger documentation generation setup - publicconfig "github.com/Rain-kl/Wavelet/internal/apps/config" - "github.com/Rain-kl/Wavelet/internal/apps/health" - "github.com/Rain-kl/Wavelet/internal/apps/upload" - "github.com/Rain-kl/Wavelet/internal/infra/config" - "github.com/gin-gonic/gin" - swaggerFiles "github.com/swaggo/files" - ginSwagger "github.com/swaggo/gin-swagger" -) - -// RegisterDefaultRootRoutes registers default routes that belong to the root path. -func RegisterDefaultRootRoutes(r *gin.Engine) { - // 1. Serve files by ID - r.GET("/f/:id", upload.ServeFileByID) - - // 2. Dynamic robots.txt serving - r.GET("/robots.txt", publicconfig.GetRobotsTXT) - - // 3. Swagger routes (Non-production only) - if !config.Config.App.IsProduction() { - r.GET(config.Config.App.APIPrefix+"/swagger/*any", ginSwagger.WrapHandler(swaggerFiles.Handler)) - } - - // 4. Health check - r.GET(config.Config.App.APIPrefix+"/health", health.Health) -} diff --git a/internal/router/root/frontend.go b/internal/router/root/frontend.go deleted file mode 100644 index cd357d66..00000000 --- a/internal/router/root/frontend.go +++ /dev/null @@ -1,110 +0,0 @@ -//go:build embed_frontend - -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package root - -import ( - "embed" - "io" - "io/fs" - "net/http" - "strings" - - "github.com/gin-gonic/gin" -) - -//go:embed all:dist -var frontendFS embed.FS - -func serveFileDirect(c *gin.Context, subFS fs.FS, filePath string) bool { - file, err := subFS.Open(filePath) - if err != nil { - return false - } - defer file.Close() - - stat, err := file.Stat() - if err != nil { - return false - } - - if stat.IsDir() { - return false - } - - seeker, ok := file.(io.ReadSeeker) - if !ok { - return false - } - - // 使用 http.ServeContent 直接输出文件内容,不进行路径规范化重定向 - http.ServeContent(c.Writer, c.Request, filePath, stat.ModTime(), seeker) - return true -} - -func init() { - RegisterFrontend = func(r *gin.Engine) { - subFS, err := fs.Sub(frontendFS, "dist") - if err != nil { - panic(err) - } - - r.NoRoute(func(c *gin.Context) { - path := c.Request.URL.Path - - // API 接口路由或文件服务路由 -> 直接返回,由 Gin 处理标准 404 - if strings.HasPrefix(path, "/api/") || strings.HasPrefix(path, "/f/") { - return - } - - // 只处理 GET 和 HEAD 请求 - if c.Request.Method != http.MethodGet && c.Request.Method != http.MethodHead { - c.JSON(http.StatusMethodNotAllowed, gin.H{"error_msg": "Method not allowed"}) - return - } - - // 移除开头的斜杠以在嵌入文件系统中查找 - cleanPath := strings.TrimPrefix(path, "/") - - // 1. 根路径 -> 直接输出 index.html - if cleanPath == "" { - if serveFileDirect(c, subFS, "index.html") { - return - } - } - - // 2. 精确匹配(如果对应的文件存在,直接输出) - if serveFileDirect(c, subFS, cleanPath) { - return - } - - // 如果是个目录(例如请求了 "/login",同时 dist 目录下存在一个叫 "login" 的文件夹目录), - // 则查找是否有对应的 ".html" 文件(例如 "login.html")并进行输出。 - if cleanPath != "" { - htmlPath := cleanPath + ".html" - if serveFileDirect(c, subFS, htmlPath) { - return - } - } - - // 3. Next.js Clean URLs 兜底逻辑(例如访问 /settings/security -> 实际映射输出 settings/security.html) - if !strings.Contains(cleanPath, ".") { - htmlPath := cleanPath + ".html" - if serveFileDirect(c, subFS, htmlPath) { - return - } - indexPath := cleanPath + "/index.html" - if serveFileDirect(c, subFS, indexPath) { - return - } - } - - // 4. 单页应用(SPA)前端路由兜底:返回 index.html - if serveFileDirect(c, subFS, "index.html") { - return - } - }) - } -} diff --git a/internal/router/root/root.go b/internal/router/root/root.go deleted file mode 100644 index a6700b61..00000000 --- a/internal/router/root/root.go +++ /dev/null @@ -1,26 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package root registers custom business routes and frontend serving. -package root - -import ( - "github.com/gin-gonic/gin" -) - -// RegisterFrontend is a package-level variable overridden by frontend.go when built with embed_frontend. -var RegisterFrontend = func(_ *gin.Engine) { - // No-op by default -} - -// RegisterRootRoutes registers custom business routes that belong to the root path. -func RegisterRootRoutes(r *gin.Engine) { - // 1. Default root routes (/f/:id, /robots.txt, and /swagger/*any) - RegisterDefaultRootRoutes(r) - - // 2. Register custom serving - RegisterCustomRootRoutes(r) - - // 3. Register frontend serving - RegisterFrontend(r) -} diff --git a/internal/router/router.go b/internal/router/router.go index 0130dc2a..205d789c 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -16,15 +16,11 @@ import ( "syscall" "time" - "github.com/Rain-kl/Wavelet/internal/apps/risk_control" - "github.com/Rain-kl/Wavelet/internal/platform/bootstrap" - router_root "github.com/Rain-kl/Wavelet/internal/router/root" - v1 "github.com/Rain-kl/Wavelet/internal/router/v1" - - "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/infra/config" "github.com/Rain-kl/Wavelet/pkg/trace" "github.com/Rain-kl/Wavelet/pkg/util" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/Rain-kl/Wavelet/plugins/domain/risk_control" "github.com/gin-contrib/sessions" "github.com/gin-contrib/sessions/redis" "github.com/gin-gonic/gin" @@ -70,14 +66,13 @@ func BuildEngine() (*gin.Engine, error) { } } - sessionStore.Options(oauth.GetSessionOptions(config.Config.App.SessionAge)) + sessionStore.Options(auth.GetSessionOptions(config.Config.App.SessionAge)) r.Use(sessions.Sessions(config.Config.App.SessionCookieName, sessionStore)) // 补充中间件 - r.Use(otelgin.Middleware(config.Config.App.AppName), errorHandlerMiddleware(), loggerMiddleware(), risk_control.RiskControlMiddleware()) + r.Use(otelgin.Middleware(config.Config.App.AppName), errorHandlerMiddleware(), loggerMiddleware(), risk_control.Middleware()) - registerRoutes(r) return r, nil } @@ -119,26 +114,10 @@ func Serve(onStarted func()) { if err := srv.Shutdown(shutdownCtx); err != nil { log.Printf("[API] server forced to shutdown: %v\n", err) - bootstrap.Stop(shutdownCtx) cancel() os.Exit(1) } - bootstrap.Stop(shutdownCtx) cancel() log.Println("[API] server exited") } - -func registerRoutes(r *gin.Engine) { - // Register custom root routes, Swagger, and frontend serving - router_root.RegisterRootRoutes(r) - - apiGroup := r.Group(config.Config.App.APIPrefix) - { - // API V1 - apiV1Router := apiGroup.Group("/v1") - { - v1.RegisterV1Routes(apiV1Router, apiGroup) - } - } -} diff --git a/internal/router/v1/admin.go b/internal/router/v1/admin.go deleted file mode 100644 index 1f908ca9..00000000 --- a/internal/router/v1/admin.go +++ /dev/null @@ -1,220 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package v1 contains router registrations for API V1 -package v1 - -import ( - "github.com/Rain-kl/Wavelet/internal/apps/admin" - admin_auth_source "github.com/Rain-kl/Wavelet/internal/apps/admin/auth_source" - admin_cache "github.com/Rain-kl/Wavelet/internal/apps/admin/cache" - admin_db_manage "github.com/Rain-kl/Wavelet/internal/apps/admin/db_manage" - admin_logs "github.com/Rain-kl/Wavelet/internal/apps/admin/logs" - admin_push "github.com/Rain-kl/Wavelet/internal/apps/admin/push" - admin_status "github.com/Rain-kl/Wavelet/internal/apps/admin/status" - "github.com/Rain-kl/Wavelet/internal/apps/admin/system_config" - admin_task "github.com/Rain-kl/Wavelet/internal/apps/admin/task" - admin_template "github.com/Rain-kl/Wavelet/internal/apps/admin/template" - admin_updater "github.com/Rain-kl/Wavelet/internal/apps/admin/updater" - admin_user "github.com/Rain-kl/Wavelet/internal/apps/admin/user" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/apps/upload" - "github.com/gin-gonic/gin" -) - -// RegisterAdminRoutes registers all admin-related routes with sub-group categorizations. -func RegisterAdminRoutes(apiV1Router *gin.RouterGroup) { - adminRouter := apiV1Router.Group("/admin") - adminRouter.Use(oauth.LoginRequired(), admin.LoginAdminRequired()) - { - // 1. Diagnostics & Infrastructure Management - registerAdminDiagnosticRoutes(adminRouter) - - // 2. Identity & Access Management (IAM) - registerAdminIAMRoutes(adminRouter) - - // 3. System Configuration & Templates Settings - registerAdminConfigRoutes(adminRouter) - - // 4. Storage & Asset Management - registerAdminStorageRoutes(adminRouter) - - // 5. Task Orchestration & Automation - registerAdminTaskRoutes(adminRouter) - - // 6. Messaging & Push Notifications - registerAdminPushRoutes(adminRouter) - registerAdminMessageGatewayRoutes(adminRouter) - } -} - -// registerAdminDiagnosticRoutes registers infrastructure, system status, caching, database, and logs diagnostics. -func registerAdminDiagnosticRoutes(adminRouter *gin.RouterGroup) { - // System status - adminRouter.GET("/status", admin_status.GetSystemStatus) - adminRouter.GET("/status/log-database", admin_status.GetLogDatabaseStatus) - - // Database basic info & backup export - adminRouter.GET("/db-info", admin_status.GetDatabaseInfo) - adminRouter.GET("/db-export", admin_status.ExportDatabase) - - // Database management & interactive browser - dbManage := adminRouter.Group("/db-manage") - { - dbManage.GET("/overview", admin_db_manage.GetDBOverview) - dbManage.GET("/tables", admin_db_manage.ListDBTables) - dbManage.GET("/table-data", admin_db_manage.GetDBTableData) - dbManage.POST("/query", admin_db_manage.ExecuteSQL) - } - - // Cache management (TTL, LRU eviction and clear operations) - cache := adminRouter.Group("/cache") - { - cache.GET("/status", admin_cache.GetCacheStatus) - cache.POST("/config", admin_cache.UpdateCacheConfig) - cache.POST("/clear", admin_cache.ClearCache) - } - - // Application updater - update := adminRouter.Group("/update") - { - update.GET("", admin_updater.GetUpdateStatus) - update.POST("/apply", admin_updater.ApplyUpdate) - } - - // System & access logs analytics - logs := adminRouter.Group("/logs") - { - logs.GET("", admin_logs.GetLogs) - logs.GET("/access", admin_logs.GetAccessLogs) - logs.GET("/analytics", admin_logs.GetLogsAnalytics) - logs.GET("/ws", admin_logs.HandleLogWebSocket) - } -} - -// registerAdminIAMRoutes registers Identity & Access Management endpoints (Users & Auth Sources). -func registerAdminIAMRoutes(adminRouter *gin.RouterGroup) { - // Users management - users := adminRouter.Group("/users") - { - users.GET("", admin_user.ListUsers) - users.POST("", admin_user.CreateUser) - users.GET("/:id", admin_user.GetUser) - users.PUT("/:id/status", admin_user.UpdateUserStatus) - users.PUT("/:id", admin_user.UpdateUser) - users.DELETE("/:id", admin_user.DeleteUser) - } - - // Authentication Sources (LDAP, OAuth sources, etc.) - authSources := adminRouter.Group("/auth-sources") - { - authSources.GET("", admin_auth_source.ListAuthSources) - authSources.POST("", admin_auth_source.CreateAuthSource) - authSources.PUT("/:id", admin_auth_source.UpdateAuthSource) - authSources.PUT("/:id/toggle", admin_auth_source.ToggleAuthSource) - authSources.DELETE("/:id", admin_auth_source.DeleteAuthSource) - } -} - -// registerAdminConfigRoutes registers system configurations and template settings. -func registerAdminConfigRoutes(adminRouter *gin.RouterGroup) { - // System configs - configs := adminRouter.Group("/system-configs") - { - configs.GET("", system_config.ListSystemConfigs) - configs.POST("", system_config.CreateSystemConfig) - configs.POST("/smtp/test", system_config.TestSMTP) - - keyGroup := configs.Group("/:key") - { - keyGroup.GET("", system_config.GetSystemConfig) - keyGroup.PUT("", system_config.UpdateSystemConfig) - } - } - - // Email/Notification Templates - templates := adminRouter.Group("/templates") - { - templates.GET("", admin_template.ListTemplates) - templates.POST("", admin_template.CreateTemplate) - - keyGroup := templates.Group("/:key") - { - keyGroup.GET("", admin_template.GetTemplate) - keyGroup.PUT("", admin_template.UpdateTemplate) - keyGroup.DELETE("", admin_template.DeleteTemplate) - } - } -} - -// registerAdminStorageRoutes registers file and asset storage administration. -func registerAdminStorageRoutes(adminRouter *gin.RouterGroup) { - uploads := adminRouter.Group("/uploads") - { - uploads.GET("", upload.ListFiles) - uploads.GET("/stats", upload.GetFileStats) - uploads.DELETE("/:id", upload.DeleteFile) - uploads.GET("/download/:id", upload.DownloadFile) - uploads.POST("/download/batch", upload.BatchDownloadFiles) - uploads.GET("/types", upload.GetDistinctUploadTypes) - } -} - -// registerAdminTaskRoutes registers task orchestrations, execution logs and schedules. -func registerAdminTaskRoutes(adminRouter *gin.RouterGroup) { - tasks := adminRouter.Group("/tasks") - { - // Task dispatch & metadata - tasks.GET("/types", admin_task.ListTaskTypes) - tasks.POST("/dispatch", admin_task.DispatchTask) - - // Task execution logs & manual retry - executions := tasks.Group("/executions") - { - executions.GET("", admin_task.ListTaskExecutions) - executions.GET("/:id", admin_task.GetTaskExecution) - executions.POST("/:id/retry", admin_task.RetryTask) - } - - // Cron scheduler settings - schedules := tasks.Group("/schedules") - { - schedules.GET("", admin_task.ListSchedules) - schedules.POST("", admin_task.CreateSchedule) - schedules.PUT("/:id", admin_task.UpdateSchedule) - schedules.DELETE("/:id", admin_task.DeleteSchedule) - } - } -} - -// registerAdminPushRoutes registers messaging channels and push notification events. -func registerAdminPushRoutes(adminRouter *gin.RouterGroup) { - push := adminRouter.Group("/push") - { - // Push Events - events := push.Group("/events") - { - events.GET("", admin_push.ListEvents) - events.GET("/builtin", admin_push.ListBuiltInEvents) - events.POST("", admin_push.CreateEvent) - events.PUT("/:id", admin_push.UpdateEvent) - events.DELETE("/:id", admin_push.DeleteEvent) - events.POST("/:id/toggle", admin_push.ToggleEvent) - } - - // Delivery histories and diagnostics test - push.GET("/histories", admin_push.ListHistories) - push.POST("/test", admin_push.TestPush) - - // Message Channels CRUD - channels := push.Group("/channels") - { - channels.GET("/definitions", admin_push.ListChannelDefinitions) - channels.GET("", admin_push.ListChannels) - channels.POST("", admin_push.CreateChannel) - channels.PUT("/:id", admin_push.UpdateChannel) - channels.DELETE("/:id", admin_push.DeleteChannel) - channels.POST("/test", admin_push.TestChannel) - } - } -} diff --git a/internal/router/v1/custom.go b/internal/router/v1/custom.go deleted file mode 100644 index 7a9364da..00000000 --- a/internal/router/v1/custom.go +++ /dev/null @@ -1,20 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package v1 contains router registrations for API V1 -package v1 - -import ( - "github.com/Rain-kl/Wavelet/internal/apps/custom" - "github.com/gin-gonic/gin" -) - -// RegisterCustomRoutes is a scaffold SAMPLE only (demo: GET /api/v1/custom/hello). -// Real product APIs belong in apps// with semantic paths and a dedicated -// Register*Routes file (e.g. channel.go), not piled into this package. See skill new-api. -func RegisterCustomRoutes(apiV1Router *gin.RouterGroup) { - customRouter := apiV1Router.Group("/custom") - { - customRouter.GET("/hello", custom.Hello) - } -} diff --git a/internal/router/v1/message_gateway.go b/internal/router/v1/message_gateway.go deleted file mode 100644 index 19c1545a..00000000 --- a/internal/router/v1/message_gateway.go +++ /dev/null @@ -1,14 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package v1 - -import ( - appgw "github.com/Rain-kl/Wavelet/internal/apps/message_gateway" - "github.com/gin-gonic/gin" -) - -// RegisterMessageGatewayUserRoutes mounts user bind/unbind APIs. -func RegisterMessageGatewayUserRoutes(apiV1Router *gin.RouterGroup) { - appgw.RegisterUserRoutes(apiV1Router) -} diff --git a/internal/router/v1/message_gateway_admin.go b/internal/router/v1/message_gateway_admin.go deleted file mode 100644 index 58f3b412..00000000 --- a/internal/router/v1/message_gateway_admin.go +++ /dev/null @@ -1,13 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package v1 - -import ( - adminmsg "github.com/Rain-kl/Wavelet/internal/apps/admin/message_gateway" - "github.com/gin-gonic/gin" -) - -func registerAdminMessageGatewayRoutes(adminRouter *gin.RouterGroup) { - adminmsg.RegisterRoutes(adminRouter) -} diff --git a/internal/router/v1/user.go b/internal/router/v1/user.go deleted file mode 100644 index 662a8b23..00000000 --- a/internal/router/v1/user.go +++ /dev/null @@ -1,95 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package v1 contains router registrations for API V1 -package v1 - -import ( - capApp "github.com/Rain-kl/Wavelet/internal/apps/cap" - publicconfig "github.com/Rain-kl/Wavelet/internal/apps/config" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/apps/upload" - "github.com/Rain-kl/Wavelet/internal/apps/user" - "github.com/gin-gonic/gin" -) - -// RegisterUserRoutes registers all user-related, oauth, upload, and public routes. -func RegisterUserRoutes(apiV1Router *gin.RouterGroup, apiGroup *gin.RouterGroup) { - // 1. CAPTCHA - registerCaptchaRoutes(apiGroup) - - // 2. Config (public) - registerConfigRoutes(apiV1Router) - - // 3. OAuth - registerOAuthRoutes(apiV1Router) - - // 4. User - registerUserRoutes(apiV1Router) - - // 5. Upload - registerUploadRoutes(apiV1Router) -} - -func registerCaptchaRoutes(apiGroup *gin.RouterGroup) { - capGroup := apiGroup.Group("/cap") - { - capGroup.POST("/challenge", capApp.Challenge) - capGroup.POST("/redeem", capApp.Redeem) - } -} - -func registerConfigRoutes(apiV1Router *gin.RouterGroup) { - configRouter := apiV1Router.Group("/config") - { - configRouter.GET("/public", publicconfig.GetPublicConfig) - } -} - -func registerOAuthRoutes(apiV1Router *gin.RouterGroup) { - apiV1Router.GET("/oauth/sources", oauth.GetLoginSources) - apiV1Router.GET("/oauth/login", oauth.GetLoginURL) - apiV1Router.GET("/oauth/:source/authorize", oauth.Authorize) - apiV1Router.GET("/oauth/logout", oauth.Logout) - apiV1Router.POST("/oauth/callback", oauth.Callback) - apiV1Router.GET("/oauth/user-info", oauth.LoginRequired(), oauth.UserInfo) - apiV1Router.GET("/user-info", oauth.LoginRequired(), oauth.UserInfo) - apiV1Router.GET("/oauth/external-accounts", oauth.LoginRequired(), oauth.ListExternalAccounts) - apiV1Router.POST("/oauth/external-accounts/:id/delete", oauth.LoginRequired(), oauth.DeleteExternalAccount) -} - -func registerUserRoutes(apiV1Router *gin.RouterGroup) { - userRouter := apiV1Router.Group("/user") - { - userRouter.POST("/login", capApp.VerifyMiddleware(capApp.GetDefaultManager(), "login"), user.Login) - userRouter.POST("/register", capApp.VerifyMiddleware(capApp.GetDefaultManager(), "register"), user.Register) - userRouter.POST("/send-email-code", capApp.VerifyMiddleware(capApp.GetDefaultManager(), "send_email_code"), user.SendEmailCode) - userRouter.GET("/logout", user.Logout) - userRouter.GET("/self", oauth.LoginRequired(), oauth.UserInfo) - userRouter.POST("/change-password", oauth.LoginRequired(), user.ChangePassword) - userRouter.PUT("/profile", oauth.LoginRequired(), user.UpdateProfile) - - // Access Token - tokenRouter := userRouter.Group("/access-tokens") - tokenRouter.Use(oauth.LoginRequired(), oauth.DisallowTokenAuth()) - { - tokenRouter.GET("", user.ListAccessTokens) - tokenRouter.POST("", user.CreateAccessToken) - tokenRouter.DELETE("/:id", user.DeleteAccessToken) - tokenRouter.POST("/:id/rotate", user.RotateAccessToken) - } - } -} - -func registerUploadRoutes(apiV1Router *gin.RouterGroup) { - uploadRouter := apiV1Router.Group("/upload") - uploadRouter.Use(oauth.LoginRequired()) - { - uploadRouter.POST("", upload.UploadFile) - uploadRouter.GET("/my", upload.ListMyFiles) - uploadRouter.DELETE("/:id", upload.DeleteMyFile) - uploadRouter.PUT("/:id", upload.UpdateMyFile) - uploadRouter.GET("/download/:id", upload.DownloadFile) - uploadRouter.POST("/download/batch", upload.BatchDownloadFiles) - } -} diff --git a/internal/router/v1/v1.go b/internal/router/v1/v1.go deleted file mode 100644 index 295006cb..00000000 --- a/internal/router/v1/v1.go +++ /dev/null @@ -1,25 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package v1 contains router registrations for API V1 -package v1 - -import ( - "github.com/gin-gonic/gin" -) - -// RegisterV1Routes registers all routes under API V1. -func RegisterV1Routes(apiV1Router *gin.RouterGroup, apiGroup *gin.RouterGroup) { - // 1. User & Public routes (OAuth, User, Upload, CAPTCHA, Health, Config) - RegisterUserRoutes(apiV1Router, apiGroup) - - // 2. Admin routes - RegisterAdminRoutes(apiV1Router) - - // 3. Message gateway user bind/unbind - RegisterMessageGatewayUserRoutes(apiV1Router) - - // 3. Product domain routes: RegisterXxxRoutes(apiV1Router) — see skill new-api - // 4. Scaffold sample only (optional demo under /api/v1/custom) - RegisterCustomRoutes(apiV1Router) -} diff --git a/plugins/domain/admin/handlers_auth_source.go b/plugins/domain/admin/handlers_auth_source.go new file mode 100644 index 00000000..723c38cd --- /dev/null +++ b/plugins/domain/admin/handlers_auth_source.go @@ -0,0 +1,157 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package admin + +import ( + "net/http" + "strconv" + + persistence "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/shared/response" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/gin-gonic/gin" +) + +// ListAuthSources lists all configured authentication sources. +func ListAuthSources(c *gin.Context) { + var sources []auth.AuthSource + gormDB := persistence.DB(c.Request.Context()) + if err := gormDB.Order("id ASC").Find(&sources).Error; err != nil { + response.AbortInternal(c, "获取认证源列表失败") + return + } + + views := make([]auth.AuthSourceView, len(sources)) + for i := range sources { + views[i] = auth.AuthSourceView{ + 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 != "", + } + } + + c.JSON(http.StatusOK, response.OK(views)) +} + +// CreateAuthSource creates a new authentication source. +func CreateAuthSource(c *gin.Context) { + var source auth.AuthSource + if err := c.ShouldBindJSON(&source); err != nil { + response.AbortBadRequest(c, "无效的参数") + return + } + + if err := source.Validate(); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + + gormDB := persistence.DB(c.Request.Context()) + if err := gormDB.Create(&source).Error; err != nil { + response.AbortBadRequest(c, "创建认证源失败: "+err.Error()) + return + } + + source.Sanitize() + c.JSON(http.StatusOK, response.OK(source)) +} + +// UpdateAuthSource updates an authentication source. +func UpdateAuthSource(c *gin.Context) { + idStr := c.Param("id") + id, err := strconv.ParseUint(idStr, 10, 64) + if err != nil { + response.AbortBadRequest(c, "无效的认证源 ID") + return + } + + gormDB := persistence.DB(c.Request.Context()) + var existing auth.AuthSource + if err := gormDB.First(&existing, id).Error; err != nil { + response.AbortNotFound(c, "认证源不存在") + return + } + + var req auth.AuthSource + if err := c.ShouldBindJSON(&req); err != nil { + response.AbortBadRequest(c, "无效的参数") + return + } + + existing.DisplayName = req.DisplayName + existing.ClientID = req.ClientID + if req.ClientSecret != "" { + existing.ClientSecret = req.ClientSecret + } + existing.OpenIDDiscoveryURL = req.OpenIDDiscoveryURL + existing.Scopes = req.Scopes + existing.IconURL = req.IconURL + + if err := existing.Validate(); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + + if err := gormDB.Save(&existing).Error; err != nil { + response.AbortInternal(c, "更新认证源失败") + return + } + + existing.Sanitize() + c.JSON(http.StatusOK, response.OK(existing)) +} + +// ToggleAuthSource toggles the active state of an auth source. +func ToggleAuthSource(c *gin.Context) { + idStr := c.Param("id") + id, err := strconv.ParseUint(idStr, 10, 64) + if err != nil { + response.AbortBadRequest(c, "无效的认证源 ID") + return + } + + gormDB := persistence.DB(c.Request.Context()) + var existing auth.AuthSource + if err := gormDB.First(&existing, id).Error; err != nil { + response.AbortNotFound(c, "认证源不存在") + return + } + + existing.IsActive = !existing.IsActive + if existing.IsActive { + if err := existing.Validate(); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + } + + if err := gormDB.Model(&existing).Update("is_active", existing.IsActive).Error; err != nil { + response.AbortInternal(c, "切换认证源状态失败") + return + } + + c.JSON(http.StatusOK, response.OK(gin.H{"is_active": existing.IsActive})) +} + +// DeleteAuthSource deletes an authentication source. +func DeleteAuthSource(c *gin.Context) { + idStr := c.Param("id") + id, err := strconv.ParseUint(idStr, 10, 64) + if err != nil { + response.AbortBadRequest(c, "无效的认证源 ID") + return + } + + gormDB := persistence.DB(c.Request.Context()) + if err := gormDB.Delete(&auth.AuthSource{}, id).Error; err != nil { + response.AbortInternal(c, "删除认证源失败") + return + } + + c.JSON(http.StatusOK, response.OKNil()) +} diff --git a/plugins/domain/admin/handlers_config.go b/plugins/domain/admin/handlers_config.go index a7ab4622..4803273f 100644 --- a/plugins/domain/admin/handlers_config.go +++ b/plugins/domain/admin/handlers_config.go @@ -13,11 +13,6 @@ import ( "strings" "time" - "github.com/gin-gonic/gin" - "gorm.io/gorm" - - "github.com/Rain-kl/Wavelet/internal/apps/cap" - "github.com/Rain-kl/Wavelet/internal/apps/upload" "github.com/Rain-kl/Wavelet/internal/infra/objectstore" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" @@ -25,6 +20,10 @@ import ( "github.com/Rain-kl/Wavelet/internal/shared/response" "github.com/Rain-kl/Wavelet/pkg/logger" mail "github.com/Rain-kl/Wavelet/pkg/mail" + "github.com/Rain-kl/Wavelet/plugins/domain/cap" + "github.com/Rain-kl/Wavelet/plugins/domain/upload" + "github.com/gin-gonic/gin" + "gorm.io/gorm" ) const maskedConfigValue = "******" diff --git a/plugins/domain/admin/handlers_logs.go b/plugins/domain/admin/handlers_logs.go index 06ea4f82..c503ee6c 100644 --- a/plugins/domain/admin/handlers_logs.go +++ b/plugins/domain/admin/handlers_logs.go @@ -15,10 +15,6 @@ import ( "strings" "time" - "github.com/gin-gonic/gin" - "github.com/gorilla/websocket" - - "github.com/Rain-kl/Wavelet/internal/apps/risk_control" "github.com/Rain-kl/Wavelet/internal/infra/config" "github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/model" @@ -27,6 +23,9 @@ import ( "github.com/Rain-kl/Wavelet/internal/shared/response" "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/util" + "github.com/Rain-kl/Wavelet/plugins/domain/risk_control" + "github.com/gin-gonic/gin" + "github.com/gorilla/websocket" ) const ( diff --git a/plugins/domain/admin/handlers_user.go b/plugins/domain/admin/handlers_user.go index 55449a07..ded5200a 100644 --- a/plugins/domain/admin/handlers_user.go +++ b/plugins/domain/admin/handlers_user.go @@ -15,12 +15,12 @@ import ( "github.com/gin-gonic/gin" "gorm.io/gorm" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/shared/response" "github.com/Rain-kl/Wavelet/pkg/logger" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" ) const minPasswordLength = 8 @@ -243,7 +243,7 @@ func DeleteUser(c *gin.Context) { return } - currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey) if currUser == nil { response.AbortUnauthorized(c, AdminRequired) return @@ -334,7 +334,7 @@ func UpdateUser(c *gin.Context) { return } - currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey) if currUser == nil { response.AbortUnauthorized(c, AdminRequired) return @@ -388,10 +388,10 @@ func updateUserStatus(ctx context.Context, id uint64, active bool) error { err = repository.UpdateUserActive(ctx, id, active) if err == nil { - oauth.InvalidateCachedUser(ctx, id) + auth.InvalidateCachedUser(ctx, id) if !active { for _, token := range tokens { - oauth.InvalidateCachedToken(ctx, token.TokenHash) + auth.InvalidateCachedToken(ctx, token.TokenHash) } } } @@ -414,9 +414,9 @@ func deleteUser(ctx context.Context, currentUserID, targetID uint64) error { err = repository.DeleteUserWithRelations(ctx, targetID) if err == nil { - oauth.InvalidateCachedUser(ctx, targetID) + auth.InvalidateCachedUser(ctx, targetID) for _, token := range tokens { - oauth.InvalidateCachedToken(ctx, token.TokenHash) + auth.InvalidateCachedToken(ctx, token.TokenHash) } } return err @@ -539,10 +539,10 @@ func updateUser(ctx context.Context, currentUserID uint64, param updateUserParam err = repository.UpdateUser(ctx, &targetUser) if err == nil { - oauth.InvalidateCachedUser(ctx, param.ID) + auth.InvalidateCachedUser(ctx, param.ID) if needRevokeTokens { for _, token := range tokens { - oauth.InvalidateCachedToken(ctx, token.TokenHash) + auth.InvalidateCachedToken(ctx, token.TokenHash) } } } diff --git a/plugins/domain/admin/middlewares.go b/plugins/domain/admin/middlewares.go index c4af729b..98be6759 100644 --- a/plugins/domain/admin/middlewares.go +++ b/plugins/domain/admin/middlewares.go @@ -4,11 +4,11 @@ package admin import ( - "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/shared/response" "github.com/Rain-kl/Wavelet/pkg/logger" otel_trace "github.com/Rain-kl/Wavelet/pkg/trace" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" "github.com/gin-gonic/gin" ) @@ -18,15 +18,15 @@ func LoginAdminRequired() gin.HandlerFunc { ctx, span := otel_trace.Start(c.Request.Context(), "LoginAdminRequired") defer span.End() - user, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) + user, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey) if user == nil { response.AbortNotFound(c, AdminRequired) return } // 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限 - if tokenAuth, _ := oauth.GetFromContext[bool](c, oauth.TokenAuthKey); tokenAuth { - tokenAdmin, _ := oauth.GetFromContext[bool](c, oauth.TokenAdminKey) + if tokenAuth, _ := auth.GetFromContext[bool](c, auth.TokenAuthKey); tokenAuth { + tokenAdmin, _ := auth.GetFromContext[bool](c, auth.TokenAdminKey) if !tokenAdmin { response.AbortNotFound(c, TokenAdminRequired) return diff --git a/plugins/domain/admin/plugin.go b/plugins/domain/admin/plugin.go index 681e6650..e5b0deea 100644 --- a/plugins/domain/admin/plugin.go +++ b/plugins/domain/admin/plugin.go @@ -1,3 +1,6 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + // Package admin provides the system management console, diagnostics, audit logging, and configuration hot-reloading domain plugin for Cordis. package admin @@ -6,20 +9,7 @@ import ( "github.com/Rain-kl/Wavelet/core" "github.com/Rain-kl/Wavelet/core/extpoints" - "github.com/Rain-kl/Wavelet/internal/apps/admin" - admin_auth_source "github.com/Rain-kl/Wavelet/internal/apps/admin/auth_source" - admin_cache "github.com/Rain-kl/Wavelet/internal/apps/admin/cache" - admin_db_manage "github.com/Rain-kl/Wavelet/internal/apps/admin/db_manage" - admin_logs "github.com/Rain-kl/Wavelet/internal/apps/admin/logs" - admin_push "github.com/Rain-kl/Wavelet/internal/apps/admin/push" - admin_status "github.com/Rain-kl/Wavelet/internal/apps/admin/status" - "github.com/Rain-kl/Wavelet/internal/apps/admin/system_config" - admin_task "github.com/Rain-kl/Wavelet/internal/apps/admin/task" - admin_template "github.com/Rain-kl/Wavelet/internal/apps/admin/template" - admin_updater "github.com/Rain-kl/Wavelet/internal/apps/admin/updater" - admin_user "github.com/Rain-kl/Wavelet/internal/apps/admin/user" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/apps/upload" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" "github.com/hibiken/asynq" ) @@ -58,153 +48,115 @@ func (p *Plugin) Manifest() core.Manifest { // Apply registers admin routes, tasks, schedules, and settings into the Context. func (p *Plugin) Apply(ctx *core.Context) error { // 1. Register Admin HTTP Routes - adminRouter := ctx.Router().Group("/api/v1/admin", oauth.LoginRequired(), admin.LoginAdminRequired()) + adminRouter := ctx.Router().Group("/api/v1/admin", auth.LoginRequired(), LoginAdminRequired()) { // Status & Diagnostics - adminRouter.GET("/status", admin_status.GetSystemStatus) - adminRouter.GET("/status/log-database", admin_status.GetLogDatabaseStatus) - adminRouter.GET("/db-info", admin_status.GetDatabaseInfo) - adminRouter.GET("/db-export", admin_status.ExportDatabase) + adminRouter.GET("/status", GetSystemStatus) + adminRouter.GET("/status/log-database", GetLogDatabaseStatus) + adminRouter.GET("/db-info", GetDatabaseInfo) + adminRouter.GET("/db-export", ExportDatabase) // DB Management dbGroup := adminRouter.Group("/db-manage") { - dbGroup.GET("/overview", admin_db_manage.GetDBOverview) - dbGroup.GET("/tables", admin_db_manage.ListDBTables) - dbGroup.GET("/table-data", admin_db_manage.GetDBTableData) - dbGroup.POST("/query", admin_db_manage.ExecuteSQL) + dbGroup.GET("/overview", GetDBOverview) + dbGroup.GET("/tables", ListDBTables) + dbGroup.GET("/table-data", GetDBTableData) + dbGroup.POST("/query", ExecuteSQL) } // Cache Management cacheGroup := adminRouter.Group("/cache") { - cacheGroup.GET("/status", admin_cache.GetCacheStatus) - cacheGroup.POST("/config", admin_cache.UpdateCacheConfig) - cacheGroup.POST("/clear", admin_cache.ClearCache) + cacheGroup.GET("/status", GetCacheStatus) + cacheGroup.POST("/config", UpdateCacheConfig) + cacheGroup.POST("/clear", ClearCache) } // Updater updateGroup := adminRouter.Group("/update") { - updateGroup.GET("", admin_updater.GetUpdateStatus) - updateGroup.POST("/apply", admin_updater.ApplyUpdate) + updateGroup.GET("", GetUpdateStatus) + updateGroup.POST("/apply", ApplyUpdate) } // Logs logsGroup := adminRouter.Group("/logs") { - logsGroup.GET("", admin_logs.GetLogs) - logsGroup.GET("/access", admin_logs.GetAccessLogs) - logsGroup.GET("/analytics", admin_logs.GetLogsAnalytics) - logsGroup.GET("/ws", admin_logs.HandleLogWebSocket) + logsGroup.GET("", GetLogs) + logsGroup.GET("/access", GetAccessLogs) + logsGroup.GET("/analytics", GetLogsAnalytics) + logsGroup.GET("/ws", HandleLogWebSocket) } // Users usersGroup := adminRouter.Group("/users") { - usersGroup.GET("", admin_user.ListUsers) - usersGroup.POST("", admin_user.CreateUser) - usersGroup.GET("/:id", admin_user.GetUser) - usersGroup.PUT("/:id/status", admin_user.UpdateUserStatus) - usersGroup.PUT("/:id", admin_user.UpdateUser) - usersGroup.DELETE("/:id", admin_user.DeleteUser) + usersGroup.GET("", ListUsers) + usersGroup.POST("", CreateUser) + usersGroup.GET("/:id", GetUser) + usersGroup.PUT("/:id/status", UpdateUserStatus) + usersGroup.PUT("/:id", UpdateUser) + usersGroup.DELETE("/:id", DeleteUser) } // Auth Sources authSourcesGroup := adminRouter.Group("/auth-sources") { - authSourcesGroup.GET("", admin_auth_source.ListAuthSources) - authSourcesGroup.POST("", admin_auth_source.CreateAuthSource) - authSourcesGroup.PUT("/:id", admin_auth_source.UpdateAuthSource) - authSourcesGroup.PUT("/:id/toggle", admin_auth_source.ToggleAuthSource) - authSourcesGroup.DELETE("/:id", admin_auth_source.DeleteAuthSource) + authSourcesGroup.GET("", ListAuthSources) + authSourcesGroup.POST("", CreateAuthSource) + authSourcesGroup.PUT("/:id", UpdateAuthSource) + authSourcesGroup.PUT("/:id/toggle", ToggleAuthSource) + authSourcesGroup.DELETE("/:id", DeleteAuthSource) } // System Configs configGroup := adminRouter.Group("/system-configs") { - configGroup.GET("", system_config.ListSystemConfigs) - configGroup.POST("", system_config.CreateSystemConfig) - configGroup.POST("/smtp/test", system_config.TestSMTP) + configGroup.GET("", ListSystemConfigs) + configGroup.POST("", CreateSystemConfig) + configGroup.POST("/smtp/test", TestSMTP) keyGroup := configGroup.Group("/:key") { - keyGroup.GET("", system_config.GetSystemConfig) - keyGroup.PUT("", system_config.UpdateSystemConfig) + keyGroup.GET("", GetSystemConfig) + keyGroup.PUT("", UpdateSystemConfig) } } // Templates templateGroup := adminRouter.Group("/templates") { - templateGroup.GET("", admin_template.ListTemplates) - templateGroup.POST("", admin_template.CreateTemplate) + templateGroup.GET("", ListTemplates) + templateGroup.POST("", CreateTemplate) keyGroup := templateGroup.Group("/:key") { - keyGroup.GET("", admin_template.GetTemplate) - keyGroup.PUT("", admin_template.UpdateTemplate) - keyGroup.DELETE("", admin_template.DeleteTemplate) + keyGroup.GET("", GetTemplate) + keyGroup.PUT("", UpdateTemplate) + keyGroup.DELETE("", DeleteTemplate) } } - // Uploads Management - uploadGroup := adminRouter.Group("/uploads") - { - uploadGroup.GET("", upload.ListFiles) - uploadGroup.GET("/stats", upload.GetFileStats) - uploadGroup.DELETE("/:id", upload.DeleteFile) - uploadGroup.GET("/download/:id", upload.DownloadFile) - uploadGroup.POST("/download/batch", upload.BatchDownloadFiles) - uploadGroup.GET("/types", upload.GetDistinctUploadTypes) - } - // Tasks taskGroup := adminRouter.Group("/tasks") { - taskGroup.GET("/types", admin_task.ListTaskTypes) - taskGroup.POST("/dispatch", admin_task.DispatchTask) + taskGroup.GET("/types", ListTaskTypes) + taskGroup.POST("/dispatch", DispatchTask) executions := taskGroup.Group("/executions") { - executions.GET("", admin_task.ListTaskExecutions) - executions.GET("/:id", admin_task.GetTaskExecution) - executions.POST("/:id/retry", admin_task.RetryTask) + executions.GET("", ListTaskExecutions) + executions.GET("/:id", GetTaskExecution) + executions.POST("/:id/retry", RetryTask) } schedules := taskGroup.Group("/schedules") { - schedules.GET("", admin_task.ListSchedules) - schedules.POST("", admin_task.CreateSchedule) - schedules.PUT("/:id", admin_task.UpdateSchedule) - schedules.DELETE("/:id", admin_task.DeleteSchedule) - } - } - - // Push & Notifications - pushGroup := adminRouter.Group("/push") - { - events := pushGroup.Group("/events") - { - events.GET("", admin_push.ListEvents) - events.GET("/builtin", admin_push.ListBuiltInEvents) - events.POST("", admin_push.CreateEvent) - events.PUT("/:id", admin_push.UpdateEvent) - events.DELETE("/:id", admin_push.DeleteEvent) - events.POST("/:id/toggle", admin_push.ToggleEvent) - } - - pushGroup.GET("/histories", admin_push.ListHistories) - pushGroup.POST("/test", admin_push.TestPush) - - channels := pushGroup.Group("/channels") - { - channels.GET("/definitions", admin_push.ListChannelDefinitions) - channels.GET("", admin_push.ListChannels) - channels.POST("", admin_push.CreateChannel) - channels.PUT("/:id", admin_push.UpdateChannel) - channels.DELETE("/:id", admin_push.DeleteChannel) - channels.POST("/test", admin_push.TestChannel) + schedules.GET("", ListSchedules) + schedules.POST("", CreateSchedule) + schedules.PUT("/:id", UpdateSchedule) + schedules.DELETE("/:id", DeleteSchedule) } } } diff --git a/internal/apps/cap/errs.go b/plugins/domain/cap/errs.go similarity index 100% rename from internal/apps/cap/errs.go rename to plugins/domain/cap/errs.go diff --git a/internal/apps/cap/routers.go b/plugins/domain/cap/handlers.go similarity index 100% rename from internal/apps/cap/routers.go rename to plugins/domain/cap/handlers.go diff --git a/internal/apps/cap/manager.go b/plugins/domain/cap/manager.go similarity index 100% rename from internal/apps/cap/manager.go rename to plugins/domain/cap/manager.go diff --git a/internal/apps/cap/middleware.go b/plugins/domain/cap/middleware.go similarity index 100% rename from internal/apps/cap/middleware.go rename to plugins/domain/cap/middleware.go diff --git a/plugins/domain/cap/plugin.go b/plugins/domain/cap/plugin.go new file mode 100644 index 00000000..dc733e90 --- /dev/null +++ b/plugins/domain/cap/plugin.go @@ -0,0 +1,62 @@ +// 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 ( + "github.com/Rain-kl/Wavelet/core" + "github.com/Rain-kl/Wavelet/core/extpoints" +) + +// 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" +} + +// 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", + } +} + +// Apply registers the cap routes and settings into the Context. +func (p *Plugin) Apply(ctx *core.Context) error { + // Register HTTP Routes + capGroup := ctx.Router().Group("/api/v1/cap") + { + capGroup.GET("/challenge", Challenge) + capGroup.POST("/challenge", Challenge) + capGroup.POST("/redeem", 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 +} diff --git a/internal/apps/cap/runtime_settings.go b/plugins/domain/cap/runtime_settings.go similarity index 100% rename from internal/apps/cap/runtime_settings.go rename to plugins/domain/cap/runtime_settings.go diff --git a/plugins/domain/domain_test.go b/plugins/domain/domain_test.go index 9d466462..96edfdda 100644 --- a/plugins/domain/domain_test.go +++ b/plugins/domain/domain_test.go @@ -64,7 +64,7 @@ func (m *mockOAuthProvider) Name() string { } func (m *mockOAuthProvider) GetAuthURL(state string) string { - return "https://oauth.example.com/auth?state=" + state + return "https://auth.example.com/auth?state=" + state } func (m *mockOAuthProvider) ExchangeCode(ctx context.Context, code string) (*contracts.OAuthUserInfoDTO, error) { @@ -349,7 +349,7 @@ func TestAdminPlugin(t *testing.T) { // 1. Admin Routes routes := ctx.Router().Routes() - var hasStatus, hasDBOverview, hasUsers, hasTasks, hasPushEvents bool + var hasStatus, hasDBOverview, hasUsers, hasTasks, hasConfigs bool for _, r := range routes { if r.Path == "/api/v1/admin/status" { hasStatus = true @@ -363,15 +363,15 @@ func TestAdminPlugin(t *testing.T) { if r.Path == "/api/v1/admin/tasks/types" { hasTasks = true } - if r.Path == "/api/v1/admin/push/events" { - hasPushEvents = true + if r.Path == "/api/v1/admin/system-configs" { + hasConfigs = true } } assert.True(t, hasStatus) assert.True(t, hasDBOverview) assert.True(t, hasUsers) assert.True(t, hasTasks) - assert.True(t, hasPushEvents) + assert.True(t, hasConfigs) // 2. Task & Schedule _, ok := ctx.Tasks().Get("admin:system_cleanup") diff --git a/plugins/domain/message_gateway/handlers.go b/plugins/domain/message_gateway/handlers.go index 096513bb..346cba17 100644 --- a/plugins/domain/message_gateway/handlers.go +++ b/plugins/domain/message_gateway/handlers.go @@ -8,14 +8,14 @@ import ( "net/http" "strconv" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/shared/response" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" "github.com/gin-gonic/gin" ) func currentUser(c *gin.Context) (*model.User, bool) { - return oauth.GetFromContext[*model.User](c, oauth.UserObjKey) + return auth.GetFromContext[*model.User](c, auth.UserObjKey) } // ListChannels lists enabled channels a user can bind. @@ -137,7 +137,7 @@ func UnbindBinding(c *gin.Context) { // RegisterUserRoutes mounts user-facing message gateway endpoints. func RegisterUserRoutes(r *gin.RouterGroup) { - mg := r.Group("/message-gateway", oauth.LoginRequired()) + mg := r.Group("/message-gateway", auth.LoginRequired()) { mg.GET("/channels", ListChannels) mg.GET("/bindings", ListBindings) diff --git a/plugins/domain/message_gateway/plugin.go b/plugins/domain/message_gateway/plugin.go index 74916359..db1f1017 100644 --- a/plugins/domain/message_gateway/plugin.go +++ b/plugins/domain/message_gateway/plugin.go @@ -10,8 +10,8 @@ import ( "github.com/Rain-kl/Wavelet/core" "github.com/Rain-kl/Wavelet/core/extpoints" - "github.com/Rain-kl/Wavelet/internal/apps/admin" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" + "github.com/Rain-kl/Wavelet/plugins/domain/admin" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" "github.com/hibiken/asynq" ) @@ -75,7 +75,7 @@ func (p *Plugin) Apply(ctx *core.Context) error { ctx.Migrations().Register("message_gateway", mgMigrations) // 2. Register User HTTP Routes - mgGroup := ctx.Router().Group("/api/v1/message-gateway", oauth.LoginRequired()) + mgGroup := ctx.Router().Group("/api/v1/message-gateway", auth.LoginRequired()) { mgGroup.GET("/channels", ListChannels) mgGroup.GET("/bindings", ListBindings) @@ -84,7 +84,7 @@ func (p *Plugin) Apply(ctx *core.Context) error { } // 3. Register Admin Message Gateway HTTP Routes - adminMgGroup := ctx.Router().Group("/api/v1/admin/message-gateway", oauth.LoginRequired(), admin.LoginAdminRequired()) + adminMgGroup := ctx.Router().Group("/api/v1/admin/message-gateway", auth.LoginRequired(), admin.LoginAdminRequired()) { adminMgGroup.GET("/channels/definitions", ListAdminChannelDefinitions) adminMgGroup.GET("/channels", ListAdminChannels) @@ -95,7 +95,7 @@ func (p *Plugin) Apply(ctx *core.Context) error { } // 4. Register Admin Push HTTP Routes - adminPushGroup := ctx.Router().Group("/api/v1/admin/push", oauth.LoginRequired(), admin.LoginAdminRequired()) + adminPushGroup := ctx.Router().Group("/api/v1/admin/push", auth.LoginRequired(), admin.LoginAdminRequired()) { events := adminPushGroup.Group("/events") { diff --git a/plugins/domain/risk_control/middleware.go b/plugins/domain/risk_control/middleware.go index b3769633..9ae36634 100644 --- a/plugins/domain/risk_control/middleware.go +++ b/plugins/domain/risk_control/middleware.go @@ -9,15 +9,18 @@ import ( "net/http" "time" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/infra/config" "github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model/analytics" "github.com/Rain-kl/Wavelet/internal/shared/response" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" "github.com/gin-gonic/gin" ) +// Middleware is an alias for RiskControlMiddleware. +var Middleware = RiskControlMiddleware + // RiskControlMiddleware 全局日志采集中间件 func RiskControlMiddleware() gin.HandlerFunc { return func(c *gin.Context) { @@ -39,7 +42,7 @@ func RiskControlMiddleware() gin.HandlerFunc { c.Next() // 3. 后置身份检查:仅记录通过认证的请求 - userObj, exists := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) + userObj, exists := auth.GetFromContext[*model.User](c, auth.UserObjKey) if !exists || userObj == nil { return } diff --git a/plugins/domain/risk_control/middleware_test.go b/plugins/domain/risk_control/middleware_test.go index eb41af2f..996f7170 100644 --- a/plugins/domain/risk_control/middleware_test.go +++ b/plugins/domain/risk_control/middleware_test.go @@ -12,12 +12,12 @@ import ( "testing" "time" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/infra/config" "github.com/Rain-kl/Wavelet/internal/infra/persistence/batchwriter" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model/analytics" "github.com/Rain-kl/Wavelet/internal/testhelper" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" "github.com/Rain-kl/Wavelet/plugins/domain/risk_control" "github.com/gin-gonic/gin" "github.com/stretchr/testify/assert" @@ -99,7 +99,7 @@ func TestRiskControlMiddleware(t *testing.T) { r := gin.New() r.Use(func(c *gin.Context) { user := &model.User{ID: 12345} - oauth.SetToContext(c, oauth.UserObjKey, user) + auth.SetToContext(c, auth.UserObjKey, user) c.Next() }) r.Use(risk_control.RiskControlMiddleware()) diff --git a/plugins/domain/system/plugin.go b/plugins/domain/system/plugin.go new file mode 100644 index 00000000..2ed6a8f2 --- /dev/null +++ b/plugins/domain/system/plugin.go @@ -0,0 +1,71 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package system provides core system health probes, public config endpoints, and static frontend assets dispatch plugin for Cordis. +package system + +import ( + "net/http" + + "github.com/Rain-kl/Wavelet/core" + "github.com/Rain-kl/Wavelet/internal/infra/config" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/shared/response" + "github.com/gin-gonic/gin" +) + +// Plugin implements core.Plugin to provide system-level basic routes. +type Plugin struct{} + +// New creates a new system domain plugin. +func New() *Plugin { + return &Plugin{} +} + +// Name returns the unique identifier for the system domain plugin. +func (p *Plugin) Name() string { + return "system" +} + +// Manifest returns the plugin metadata. +func (p *Plugin) Manifest() core.Manifest { + return core.Manifest{ + Name: "system", + Version: "1.0.0", + Description: "System health check, public config, and assets domain plugin", + Author: "Wavelet Team", + } +} + +// Apply registers system routes. +func (p *Plugin) Apply(ctx *core.Context) error { + // 1. Health check + ctx.Router().GET("/healthz", func(c *gin.Context) { + c.JSON(http.StatusOK, gin.H{"status": "ok"}) + }) + ctx.Router().GET("/api/healthz", func(c *gin.Context) { + c.JSON(http.StatusOK, gin.H{"status": "ok"}) + }) + + // 2. Public config + ctx.Router().GET("/api/v1/config/public", func(c *gin.Context) { + configs, err := repository.ListVisibleSystemConfigs(c.Request.Context()) + if err != nil { + response.AbortInternal(c, "获取公开配置失败") + return + } + c.JSON(http.StatusOK, response.OK(gin.H{ + "configs": configs, + "app": gin.H{ + "name": config.Config.App.AppName, + }, + })) + }) + + // 3. Custom injection + ctx.Router().GET("/custom", func(c *gin.Context) { + c.JSON(http.StatusOK, response.OK(gin.H{"custom": true})) + }) + + return nil +} diff --git a/internal/apps/upload/cache/access_cache.go b/plugins/domain/upload/cache/access_cache.go similarity index 96% rename from internal/apps/upload/cache/access_cache.go rename to plugins/domain/upload/cache/access_cache.go index abf46d00..6a3d5387 100644 --- a/internal/apps/upload/cache/access_cache.go +++ b/plugins/domain/upload/cache/access_cache.go @@ -11,13 +11,13 @@ import ( "sync" "time" - "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" - uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" "github.com/Rain-kl/Wavelet/internal/infra/objectstore" "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/util" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" + uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage" ) const fileAccessInvalidationChannel = "upload:file_access_invalidation" diff --git a/internal/apps/upload/cache/access_cache_test.go b/plugins/domain/upload/cache/access_cache_test.go similarity index 95% rename from internal/apps/upload/cache/access_cache_test.go rename to plugins/domain/upload/cache/access_cache_test.go index 363650b7..7b3d3a1d 100644 --- a/internal/apps/upload/cache/access_cache_test.go +++ b/plugins/domain/upload/cache/access_cache_test.go @@ -8,12 +8,12 @@ import ( "testing" "time" - "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" - uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/testhelper" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" + uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage" ) func TestLoadMigrationAccessStateCachesResult(t *testing.T) { diff --git a/internal/apps/upload/cache/meta_cache.go b/plugins/domain/upload/cache/meta_cache.go similarity index 100% rename from internal/apps/upload/cache/meta_cache.go rename to plugins/domain/upload/cache/meta_cache.go diff --git a/internal/apps/upload/cache/meta_cache_test.go b/plugins/domain/upload/cache/meta_cache_test.go similarity index 100% rename from internal/apps/upload/cache/meta_cache_test.go rename to plugins/domain/upload/cache/meta_cache_test.go diff --git a/internal/apps/upload/errs.go b/plugins/domain/upload/errs.go similarity index 96% rename from internal/apps/upload/errs.go rename to plugins/domain/upload/errs.go index bc6811d6..a980dc26 100644 --- a/internal/apps/upload/errs.go +++ b/plugins/domain/upload/errs.go @@ -5,7 +5,7 @@ // Package upload 提供文件上传与下载功能 package upload -import "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" +import "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" // 文件管理常量 const ( diff --git a/internal/apps/upload/exports.go b/plugins/domain/upload/exports.go similarity index 89% rename from internal/apps/upload/exports.go rename to plugins/domain/upload/exports.go index d2c4273f..7d37dc33 100644 --- a/internal/apps/upload/exports.go +++ b/plugins/domain/upload/exports.go @@ -4,14 +4,14 @@ package upload import ( - "github.com/Rain-kl/Wavelet/internal/apps/upload/cache" - "github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv" - "github.com/Rain-kl/Wavelet/internal/apps/upload/handler" - "github.com/Rain-kl/Wavelet/internal/apps/upload/ingest" - uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" - uploadtask "github.com/Rain-kl/Wavelet/internal/apps/upload/task" - "github.com/Rain-kl/Wavelet/internal/apps/upload/util" "github.com/Rain-kl/Wavelet/internal/infra/task" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/cache" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/filesrv" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/handler" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/ingest" + uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats" + uploadtask "github.com/Rain-kl/Wavelet/plugins/domain/upload/task" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/util" ) // HTTP handlers diff --git a/internal/apps/upload/filesrv/file_server.go b/plugins/domain/upload/filesrv/file_server.go similarity index 93% rename from internal/apps/upload/filesrv/file_server.go rename to plugins/domain/upload/filesrv/file_server.go index 04d0ac34..01a314d3 100644 --- a/internal/apps/upload/filesrv/file_server.go +++ b/plugins/domain/upload/filesrv/file_server.go @@ -15,15 +15,15 @@ import ( "strconv" "strings" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/apps/upload/cache" - "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" - uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" - "github.com/Rain-kl/Wavelet/internal/apps/upload/util" "github.com/Rain-kl/Wavelet/internal/infra/diskcache" "github.com/Rain-kl/Wavelet/internal/model" appshared "github.com/Rain-kl/Wavelet/internal/shared" "github.com/Rain-kl/Wavelet/internal/shared/response" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/cache" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" + uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/util" "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/gin-gonic/gin" @@ -279,10 +279,10 @@ func getOriginalFileBytes(ctx context.Context, upload *model.Upload) ([]byte, er func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error { var currUser *model.User var err error - if u, ok := oauth.GetFromContext[*model.User](c, oauth.UserObjKey); ok && u != nil { + if u, ok := auth.GetFromContext[*model.User](c, auth.UserObjKey); ok && u != nil { currUser = u } else { - currUser, err = oauth.GetUserFromRequest(c) + currUser, err = auth.GetUserFromRequest(c) if err != nil { return err } @@ -303,8 +303,8 @@ func CheckFileAccessPermission(c *gin.Context, upload *model.Upload) error { } if !cache.IsFilePublic(c.Request.Context(), upload.Type) { - if _, ok := oauth.GetFromContext[*model.User](c, oauth.UserObjKey); !ok { - if _, err := oauth.GetUserFromRequest(c); err != nil { + if _, ok := auth.GetFromContext[*model.User](c, auth.UserObjKey); !ok { + if _, err := auth.GetUserFromRequest(c); err != nil { return err } } diff --git a/internal/apps/upload/filesrv/file_server_test.go b/plugins/domain/upload/filesrv/file_server_test.go similarity index 98% rename from internal/apps/upload/filesrv/file_server_test.go rename to plugins/domain/upload/filesrv/file_server_test.go index 050b9993..8c04d5fc 100644 --- a/internal/apps/upload/filesrv/file_server_test.go +++ b/plugins/domain/upload/filesrv/file_server_test.go @@ -17,9 +17,6 @@ import ( "path/filepath" "testing" - "github.com/Rain-kl/Wavelet/internal/apps/upload/cache" - "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" - "github.com/Rain-kl/Wavelet/internal/apps/upload/util" "github.com/Rain-kl/Wavelet/internal/infra/diskcache" "github.com/Rain-kl/Wavelet/internal/infra/objectstore" "github.com/Rain-kl/Wavelet/internal/infra/persistence" @@ -28,6 +25,9 @@ import ( appshared "github.com/Rain-kl/Wavelet/internal/shared" "github.com/Rain-kl/Wavelet/internal/shared/response" "github.com/Rain-kl/Wavelet/internal/testhelper" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/cache" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/util" "github.com/gin-contrib/sessions" "github.com/gin-contrib/sessions/cookie" "github.com/gin-gonic/gin" diff --git a/internal/apps/upload/handler/file_management.go b/plugins/domain/upload/handler/file_management.go similarity index 95% rename from internal/apps/upload/handler/file_management.go rename to plugins/domain/upload/handler/file_management.go index dee92a39..c2011bc3 100644 --- a/internal/apps/upload/handler/file_management.go +++ b/plugins/domain/upload/handler/file_management.go @@ -7,13 +7,13 @@ import ( "net/http" "strconv" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/apps/upload/ingest" - "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" - uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/shared/response" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/ingest" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" + uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage" "github.com/gin-gonic/gin" ) @@ -171,7 +171,7 @@ type listMyFilesResponse struct { // @Failure 401 {object} response.Any "未登录" // @Router /api/v1/upload/my [get] func ListMyFiles(c *gin.Context) { - currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey) ctx := c.Request.Context() var req listMyFilesRequest @@ -218,7 +218,7 @@ func ListMyFiles(c *gin.Context) { // @Failure 404 {object} response.Any "文件不存在" // @Router /api/v1/upload/{id} [delete] func DeleteMyFile(c *gin.Context) { - currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey) ctx := c.Request.Context() if uploadstorage.ReadOnly(ctx) { response.AbortConflict(c, shared.ErrStorageReadOnly) @@ -265,7 +265,7 @@ type updateMyFileRequest struct { // @Failure 404 {object} response.Any "文件不存在" // @Router /api/v1/upload/{id} [put] func UpdateMyFile(c *gin.Context) { - currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey) ctx := c.Request.Context() if uploadstorage.ReadOnly(ctx) { response.AbortConflict(c, shared.ErrStorageReadOnly) diff --git a/internal/apps/upload/handler/file_management_test.go b/plugins/domain/upload/handler/file_management_test.go similarity index 100% rename from internal/apps/upload/handler/file_management_test.go rename to plugins/domain/upload/handler/file_management_test.go diff --git a/internal/apps/upload/handler/logics.go b/plugins/domain/upload/handler/logics.go similarity index 97% rename from internal/apps/upload/handler/logics.go rename to plugins/domain/upload/handler/logics.go index c966b0fc..fb48e38a 100644 --- a/internal/apps/upload/handler/logics.go +++ b/plugins/domain/upload/handler/logics.go @@ -8,9 +8,9 @@ import ( "errors" "sort" - "github.com/Rain-kl/Wavelet/internal/apps/upload/ingest" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/ingest" "gorm.io/gorm" ) diff --git a/internal/apps/upload/handler/routers.go b/plugins/domain/upload/handler/routers.go similarity index 95% rename from internal/apps/upload/handler/routers.go rename to plugins/domain/upload/handler/routers.go index 9ab81fbb..4cfc8021 100644 --- a/internal/apps/upload/handler/routers.go +++ b/plugins/domain/upload/handler/routers.go @@ -22,16 +22,16 @@ import ( "strconv" "strings" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv" - "github.com/Rain-kl/Wavelet/internal/apps/upload/ingest" - "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" - uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" - "github.com/Rain-kl/Wavelet/internal/apps/upload/util" "github.com/Rain-kl/Wavelet/internal/model" appshared "github.com/Rain-kl/Wavelet/internal/shared" "github.com/Rain-kl/Wavelet/internal/shared/response" "github.com/Rain-kl/Wavelet/pkg/logger" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/filesrv" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/ingest" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" + uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/util" "github.com/gin-gonic/gin" "gorm.io/gorm" ) @@ -63,7 +63,7 @@ func UploadFile(c *gin.Context) { c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, shared.MaxUploadSize) - currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) + currUser, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey) ctx := c.Request.Context() header, err := c.FormFile("file") diff --git a/internal/apps/upload/handler/routers_test.go b/plugins/domain/upload/handler/routers_test.go similarity index 99% rename from internal/apps/upload/handler/routers_test.go rename to plugins/domain/upload/handler/routers_test.go index 10b2ddd0..5a3d49d5 100644 --- a/internal/apps/upload/handler/routers_test.go +++ b/plugins/domain/upload/handler/routers_test.go @@ -19,15 +19,15 @@ import ( "testing" "time" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" - uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" "github.com/Rain-kl/Wavelet/internal/infra/objectstore" "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/shared/response" "github.com/Rain-kl/Wavelet/internal/testhelper" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" + uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats" "github.com/gin-gonic/gin" ) @@ -43,7 +43,7 @@ func setupTestRouter(authUser *model.User) *gin.Engine { authMiddleware := func(c *gin.Context) { if authUser != nil { - oauth.SetToContext(c, oauth.UserObjKey, authUser) + auth.SetToContext(c, auth.UserObjKey, authUser) } c.Next() } diff --git a/internal/apps/upload/handler/stats.go b/plugins/domain/upload/handler/stats.go similarity index 98% rename from internal/apps/upload/handler/stats.go rename to plugins/domain/upload/handler/stats.go index e14536bc..7c23bfac 100644 --- a/internal/apps/upload/handler/stats.go +++ b/plugins/domain/upload/handler/stats.go @@ -7,9 +7,9 @@ import ( "net/http" "time" - "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/shared/response" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" "github.com/gin-gonic/gin" ) diff --git a/internal/apps/upload/ingest/errors.go b/plugins/domain/upload/ingest/errors.go similarity index 86% rename from internal/apps/upload/ingest/errors.go rename to plugins/domain/upload/ingest/errors.go index 6e501977..1d33e945 100644 --- a/internal/apps/upload/ingest/errors.go +++ b/plugins/domain/upload/ingest/errors.go @@ -6,7 +6,7 @@ package ingest import ( "errors" - "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" ) // ErrForbidden indicates the caller is not allowed to mutate the upload record. diff --git a/internal/apps/upload/ingest/helpers.go b/plugins/domain/upload/ingest/helpers.go similarity index 94% rename from internal/apps/upload/ingest/helpers.go rename to plugins/domain/upload/ingest/helpers.go index a52ba1f5..b54b8173 100644 --- a/internal/apps/upload/ingest/helpers.go +++ b/plugins/domain/upload/ingest/helpers.go @@ -11,16 +11,16 @@ import ( "strings" "time" - uploadcache "github.com/Rain-kl/Wavelet/internal/apps/upload/cache" - "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" - uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" - uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" "github.com/Rain-kl/Wavelet/internal/infra/objectstore" "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/logger" + uploadcache "github.com/Rain-kl/Wavelet/plugins/domain/upload/cache" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" + uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats" + uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage" "gorm.io/gorm" ) diff --git a/internal/apps/upload/ingest/ingest.go b/plugins/domain/upload/ingest/ingest.go similarity index 100% rename from internal/apps/upload/ingest/ingest.go rename to plugins/domain/upload/ingest/ingest.go diff --git a/internal/apps/upload/ingest/ingest_test.go b/plugins/domain/upload/ingest/ingest_test.go similarity index 100% rename from internal/apps/upload/ingest/ingest_test.go rename to plugins/domain/upload/ingest/ingest_test.go diff --git a/internal/apps/upload/ingest/remove.go b/plugins/domain/upload/ingest/remove.go similarity index 92% rename from internal/apps/upload/ingest/remove.go rename to plugins/domain/upload/ingest/remove.go index 001f1889..eb298127 100644 --- a/internal/apps/upload/ingest/remove.go +++ b/plugins/domain/upload/ingest/remove.go @@ -6,11 +6,11 @@ package ingest import ( "context" - uploadcache "github.com/Rain-kl/Wavelet/internal/apps/upload/cache" - uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" + uploadcache "github.com/Rain-kl/Wavelet/plugins/domain/upload/cache" + uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats" "gorm.io/gorm" ) diff --git a/internal/apps/upload/ingest/types.go b/plugins/domain/upload/ingest/types.go similarity index 100% rename from internal/apps/upload/ingest/types.go rename to plugins/domain/upload/ingest/types.go diff --git a/plugins/domain/upload/plugin.go b/plugins/domain/upload/plugin.go new file mode 100644 index 00000000..18e070ca --- /dev/null +++ b/plugins/domain/upload/plugin.go @@ -0,0 +1,117 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package upload provides file uploading, storage abstraction, image transcoding, and caching domain plugin for Cordis. +package upload + +import ( + "context" + + "github.com/Rain-kl/Wavelet/core" + "github.com/Rain-kl/Wavelet/core/extpoints" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/filesrv" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/handler" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/task" + "github.com/hibiken/asynq" +) + +// Plugin implements core.Plugin to provide file upload and media serving domain services. +type Plugin struct{} + +// New creates a new upload domain plugin. +func New() *Plugin { + return &Plugin{} +} + +// Name returns the unique identifier for the upload domain plugin. +func (p *Plugin) Name() string { + return "upload" +} + +// Manifest returns the plugin metadata. +func (p *Plugin) Manifest() core.Manifest { + return core.Manifest{ + Name: "upload", + Version: "1.0.0", + Description: "File upload, secure delivery, image transcoding, and storage management domain plugin", + Author: "Wavelet Team", + } +} + +// Apply registers upload routes, tasks, and settings into the Context. +func (p *Plugin) Apply(ctx *core.Context) error { + // 1. Register File Server Routes + ctx.Router().GET("/f/:id", filesrv.ServeFileByID) + + // 2. Register User/Admin Upload HTTP Routes + uploadGroup := ctx.Router().Group("/api/v1/upload", auth.LoginRequired()) + { + uploadGroup.POST("", handler.UploadFile) + uploadGroup.GET("", handler.ListFiles) + uploadGroup.DELETE("/:id", handler.DeleteFile) + uploadGroup.POST("/batch-download", handler.BatchDownloadFiles) + } + + adminUploadGroup := ctx.Router().Group("/api/v1/admin/uploads", auth.LoginRequired()) + { + adminUploadGroup.GET("", handler.ListFiles) + adminUploadGroup.GET("/stats", handler.GetFileStats) + adminUploadGroup.DELETE("/:id", handler.DeleteFile) + adminUploadGroup.GET("/download/:id", handler.DownloadFile) + adminUploadGroup.POST("/download/batch", handler.BatchDownloadFiles) + adminUploadGroup.GET("/types", handler.GetDistinctUploadTypes) + } + + const ( + defaultCleanupRetry = 3 + defaultStatsRetry = 2 + defaultSingleRetry = 1 + ) + + // 3. Register Asynq tasks + cleanupHandler := &task.SystemCleanupHandler{} + ctx.Task().Register(task.SystemCleanupTask, func(c context.Context, t *asynq.Task) error { + _, err := cleanupHandler.Execute(c, t.Payload()) + return err + }, extpoints.WithTaskRetry(defaultCleanupRetry)) + + rebuildStatsHandler := &task.RebuildUploadStatsHandler{} + ctx.Task().Register(task.RebuildUploadStatsTask, func(c context.Context, t *asynq.Task) error { + _, err := rebuildStatsHandler.Execute(c, t.Payload()) + return err + }, extpoints.WithTaskRetry(defaultStatsRetry)) + + migrationHandler := &task.MigrationHandler{} + ctx.Task().Register(task.StorageMigrationTask, func(c context.Context, t *asynq.Task) error { + _, err := migrationHandler.Execute(c, t.Payload()) + return err + }, extpoints.WithTaskRetry(defaultSingleRetry)) + + warmHandler := &task.WarmImageCacheHandler{} + ctx.Task().Register(task.WarmImageCacheTask, func(c context.Context, t *asynq.Task) error { + _, err := warmHandler.Execute(c, t.Payload()) + return err + }, extpoints.WithTaskRetry(1)) + + // 4. Register Cron Schedule + ctx.Schedule().RegisterCron("0 3 * * *", task.SystemCleanupTask, nil) + + // 5. Register Settings Schemas + ctx.Settings().Register(extpoints.SettingSchema{ + Key: "upload.max_file_size_mb", + Default: 100, + Description: "Maximum upload file size limit in MB", + Type: "integer", + Category: "storage", + }) + ctx.Settings().Register(extpoints.SettingSchema{ + Key: "upload.allowed_extensions", + Default: "png,jpg,jpeg,gif,webp,svg,pdf,zip,tar,gz", + Description: "Allowed upload file extensions separated by commas", + Type: "string", + Category: "storage", + }) + + return nil +} diff --git a/internal/apps/upload/shared/constants.go b/plugins/domain/upload/shared/constants.go similarity index 100% rename from internal/apps/upload/shared/constants.go rename to plugins/domain/upload/shared/constants.go diff --git a/internal/apps/upload/shared/errs.go b/plugins/domain/upload/shared/errs.go similarity index 100% rename from internal/apps/upload/shared/errs.go rename to plugins/domain/upload/shared/errs.go diff --git a/internal/apps/upload/stats/category.go b/plugins/domain/upload/stats/category.go similarity index 94% rename from internal/apps/upload/stats/category.go rename to plugins/domain/upload/stats/category.go index 52fdc583..466dafe8 100644 --- a/internal/apps/upload/stats/category.go +++ b/plugins/domain/upload/stats/category.go @@ -7,7 +7,7 @@ package stats import ( "strings" - "github.com/Rain-kl/Wavelet/internal/apps/upload/util" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/util" ) const ( diff --git a/internal/apps/upload/stats/stats_counter.go b/plugins/domain/upload/stats/stats_counter.go similarity index 100% rename from internal/apps/upload/stats/stats_counter.go rename to plugins/domain/upload/stats/stats_counter.go diff --git a/internal/apps/upload/stats/stats_counter_test.go b/plugins/domain/upload/stats/stats_counter_test.go similarity index 100% rename from internal/apps/upload/stats/stats_counter_test.go rename to plugins/domain/upload/stats/stats_counter_test.go diff --git a/internal/apps/upload/storage/access_state.go b/plugins/domain/upload/storage/access_state.go similarity index 97% rename from internal/apps/upload/storage/access_state.go rename to plugins/domain/upload/storage/access_state.go index dceda766..ad45c22b 100644 --- a/internal/apps/upload/storage/access_state.go +++ b/plugins/domain/upload/storage/access_state.go @@ -9,9 +9,9 @@ import ( "sync" "time" - "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" "github.com/Rain-kl/Wavelet/internal/infra/objectstore" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" ) // MigrationAccessState captures cached migration maintenance state. diff --git a/internal/apps/upload/storage/migration.go b/plugins/domain/upload/storage/migration.go similarity index 100% rename from internal/apps/upload/storage/migration.go rename to plugins/domain/upload/storage/migration.go diff --git a/internal/apps/upload/storage/storage_ops.go b/plugins/domain/upload/storage/storage_ops.go similarity index 100% rename from internal/apps/upload/storage/storage_ops.go rename to plugins/domain/upload/storage/storage_ops.go diff --git a/internal/apps/upload/task/cleanup.go b/plugins/domain/upload/task/cleanup.go similarity index 95% rename from internal/apps/upload/task/cleanup.go rename to plugins/domain/upload/task/cleanup.go index 3e96351c..76881a50 100644 --- a/internal/apps/upload/task/cleanup.go +++ b/plugins/domain/upload/task/cleanup.go @@ -10,10 +10,6 @@ import ( "fmt" "time" - uploadcache "github.com/Rain-kl/Wavelet/internal/apps/upload/cache" - "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" - uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" - uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" "github.com/Rain-kl/Wavelet/internal/infra/objectstore" "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/task" @@ -21,6 +17,10 @@ import ( "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/repository/logstore" "github.com/Rain-kl/Wavelet/pkg/logger" + uploadcache "github.com/Rain-kl/Wavelet/plugins/domain/upload/cache" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" + uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats" + uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage" "gorm.io/gorm" ) diff --git a/internal/apps/upload/task/rebuild_stats.go b/plugins/domain/upload/task/rebuild_stats.go similarity index 97% rename from internal/apps/upload/task/rebuild_stats.go rename to plugins/domain/upload/task/rebuild_stats.go index 605a506b..94229eb4 100644 --- a/internal/apps/upload/task/rebuild_stats.go +++ b/plugins/domain/upload/task/rebuild_stats.go @@ -7,10 +7,10 @@ import ( "context" "fmt" - uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/model" + uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats" ) const ( diff --git a/internal/apps/upload/task/rebuild_stats_test.go b/plugins/domain/upload/task/rebuild_stats_test.go similarity index 100% rename from internal/apps/upload/task/rebuild_stats_test.go rename to plugins/domain/upload/task/rebuild_stats_test.go diff --git a/internal/apps/upload/task/storage_migration.go b/plugins/domain/upload/task/storage_migration.go similarity index 98% rename from internal/apps/upload/task/storage_migration.go rename to plugins/domain/upload/task/storage_migration.go index e2b1f128..513dd347 100644 --- a/internal/apps/upload/task/storage_migration.go +++ b/plugins/domain/upload/task/storage_migration.go @@ -15,13 +15,13 @@ import ( "sync/atomic" "time" - uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" - uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" "github.com/Rain-kl/Wavelet/internal/infra/objectstore" "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/pkg/util" + uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats" + uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage" "golang.org/x/sync/errgroup" ) diff --git a/internal/apps/upload/task/storage_migration_task_test.go b/plugins/domain/upload/task/storage_migration_task_test.go similarity index 100% rename from internal/apps/upload/task/storage_migration_task_test.go rename to plugins/domain/upload/task/storage_migration_task_test.go diff --git a/internal/apps/upload/task/tasks.go b/plugins/domain/upload/task/tasks.go similarity index 97% rename from internal/apps/upload/task/tasks.go rename to plugins/domain/upload/task/tasks.go index 7f69d9a5..8c36ea6c 100644 --- a/internal/apps/upload/task/tasks.go +++ b/plugins/domain/upload/task/tasks.go @@ -12,11 +12,11 @@ import ( "strings" "sync" - "github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv" - "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/filesrv" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" ) const ( diff --git a/internal/apps/upload/task/tasks_test.go b/plugins/domain/upload/task/tasks_test.go similarity index 99% rename from internal/apps/upload/task/tasks_test.go rename to plugins/domain/upload/task/tasks_test.go index 2999d261..ded2de0e 100644 --- a/internal/apps/upload/task/tasks_test.go +++ b/plugins/domain/upload/task/tasks_test.go @@ -17,8 +17,6 @@ import ( "testing" "time" - "github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv" - "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" "github.com/Rain-kl/Wavelet/internal/infra/diskcache" "github.com/Rain-kl/Wavelet/internal/infra/objectstore" "github.com/Rain-kl/Wavelet/internal/infra/persistence" @@ -26,6 +24,8 @@ import ( "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/testhelper" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/filesrv" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) diff --git a/internal/apps/upload/util/media.go b/plugins/domain/upload/util/media.go similarity index 95% rename from internal/apps/upload/util/media.go rename to plugins/domain/upload/util/media.go index eba2d864..121b1d0e 100644 --- a/internal/apps/upload/util/media.go +++ b/plugins/domain/upload/util/media.go @@ -6,7 +6,7 @@ package util import ( "strings" - "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" ) // IsImageExtension reports whether ext is a common image format. diff --git a/internal/apps/upload/util/utils.go b/plugins/domain/upload/util/utils.go similarity index 96% rename from internal/apps/upload/util/utils.go rename to plugins/domain/upload/util/utils.go index 1c9ffc31..d10d3914 100644 --- a/internal/apps/upload/util/utils.go +++ b/plugins/domain/upload/util/utils.go @@ -16,7 +16,7 @@ import ( "io" "strings" - "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" "github.com/deepteams/webp" _ "golang.org/x/image/webp" // Register WebP decoder for image.Decode ) diff --git a/plugins/domain/user/errs.go b/plugins/domain/user/errs.go new file mode 100644 index 00000000..d72c20ec --- /dev/null +++ b/plugins/domain/user/errs.go @@ -0,0 +1,15 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package user + +const ( + errInvalidParams = "无效的请求参数" + errUserNotFound = "用户不存在" + //nolint:gosec // error message, not hardcoded credentials + errPasswordMismatch = "用户名或密码错误" + //nolint:gosec // error message, not hardcoded credentials + errOldPasswordIncorrect = "原密码不正确" + //nolint:gosec // error message, not hardcoded credentials + errTokenNotFound = "访问令牌不存在" +) diff --git a/plugins/domain/user/handlers.go b/plugins/domain/user/handlers.go new file mode 100644 index 00000000..f0ab70c2 --- /dev/null +++ b/plugins/domain/user/handlers.go @@ -0,0 +1,297 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package user + +import ( + "crypto/rand" + "crypto/sha256" + "encoding/hex" + "net/http" + "strconv" + "time" + + persistence "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/shared/response" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/gin-contrib/sessions" + "github.com/gin-gonic/gin" +) + +type loginRequest struct { + Username string `json:"username" binding:"required"` + Password string `json:"password" binding:"required"` +} + +type registerRequest struct { + Username string `json:"username" binding:"required"` + Password string `json:"password" binding:"required"` + Email string `json:"email"` +} + +type changePasswordRequest struct { + OldPassword string `json:"old_password" binding:"required"` + NewPassword string `json:"new_password" binding:"required"` +} + +type updateProfileRequest struct { + Nickname string `json:"nickname"` + AvatarURL string `json:"avatar_url"` + Bio string `json:"bio"` + Phone string `json:"phone"` + Gender string `json:"gender"` + Website string `json:"website"` + Location string `json:"location"` +} + +type createAccessTokenRequest struct { + Name string `json:"name" binding:"required"` + ExpiresAt *time.Time `json:"expires_at"` + IsAdmin bool `json:"is_admin"` +} + +// Login handles username and password authentication. +func Login(c *gin.Context) { + var req loginRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.AbortBadRequest(c, errInvalidParams) + return + } + + user, err := repository.GetUserByUsername(c.Request.Context(), req.Username) + if err != nil { + response.AbortUnauthorized(c, errPasswordMismatch) + return + } + + if !user.CheckPassword(req.Password) { + response.AbortUnauthorized(c, errPasswordMismatch) + return + } + + sess := sessions.Default(c) + sess.Set(auth.UserIDKey, user.ID) + sess.Set(auth.UserNameKey, user.Username) + _ = sess.Save() + + c.JSON(http.StatusOK, response.OK(user)) +} + +// Register registers a new user. +func Register(c *gin.Context) { + var req registerRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.AbortBadRequest(c, errInvalidParams) + return + } + + newUser := &model.User{ + Username: req.Username, + Email: req.Email, + IsActive: true, + } + if err := newUser.SetEncryptedPassword(req.Password); err != nil { + response.AbortInternal(c, "密码加密失败") + return + } + + gormDB := persistence.DB(c.Request.Context()) + if err := gormDB.Create(newUser).Error; err != nil { + response.AbortBadRequest(c, "创建用户失败: "+err.Error()) + return + } + + c.JSON(http.StatusOK, response.OK(newUser)) +} + +// Logout logs out the current session. +func Logout(c *gin.Context) { + sess := sessions.Default(c) + sess.Clear() + _ = sess.Save() + c.JSON(http.StatusOK, response.OKNil()) +} + +// SendEmailCode sends an email verification code. +func SendEmailCode(c *gin.Context) { + c.JSON(http.StatusOK, response.OK(gin.H{"sent": true})) +} + +// ChangePassword changes the current user password. +func ChangePassword(c *gin.Context) { + var req changePasswordRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.AbortBadRequest(c, errInvalidParams) + return + } + + userID := auth.GetUserIDFromContext(c) + user, err := repository.GetUserByID(c.Request.Context(), userID) + if err != nil { + response.AbortNotFound(c, errUserNotFound) + return + } + + if !user.CheckPassword(req.OldPassword) { + response.AbortBadRequest(c, errOldPasswordIncorrect) + return + } + + if err := user.SetEncryptedPassword(req.NewPassword); err != nil { + response.AbortInternal(c, "密码更新失败") + return + } + + gormDB := persistence.DB(c.Request.Context()) + _ = gormDB.Save(&user) + auth.InvalidateCachedUser(c.Request.Context(), user.ID) + + c.JSON(http.StatusOK, response.OKNil()) +} + +// UpdateProfile updates profile info. +func UpdateProfile(c *gin.Context) { + var req updateProfileRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.AbortBadRequest(c, errInvalidParams) + return + } + + userID := auth.GetUserIDFromContext(c) + user, err := repository.GetUserByID(c.Request.Context(), userID) + if err != nil { + response.AbortNotFound(c, errUserNotFound) + return + } + + user.Nickname = req.Nickname + user.AvatarURL = req.AvatarURL + user.Bio = req.Bio + user.Phone = req.Phone + user.Gender = req.Gender + user.Website = req.Website + user.Location = req.Location + + gormDB := persistence.DB(c.Request.Context()) + _ = gormDB.Save(&user) + auth.InvalidateCachedUser(c.Request.Context(), user.ID) + + c.JSON(http.StatusOK, response.OK(user)) +} + +// ListAccessTokens lists access tokens for the current user. +func ListAccessTokens(c *gin.Context) { + userID := auth.GetUserIDFromContext(c) + var tokens []model.AccessToken + gormDB := persistence.DB(c.Request.Context()) + _ = gormDB.Where("user_id = ?", userID).Find(&tokens).Error + c.JSON(http.StatusOK, response.OK(tokens)) +} + +const ( + tokenEntropyByteLength = 24 + tokenMaskMinLength = 8 +) + +// CreateAccessToken generates a new access token. +func CreateAccessToken(c *gin.Context) { + var req createAccessTokenRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.AbortBadRequest(c, errInvalidParams) + return + } + + userID := auth.GetUserIDFromContext(c) + rawBytes := make([]byte, tokenEntropyByteLength) + _, _ = rand.Read(rawBytes) + rawToken := "wvt_" + hex.EncodeToString(rawBytes) + hash := sha256.Sum256([]byte(rawToken)) + tokenHash := hex.EncodeToString(hash[:]) + + masked := rawToken + if len(rawToken) > tokenMaskMinLength { + masked = rawToken[:4] + "..." + rawToken[len(rawToken)-4:] + } + + token := model.AccessToken{ + UserID: userID, + Name: req.Name, + TokenHash: tokenHash, + MaskedToken: masked, + IsAdmin: req.IsAdmin, + } + + gormDB := persistence.DB(c.Request.Context()) + if err := gormDB.Create(&token).Error; err != nil { + response.AbortInternal(c, "创建令牌失败") + return + } + + c.JSON(http.StatusOK, response.OK(gin.H{ + "token": token, + "raw_token": rawToken, + })) +} + +// DeleteAccessToken deletes a specific access token. +func DeleteAccessToken(c *gin.Context) { + idStr := c.Param("id") + id, err := strconv.ParseUint(idStr, 10, 64) + if err != nil { + response.AbortBadRequest(c, errInvalidParams) + return + } + + userID := auth.GetUserIDFromContext(c) + var token model.AccessToken + gormDB := persistence.DB(c.Request.Context()) + if err := gormDB.Where("id = ? AND user_id = ?", id, userID).First(&token).Error; err != nil { + response.AbortNotFound(c, errTokenNotFound) + return + } + + _ = gormDB.Delete(&token) + auth.InvalidateCachedToken(c.Request.Context(), token.TokenHash) + c.JSON(http.StatusOK, response.OKNil()) +} + +// RotateAccessToken rotates an access token value. +func RotateAccessToken(c *gin.Context) { + idStr := c.Param("id") + id, err := strconv.ParseUint(idStr, 10, 64) + if err != nil { + response.AbortBadRequest(c, errInvalidParams) + return + } + + userID := auth.GetUserIDFromContext(c) + var token model.AccessToken + gormDB := persistence.DB(c.Request.Context()) + if err := gormDB.Where("id = ? AND user_id = ?", id, userID).First(&token).Error; err != nil { + response.AbortNotFound(c, errTokenNotFound) + return + } + + auth.InvalidateCachedToken(c.Request.Context(), token.TokenHash) + + rawBytes := make([]byte, tokenEntropyByteLength) + _, _ = rand.Read(rawBytes) + rawToken := "wvt_" + hex.EncodeToString(rawBytes) + hash := sha256.Sum256([]byte(rawToken)) + token.TokenHash = hex.EncodeToString(hash[:]) + + masked := rawToken + if len(rawToken) > tokenMaskMinLength { + masked = rawToken[:4] + "..." + rawToken[len(rawToken)-4:] + } + token.MaskedToken = masked + + _ = gormDB.Save(&token) + + c.JSON(http.StatusOK, response.OK(gin.H{ + "token": token, + "raw_token": rawToken, + })) +} diff --git a/plugins/domain/user/plugin.go b/plugins/domain/user/plugin.go index 432b1eaf..8dcf7e25 100644 --- a/plugins/domain/user/plugin.go +++ b/plugins/domain/user/plugin.go @@ -1,3 +1,6 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + // Package user provides the user profile, credential management, role management, and access token domain plugin for Cordis. package user @@ -8,8 +11,7 @@ import ( "github.com/Rain-kl/Wavelet/core" "github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/core/extpoints" - "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/apps/user" + "github.com/Rain-kl/Wavelet/plugins/domain/auth" "github.com/hibiken/asynq" ) @@ -71,20 +73,20 @@ func (p *Plugin) Apply(ctx *core.Context) error { // 3. Register HTTP Routes userGroup := ctx.Router().Group("/api/v1/user") { - userGroup.POST("/login", user.Login) - userGroup.POST("/register", user.Register) - userGroup.GET("/logout", user.Logout) - userGroup.POST("/send-email-code", user.SendEmailCode) - userGroup.POST("/change-password", oauth.LoginRequired(), user.ChangePassword) - userGroup.PUT("/profile", oauth.LoginRequired(), user.UpdateProfile) + userGroup.POST("/login", Login) + userGroup.POST("/register", Register) + userGroup.GET("/logout", Logout) + userGroup.POST("/send-email-code", SendEmailCode) + userGroup.POST("/change-password", auth.LoginRequired(), ChangePassword) + userGroup.PUT("/profile", auth.LoginRequired(), UpdateProfile) // Access Tokens - tokensGroup := userGroup.Group("/access-tokens", oauth.LoginRequired(), oauth.DisallowTokenAuth()) + tokensGroup := userGroup.Group("/access-tokens", auth.LoginRequired(), auth.DisallowTokenAuth()) { - tokensGroup.GET("", user.ListAccessTokens) - tokensGroup.POST("", user.CreateAccessToken) - tokensGroup.DELETE("/:id", user.DeleteAccessToken) - tokensGroup.POST("/:id/rotate", user.RotateAccessToken) + tokensGroup.GET("", ListAccessTokens) + tokensGroup.POST("", CreateAccessToken) + tokensGroup.DELETE("/:id", DeleteAccessToken) + tokensGroup.POST("/:id/rotate", RotateAccessToken) } } diff --git a/plugins/infra/storage/plugin.go b/plugins/infra/storage/plugin.go index ea572aee..469d7e46 100644 --- a/plugins/infra/storage/plugin.go +++ b/plugins/infra/storage/plugin.go @@ -8,9 +8,9 @@ import ( "github.com/Rain-kl/Wavelet/core" "github.com/Rain-kl/Wavelet/core/contracts" - "github.com/Rain-kl/Wavelet/internal/apps/upload/ingest" "github.com/Rain-kl/Wavelet/internal/infra/objectstore" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/plugins/domain/upload/ingest" ) // Option configures the storage plugin.