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 := `
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/