mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-07 08:06:37 +08:00
merge: replace legacy openflare-server with Wavelet rename
This commit is contained in:
@@ -0,0 +1,245 @@
|
||||
// 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/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/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 := model.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 := model.CreateAuthSource(c.Request.Context(), &source); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
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, _ := model.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 := model.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))
|
||||
|
||||
updated, err := model.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 := model.ToggleAuthSource(c.Request.Context(), id, req.IsActive); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
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 := model.DeleteAuthSource(c.Request.Context(), id); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
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 := model.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
|
||||
}
|
||||
@@ -0,0 +1,333 @@
|
||||
// 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/common/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")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,14 @@
|
||||
// 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)
|
||||
}
|
||||
@@ -0,0 +1,100 @@
|
||||
// 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/diskcache"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/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())
|
||||
}
|
||||
@@ -0,0 +1,498 @@
|
||||
// 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/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/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/openflare.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))
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
// 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"
|
||||
)
|
||||
@@ -0,0 +1,46 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package admin
|
||||
|
||||
import (
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"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()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,278 @@
|
||||
// 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/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
// 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
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
// 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{},
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
// 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"
|
||||
)
|
||||
@@ -0,0 +1,38 @@
|
||||
// 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)
|
||||
}
|
||||
@@ -0,0 +1,175 @@
|
||||
// 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/listener"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
"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, dbConn.Save(&event).Error)
|
||||
}
|
||||
|
||||
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, dbConn.Save(&event).Error)
|
||||
|
||||
listener.EmitAdminLoggedIn(context.Background(), adminUser, "10.0.0.1")
|
||||
waitForAsyncTrigger(t)
|
||||
|
||||
assert.Equal(t, int64(0), countPushTasks(t, dbConn))
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
// 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)
|
||||
}
|
||||
@@ -0,0 +1,343 @@
|
||||
// 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/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
||||
"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)
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,413 @@
|
||||
// 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/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,759 @@
|
||||
// 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/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
"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"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
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) 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 = dbConn.Save(&event).Error
|
||||
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 = dbConn.Save(&event).Error
|
||||
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"}
|
||||
dbConn.Save(&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)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,292 @@
|
||||
// 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/common/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())
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
"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 ""
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
// 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/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()
|
||||
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 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)
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,355 @@
|
||||
// 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/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/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/openflare.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/openflare.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="openflare.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, "openflare.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("openflare_%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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
// 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 = "配置键已存在"
|
||||
StorageDriverSwitchRequiresMigration = "存在存量文件,请通过存储迁移任务切换存储引擎"
|
||||
)
|
||||
@@ -0,0 +1,128 @@
|
||||
// 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/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/storage"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func createSystemConfig(ctx context.Context, req CreateSystemConfigRequest) error {
|
||||
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 {
|
||||
config, err := repository.GetAdminSystemConfigByKey(ctx, key)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var originalDriver storage.Driver
|
||||
if key == model.ConfigKeyStorageConfig {
|
||||
var currentCfg storage.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 storage.Driver,
|
||||
newValue string,
|
||||
) {
|
||||
if key != model.ConfigKeyStorageConfig || originalDriver == "" {
|
||||
return
|
||||
}
|
||||
|
||||
var newCfg storage.Config
|
||||
if err := json.Unmarshal([]byte(newValue), &newCfg); err != nil {
|
||||
return
|
||||
}
|
||||
if newCfg.Driver != originalDriver {
|
||||
return
|
||||
}
|
||||
|
||||
if err := tx.Model(&model.TaskExecution{}).
|
||||
Where("task_type = ? AND status = ?", "storage:migrate", model.TaskExecutionStatusFailed).
|
||||
Updates(map[string]any{
|
||||
"status": model.TaskExecutionStatusSucceeded,
|
||||
"result": "存储配置直接更新,故障迁移任务自动标记为已解决",
|
||||
"finished_at": time.Now(),
|
||||
}).Error; err != nil {
|
||||
logger.ErrorF(ctx, "自动更新迁移任务状态失败: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,365 @@
|
||||
// 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/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/storage"
|
||||
"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 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 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)
|
||||
storage.ResetCache()
|
||||
storage.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 := "OpenFlare SMTP Test Mail"
|
||||
body := `<h3>SMTP Mail Connection Test</h3>
|
||||
<p>If you received this message, your SMTP configuration is correct and mail sending is working properly.</p>
|
||||
<p>Sent from OpenFlare.</p>`
|
||||
|
||||
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 storage.Config
|
||||
if err := json.Unmarshal([]byte(value), &cfg); err == nil {
|
||||
masked := storage.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 storage.Config
|
||||
if err := json.Unmarshal([]byte(currentConfig), ¤tCfg); err != nil {
|
||||
return "", fmt.Errorf("解析当前存储配置失败: %w", err)
|
||||
}
|
||||
|
||||
var newCfg storage.Config
|
||||
if err := json.Unmarshal([]byte(value), &newCfg); err != nil {
|
||||
return "", fmt.Errorf("解析目标存储配置失败: %w", err)
|
||||
}
|
||||
|
||||
// 合并被掩码屏蔽的敏感信息,获取完整的真实配置
|
||||
targetCfg := storage.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 storage.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 := storage.ValidateConfig(targetCfg); err != nil {
|
||||
return fmt.Errorf("验证存储配置参数失败: %w", err)
|
||||
}
|
||||
return testStorageBackend(ctx, targetCfg, targetCfg.Driver)
|
||||
}
|
||||
|
||||
func validateDriverConfig(cfg storage.Config, driver storage.Driver) error {
|
||||
cfg.Driver = driver
|
||||
return storage.ValidateConfig(cfg)
|
||||
}
|
||||
|
||||
func testStorageBackend(ctx context.Context, cfg storage.Config, driver storage.Driver) error {
|
||||
testBackend, err := storage.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
|
||||
}
|
||||
@@ -0,0 +1,543 @@
|
||||
// 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/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/storage"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
const expectedDefaultConfigsCount = 30
|
||||
|
||||
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) != 1 || configs[0].Key != model.ConfigKeyMaxAPIKeysPerUser {
|
||||
t.Errorf("expected 1 business config (max_api_keys_per_user), got %d: %v", len(configs), configs)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
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 != "OpenFlare" {
|
||||
t.Errorf("expected 'OpenFlare', 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 := storage.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 storage.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 := storage.DefaultConfig()
|
||||
cfg.Driver = storage.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 := storage.DefaultConfig()
|
||||
cfg.Driver = storage.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 := storage.DefaultConfig()
|
||||
activeCfg.Driver = storage.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 = storage.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 storage.Config
|
||||
if err := json.Unmarshal([]byte(dbCfg.Value), &savedCfg); err != nil {
|
||||
t.Fatalf("parse saved storage config failed: %v", err)
|
||||
}
|
||||
if savedCfg.Driver != storage.DriverLocal {
|
||||
t.Fatalf("active driver = %q, want %q after save", savedCfg.Driver, storage.DriverLocal)
|
||||
}
|
||||
if savedCfg.Local.Root != tempDir {
|
||||
t.Fatalf("staged local root = %q, want %q", savedCfg.Local.Root, tempDir)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
// 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 = "删除定时任务失败"
|
||||
)
|
||||
@@ -0,0 +1,416 @@
|
||||
// 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/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
"github.com/Rain-kl/Wavelet/internal/task/scheduler"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/robfig/cron/v3"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/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 := model.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 := model.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 := model.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 := model.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 := model.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 := model.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 := model.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())
|
||||
}
|
||||
@@ -0,0 +1,495 @@
|
||||
// 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/bootstrap"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
"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/common/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 := model.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 := model.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 := model.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 := model.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 := model.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 := model.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)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package template 提供模板管理功能
|
||||
package template
|
||||
|
||||
// 模板管理相关错误消息
|
||||
const (
|
||||
TemplateNotFound = "模板不存在"
|
||||
TemplateKeyRequired = "模板标识符不能为空"
|
||||
TemplateNameRequired = "模板名称不能为空"
|
||||
TemplateContentRequired = "模板内容不能为空"
|
||||
TemplateKeyExists = "模板标识符已存在"
|
||||
SystemTemplateCannotDelete = "系统预置模板不可删除"
|
||||
SystemTemplateCannotModifyKey = "系统预置模板不可修改标识符"
|
||||
)
|
||||
@@ -0,0 +1,78 @@
|
||||
// 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)
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
// 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/common/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())
|
||||
}
|
||||
@@ -0,0 +1,242 @@
|
||||
// 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/common/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")
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
// 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 = "当前平台暂不支持自动替换二进制"
|
||||
)
|
||||
@@ -0,0 +1,40 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package updater
|
||||
|
||||
import "context"
|
||||
|
||||
// GetStatus returns the current build and the newest compatible upstream release.
|
||||
func GetStatus(ctx context.Context) (Status, error) {
|
||||
status, _, err := defaultManager.status(ctx)
|
||||
return status, err
|
||||
}
|
||||
|
||||
// PrepareUpgrade downloads and stages the upgrade binary for the current platform.
|
||||
func PrepareUpgrade(ctx context.Context) (executable string, stagedBinary string, status Status, err error) {
|
||||
status, _, err = defaultManager.status(ctx)
|
||||
if err != nil {
|
||||
return "", "", Status{}, err
|
||||
}
|
||||
|
||||
executable, stagedBinary, err = defaultManager.prepareUpgrade(ctx)
|
||||
return executable, stagedBinary, status, err
|
||||
}
|
||||
|
||||
// ApplyPreparedUpgrade replaces the running binary and restarts the process.
|
||||
func ApplyPreparedUpgrade(executable, stagedBinary string) error {
|
||||
return replaceAndRestart(executable, stagedBinary)
|
||||
}
|
||||
|
||||
// FinishUpgrade clears the in-progress upgrade flag after a failed restart.
|
||||
func FinishUpgrade() {
|
||||
defaultManager.finishUpgrade()
|
||||
}
|
||||
|
||||
// IsUpgrading reports whether an upgrade task is currently running.
|
||||
func IsUpgrading() bool {
|
||||
defaultManager.mu.Lock()
|
||||
defer defaultManager.mu.Unlock()
|
||||
return defaultManager.upgrading
|
||||
}
|
||||
@@ -0,0 +1,647 @@
|
||||
// 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", "OpenFlare-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", "OpenFlare-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
|
||||
}
|
||||
@@ -0,0 +1,130 @@
|
||||
// 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/OpenFlare", want: "Rain-kl/OpenFlare"},
|
||||
{name: "GitHub URL", input: "https://github.com/Rain-kl/OpenFlare.git", want: "Rain-kl/OpenFlare"},
|
||||
{name: "unsupported host", input: "https://example.com/Rain-kl/OpenFlare", wantErr: true},
|
||||
{name: "missing owner", input: "OpenFlare", 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/OpenFlare", 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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
//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.
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
//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)
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
// 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/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/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())
|
||||
|
||||
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)
|
||||
}
|
||||
}()
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
// 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 = "邮箱已被注册"
|
||||
)
|
||||
@@ -0,0 +1,106 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package user
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db/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),
|
||||
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)
|
||||
}
|
||||
return repository.UpdateUserActive(ctx, id, active)
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
return repository.DeleteUserWithRelations(ctx, targetID)
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,288 @@
|
||||
// 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/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/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"`
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
}
|
||||
response.AbortInternal(c, msg)
|
||||
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 {
|
||||
response.AbortInternal(c, err.Error())
|
||||
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)))
|
||||
}
|
||||
@@ -0,0 +1,582 @@
|
||||
// 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/common/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())
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package cap 提供人机验证中间件
|
||||
package cap
|
||||
|
||||
const (
|
||||
errCapTokenMissing = "验证码验证失败,缺少验证码凭证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errCapTokenInvalidOrExpired = "验证码校验失败或已过期,请重试" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
)
|
||||
@@ -0,0 +1,196 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package cap provides CAPTCHA and proof-of-work (PoW) verification services.
|
||||
package cap
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
pkgcap "github.com/Rain-kl/Wavelet/pkg/cap"
|
||||
)
|
||||
|
||||
const (
|
||||
redeemTokenIDLength = 8 // 兑换 Token ID 字节长度
|
||||
redeemVerTokenLength = 15 // 兑换验证 Token 字节长度
|
||||
tokenPartsCount = 2 // 兑换 Token 由两部分组成
|
||||
valuePartsCount = 2 // 存储值由 scope 和过期时间组成
|
||||
)
|
||||
|
||||
// Manager orchestrates challenge generation and solution validation.
|
||||
type Manager struct {
|
||||
secret []byte
|
||||
store pkgcap.Store
|
||||
}
|
||||
|
||||
// NewManager creates a new CAPTCHA Manager.
|
||||
func NewManager(secret []byte, store pkgcap.Store) *Manager {
|
||||
return &Manager{
|
||||
secret: secret,
|
||||
store: store,
|
||||
}
|
||||
}
|
||||
|
||||
// Generate creates a challenge response.
|
||||
func (m *Manager) Generate(ctx context.Context, scope string) (*pkgcap.ChallengeResponse, error) {
|
||||
settings, err := CurrentSettings(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
challengeConfig := pkgcap.ChallengeConfig{
|
||||
Count: settings.ChallengeCount,
|
||||
Size: settings.ChallengeSize,
|
||||
Difficulty: settings.ChallengeDifficulty,
|
||||
Expires: settings.ChallengeTTL,
|
||||
}
|
||||
return pkgcap.GenerateChallenge(m.secret, challengeConfig, scope)
|
||||
}
|
||||
|
||||
// RedeemResponse is returned to the client on redeem.
|
||||
type RedeemResponse struct {
|
||||
Success bool `json:"success"`
|
||||
Token string `json:"token,omitempty"`
|
||||
Expires int64 `json:"expires,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
// Redeem verifies PoW solutions and returns a one-time redeem token.
|
||||
func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, scope string) (*RedeemResponse, error) {
|
||||
sigHex := pkgcap.JwtSigHex(token)
|
||||
if sigHex == "" {
|
||||
return &RedeemResponse{Success: false, Error: "invalid_token"}, nil
|
||||
}
|
||||
|
||||
nonceKey := "cap:nonce:" + sigHex
|
||||
|
||||
payload, err := pkgcap.VerifyChallengeSolutions(token, solutions, m.secret, scope)
|
||||
if err != nil {
|
||||
return &RedeemResponse{Success: false, Error: err.Error()}, nil //nolint:nilerr // validation errors are returned as response, not system errors
|
||||
}
|
||||
|
||||
now := time.Now().UnixNano() / int64(time.Millisecond)
|
||||
nonceTTL := time.Duration(payload.Expires-now) * time.Millisecond
|
||||
if nonceTTL < time.Second {
|
||||
nonceTTL = time.Second
|
||||
}
|
||||
|
||||
set, err := m.store.SetNX(ctx, nonceKey, "1", nonceTTL)
|
||||
if err != nil {
|
||||
return &RedeemResponse{Success: false, Error: "nonce_store_error"}, err
|
||||
}
|
||||
if !set {
|
||||
return &RedeemResponse{Success: false, Error: "already_redeemed"}, nil
|
||||
}
|
||||
|
||||
settings, err := CurrentSettings(ctx)
|
||||
if err != nil {
|
||||
return &RedeemResponse{Success: false, Error: "settings_load_error"}, err
|
||||
}
|
||||
|
||||
id := pkgcap.RandomHex(redeemTokenIDLength)
|
||||
verToken := pkgcap.RandomHex(redeemVerTokenLength)
|
||||
verHashBytes := sha256.Sum256([]byte(verToken))
|
||||
verHashHex := hex.EncodeToString(verHashBytes[:])
|
||||
|
||||
tokenKey := "cap:token:" + id + ":" + verHashHex
|
||||
tokenExpires := time.Now().Add(settings.TokenTTL)
|
||||
storeVal := strconv.FormatInt(tokenExpires.UnixNano(), 10) + "|" + scope
|
||||
|
||||
if err := m.store.Set(ctx, tokenKey, storeVal, settings.TokenTTL); err != nil {
|
||||
return &RedeemResponse{Success: false, Error: "token_store_error"}, err
|
||||
}
|
||||
|
||||
return &RedeemResponse{
|
||||
Success: true,
|
||||
Token: id + ":" + verToken,
|
||||
Expires: tokenExpires.UnixNano() / int64(time.Millisecond),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// VerifyToken validates and consumes the redeem token (single-use).
|
||||
func (m *Manager) VerifyToken(ctx context.Context, token string, expectedScope string) (bool, error) {
|
||||
if token == "" {
|
||||
return false, nil
|
||||
}
|
||||
parts := strings.Split(token, ":")
|
||||
if len(parts) != tokenPartsCount {
|
||||
return false, nil
|
||||
}
|
||||
id := parts[0]
|
||||
verToken := parts[1]
|
||||
|
||||
verHashBytes := sha256.Sum256([]byte(verToken))
|
||||
verHashHex := hex.EncodeToString(verHashBytes[:])
|
||||
|
||||
tokenKey := "cap:token:" + id + ":" + verHashHex
|
||||
|
||||
val, exists, err := sGetAndDelete(ctx, m.store, tokenKey)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if !exists {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
valParts := strings.Split(val, "|")
|
||||
if len(valParts) != valuePartsCount {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
expNano, err := strconv.ParseInt(valParts[0], 10, 64)
|
||||
if err != nil {
|
||||
return false, nil //nolint:nilerr // invalid format is treated as validation failure
|
||||
}
|
||||
tokenScope := valParts[1]
|
||||
|
||||
if expectedScope != "" && tokenScope != expectedScope {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
if time.Now().UnixNano() > expNano {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func sGetAndDelete(ctx context.Context, store pkgcap.Store, key string) (string, bool, error) {
|
||||
if store == nil {
|
||||
return "", false, nil
|
||||
}
|
||||
return store.GetAndDelete(ctx, key)
|
||||
}
|
||||
|
||||
var (
|
||||
defaultManager *Manager
|
||||
once sync.Once
|
||||
)
|
||||
|
||||
// GetDefaultManager yields the global singleton CAPTCHA manager.
|
||||
func GetDefaultManager() *Manager {
|
||||
once.Do(func() {
|
||||
secret := []byte("default-captcha-secret-key-at-least-16-bytes")
|
||||
if config.Config != nil && config.Config.App.SessionSecret != "" {
|
||||
secret = []byte(config.Config.App.SessionSecret)
|
||||
}
|
||||
|
||||
var store pkgcap.Store
|
||||
if config.Config != nil && config.Config.Redis.Enabled && db.Redis != nil {
|
||||
store = pkgcap.NewRedisStore(db.Redis)
|
||||
} else {
|
||||
store = pkgcap.NewMemoryStore(1 * time.Minute)
|
||||
}
|
||||
|
||||
defaultManager = NewManager(secret, store)
|
||||
})
|
||||
return defaultManager
|
||||
}
|
||||
@@ -0,0 +1,162 @@
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
|
||||
import (
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
// VerifyMiddleware returns a Gin middleware that checks and consumes the X-Cap-Token header.
|
||||
func VerifyMiddleware(mgr *Manager, scope string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if !ProtectionEnabled(c.Request.Context()) {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
token := c.GetHeader("X-Cap-Token")
|
||||
if token == "" {
|
||||
response.AbortUnauthorized(c, errCapTokenMissing)
|
||||
return
|
||||
}
|
||||
|
||||
valid, err := mgr.VerifyToken(c.Request.Context(), token, scope)
|
||||
if err != nil || !valid {
|
||||
response.AbortUnauthorized(c, errCapTokenInvalidOrExpired)
|
||||
return
|
||||
}
|
||||
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type challengeRequest struct {
|
||||
Scope string `json:"scope" form:"scope"`
|
||||
}
|
||||
|
||||
type redeemRequest struct {
|
||||
Token string `json:"token" binding:"required"`
|
||||
Solutions []int `json:"solutions" binding:"required"`
|
||||
Scope string `json:"scope" form:"scope"`
|
||||
}
|
||||
|
||||
// Challenge 生成 PoW 人机验证难题
|
||||
// @Summary 生成人机验证难题
|
||||
// @Description 客户端获取 PoW 难题和签名的 JWT Token,并在后台计算。
|
||||
// @Tags cap
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body challengeRequest false "可选范围限制参数"
|
||||
// @Success 200 {object} cap.ChallengeResponse "成功返回 PoW 难题"
|
||||
// @Failure 500 {object} RedeemResponse "内部服务错误"
|
||||
// @Router /api/cap/challenge [post]
|
||||
func Challenge(c *gin.Context) {
|
||||
var req challengeRequest
|
||||
_ = c.ShouldBind(&req) // 允许不传 body,默认使用 login scope
|
||||
|
||||
if req.Scope == "" {
|
||||
req.Scope = "login"
|
||||
}
|
||||
|
||||
mgr := GetDefaultManager()
|
||||
resp, err := mgr.Generate(c.Request.Context(), req.Scope)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, RedeemResponse{
|
||||
Success: false,
|
||||
Error: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, resp)
|
||||
}
|
||||
|
||||
// Redeem 提交 PoW 解答并兑换一次性凭证 Token
|
||||
// @Summary 校验人机验证解答
|
||||
// @Description 提交 PoW 解答进行核销,成功后返回一次性 X-Cap-Token 凭证
|
||||
// @Tags cap
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body redeemRequest true "难题 Token 与解答 solutions 数组"
|
||||
// @Success 200 {object} RedeemResponse "核销成功,返回 X-Cap-Token"
|
||||
// @Failure 400 {object} RedeemResponse "参数错误或核销失败"
|
||||
// @Failure 500 {object} RedeemResponse "内部服务错误"
|
||||
// @Router /api/cap/redeem [post]
|
||||
func Redeem(c *gin.Context) {
|
||||
var req redeemRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, RedeemResponse{
|
||||
Success: false,
|
||||
Error: "无效的参数",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if req.Scope == "" {
|
||||
req.Scope = "login"
|
||||
}
|
||||
|
||||
mgr := GetDefaultManager()
|
||||
resp, err := mgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, RedeemResponse{
|
||||
Success: false,
|
||||
Error: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if !resp.Success {
|
||||
c.JSON(http.StatusBadRequest, resp)
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, resp)
|
||||
}
|
||||
@@ -0,0 +1,128 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
pkgcap "github.com/Rain-kl/Wavelet/pkg/cap"
|
||||
)
|
||||
|
||||
func TestCapEndpointsAndMiddleware(t *testing.T) {
|
||||
sqliteDB, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
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"))
|
||||
})
|
||||
|
||||
// 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())
|
||||
}
|
||||
|
||||
var challengeResp pkgcap.ChallengeResponse
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &challengeResp); err != nil {
|
||||
t.Fatalf("failed to unmarshal challenge response: %v", err)
|
||||
}
|
||||
|
||||
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())
|
||||
}
|
||||
|
||||
var redeemResp RedeemResponse
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &redeemResp); err != nil {
|
||||
t.Fatalf("failed to unmarshal redeem response: %v", err)
|
||||
}
|
||||
|
||||
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())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,212 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strconv"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"golang.org/x/sync/singleflight"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultChallengeCount = 1
|
||||
defaultChallengeSize = 32
|
||||
defaultChallengeDifficulty = 4
|
||||
defaultChallengeTTL = 10 * time.Minute
|
||||
defaultTokenTTL = 20 * time.Minute
|
||||
)
|
||||
|
||||
// RuntimeSettings is the parsed CAPTCHA runtime configuration loaded from system_configs.
|
||||
type RuntimeSettings struct {
|
||||
LoginEnabled bool
|
||||
ChallengeCount int
|
||||
ChallengeSize int
|
||||
ChallengeDifficulty int
|
||||
ChallengeTTL time.Duration
|
||||
TokenTTL time.Duration
|
||||
}
|
||||
|
||||
var runtimeConfigKeys = []string{
|
||||
model.ConfigKeyCapLoginEnabled,
|
||||
model.ConfigKeyCapChallengeCount,
|
||||
model.ConfigKeyCapChallengeSize,
|
||||
model.ConfigKeyCapChallengeDifficulty,
|
||||
model.ConfigKeyCapChallengeTTL,
|
||||
model.ConfigKeyCapTokenTTL,
|
||||
}
|
||||
|
||||
var runtimeConfigKeySet = func() map[string]struct{} {
|
||||
set := make(map[string]struct{}, len(runtimeConfigKeys))
|
||||
for _, key := range runtimeConfigKeys {
|
||||
set[key] = struct{}{}
|
||||
}
|
||||
return set
|
||||
}()
|
||||
|
||||
type runtimeSettingsStore struct {
|
||||
snapshot atomic.Pointer[RuntimeSettings]
|
||||
loadGroup singleflight.Group
|
||||
listenerOnce sync.Once
|
||||
}
|
||||
|
||||
var settingsStore = &runtimeSettingsStore{}
|
||||
|
||||
// IsRuntimeConfigKey reports whether a system config key affects CAPTCHA runtime settings.
|
||||
func IsRuntimeConfigKey(key string) bool {
|
||||
_, ok := runtimeConfigKeySet[key]
|
||||
return ok
|
||||
}
|
||||
|
||||
// CurrentSettings returns the cached CAPTCHA runtime settings snapshot.
|
||||
func CurrentSettings(ctx context.Context) (RuntimeSettings, error) {
|
||||
return settingsStore.current(ctx)
|
||||
}
|
||||
|
||||
// ProtectionEnabled reports whether CAPTCHA verification is required for protected routes.
|
||||
func ProtectionEnabled(ctx context.Context) bool {
|
||||
settings, err := CurrentSettings(ctx)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return settings.LoginEnabled
|
||||
}
|
||||
|
||||
// InvalidateRuntimeSettings drops the in-process CAPTCHA settings snapshot.
|
||||
func InvalidateRuntimeSettings() {
|
||||
settingsStore.snapshot.Store(nil)
|
||||
}
|
||||
|
||||
// ResetRuntimeSettingsForTest clears the CAPTCHA runtime snapshot.
|
||||
func ResetRuntimeSettingsForTest() {
|
||||
InvalidateRuntimeSettings()
|
||||
}
|
||||
|
||||
// InstallTestRuntimeSettings installs a fixed snapshot for unit tests.
|
||||
func InstallTestRuntimeSettings(settings RuntimeSettings) func() {
|
||||
snapshot := settings
|
||||
settingsStore.snapshot.Store(&snapshot)
|
||||
return InvalidateRuntimeSettings
|
||||
}
|
||||
|
||||
func (s *runtimeSettingsStore) current(ctx context.Context) (RuntimeSettings, error) {
|
||||
s.ensureInvalidationListener()
|
||||
|
||||
if snapshot := s.snapshot.Load(); snapshot != nil {
|
||||
return *snapshot, nil
|
||||
}
|
||||
|
||||
loaded, err, _ := s.loadGroup.Do("cap-runtime-settings", func() (any, error) {
|
||||
if snapshot := s.snapshot.Load(); snapshot != nil {
|
||||
return *snapshot, nil
|
||||
}
|
||||
|
||||
settings, loadErr := loadRuntimeSettings(ctx)
|
||||
if loadErr != nil {
|
||||
return RuntimeSettings{}, loadErr
|
||||
}
|
||||
|
||||
s.snapshot.Store(&settings)
|
||||
return settings, nil
|
||||
})
|
||||
if err != nil {
|
||||
return RuntimeSettings{}, err
|
||||
}
|
||||
|
||||
settings, ok := loaded.(RuntimeSettings)
|
||||
if !ok {
|
||||
return RuntimeSettings{}, errors.New("cap runtime settings loader returned unexpected type")
|
||||
}
|
||||
return settings, nil
|
||||
}
|
||||
|
||||
func loadRuntimeSettings(ctx context.Context) (RuntimeSettings, error) {
|
||||
configs, err := repository.ListSystemConfigsByKeys(ctx, runtimeConfigKeys)
|
||||
if err != nil {
|
||||
return RuntimeSettings{}, err
|
||||
}
|
||||
return parseRuntimeSettings(configs), nil
|
||||
}
|
||||
|
||||
func parseRuntimeSettings(configs map[string]model.SystemConfig) RuntimeSettings {
|
||||
settings := RuntimeSettings{
|
||||
ChallengeCount: defaultChallengeCount,
|
||||
ChallengeSize: defaultChallengeSize,
|
||||
ChallengeDifficulty: defaultChallengeDifficulty,
|
||||
ChallengeTTL: defaultChallengeTTL,
|
||||
TokenTTL: defaultTokenTTL,
|
||||
}
|
||||
|
||||
if sc, ok := configs[model.ConfigKeyCapLoginEnabled]; ok {
|
||||
if enabled, err := strconv.ParseBool(sc.Value); err == nil {
|
||||
settings.LoginEnabled = enabled
|
||||
}
|
||||
}
|
||||
if sc, ok := configs[model.ConfigKeyCapChallengeCount]; ok {
|
||||
if count, err := strconv.Atoi(sc.Value); err == nil && count > 0 {
|
||||
settings.ChallengeCount = count
|
||||
}
|
||||
}
|
||||
if sc, ok := configs[model.ConfigKeyCapChallengeSize]; ok {
|
||||
if size, err := strconv.Atoi(sc.Value); err == nil && size > 0 {
|
||||
settings.ChallengeSize = size
|
||||
}
|
||||
}
|
||||
if sc, ok := configs[model.ConfigKeyCapChallengeDifficulty]; ok {
|
||||
if difficulty, err := strconv.Atoi(sc.Value); err == nil && difficulty > 0 {
|
||||
settings.ChallengeDifficulty = difficulty
|
||||
}
|
||||
}
|
||||
if sc, ok := configs[model.ConfigKeyCapChallengeTTL]; ok {
|
||||
if ttlSeconds, err := strconv.Atoi(sc.Value); err == nil && ttlSeconds > 0 {
|
||||
settings.ChallengeTTL = time.Duration(ttlSeconds) * time.Second
|
||||
}
|
||||
}
|
||||
if sc, ok := configs[model.ConfigKeyCapTokenTTL]; ok {
|
||||
if ttlSeconds, err := strconv.Atoi(sc.Value); err == nil && ttlSeconds > 0 {
|
||||
settings.TokenTTL = time.Duration(ttlSeconds) * time.Second
|
||||
}
|
||||
}
|
||||
|
||||
return settings
|
||||
}
|
||||
|
||||
func (s *runtimeSettingsStore) ensureInvalidationListener() {
|
||||
s.listenerOnce.Do(startRuntimeSettingsInvalidationListener)
|
||||
}
|
||||
|
||||
func startRuntimeSettingsInvalidationListener() {
|
||||
if db.Redis == nil {
|
||||
return
|
||||
}
|
||||
|
||||
go func() {
|
||||
pubsub := db.Redis.Subscribe(context.Background(), repository.SystemConfigInvalidationChannel)
|
||||
defer func() {
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
for msg := range pubsub.Channel() {
|
||||
var payload struct {
|
||||
Key string `json:"key"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil {
|
||||
InvalidateRuntimeSettings()
|
||||
continue
|
||||
}
|
||||
if payload.Key == "" || payload.Key == "*" || IsRuntimeConfigKey(payload.Key) {
|
||||
InvalidateRuntimeSettings()
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
@@ -0,0 +1,119 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"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() = false, want true from seed defaults")
|
||||
}
|
||||
|
||||
if err := db.DB(ctx).Model(&model.SystemConfig{}).
|
||||
Where("key = ?", model.ConfigKeyCapLoginEnabled).
|
||||
Update("value", "false").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() = true, want false 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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
// 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)
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
)
|
||||
|
||||
func TestListVisibleSystemConfigsUsesRedisCache(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
||||
t.Fatalf("InvalidateVisibleSystemConfigsCache() error = %v", err)
|
||||
}
|
||||
|
||||
if _, err := repository.ListVisibleSystemConfigs(ctx); err != nil {
|
||||
t.Fatalf("ListVisibleSystemConfigs() warm cache error = %v", err)
|
||||
}
|
||||
|
||||
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, 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")
|
||||
}
|
||||
}
|
||||
|
||||
exists, err := db.Redis.Exists(ctx, db.PrefixedKey(repository.SystemConfigVisibleListRedisKey)).Result()
|
||||
if err != nil {
|
||||
t.Fatalf("Redis.Exists() error = %v", err)
|
||||
}
|
||||
if exists == 0 {
|
||||
t.Fatal("expected visible config list cache key to exist")
|
||||
}
|
||||
|
||||
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
||||
t.Fatalf("InvalidateVisibleSystemConfigsCache() error = %v", err)
|
||||
}
|
||||
|
||||
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")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,57 @@
|
||||
// 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/common/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))
|
||||
}
|
||||
@@ -0,0 +1,70 @@
|
||||
// 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/common/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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,108 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"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 != "OpenFlare" {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", warm.Value, "OpenFlare")
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
if err := db.HDel(ctx, repository.SystemConfigRedisHashKey, model.ConfigKeySiteName); err != nil {
|
||||
t.Fatalf("HDel(site_name) error = %v", err)
|
||||
}
|
||||
|
||||
cached, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name) cached error = %v", err)
|
||||
}
|
||||
if cached.Value != "OpenFlare" {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want stale RAM value %q", cached.Value, "OpenFlare")
|
||||
}
|
||||
|
||||
if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil {
|
||||
t.Fatalf("InvalidateSystemConfigCache(site_name) error = %v", err)
|
||||
}
|
||||
|
||||
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")
|
||||
}
|
||||
|
||||
exists, err := db.Redis.HExists(ctx, db.PrefixedKey(repository.SystemConfigRedisHashKey), model.ConfigKeySiteName).Result()
|
||||
if err != nil {
|
||||
t.Fatalf("HExists(site_name) error = %v", err)
|
||||
}
|
||||
if !exists {
|
||||
t.Fatal("expected redis hash field to be repopulated after refresh")
|
||||
}
|
||||
}
|
||||
|
||||
func TestInvalidateSystemConfigCacheClearsRedisField(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name) error = %v", err)
|
||||
}
|
||||
_ = sc
|
||||
|
||||
if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil {
|
||||
t.Fatalf("InvalidateSystemConfigCache(site_name) error = %v", err)
|
||||
}
|
||||
|
||||
_, err = db.Redis.HGet(ctx, db.PrefixedKey(repository.SystemConfigRedisHashKey), model.ConfigKeySiteName).Result()
|
||||
if !errors.Is(err, redis.Nil) {
|
||||
t.Fatalf("HGet(site_name) error = %v, want redis.Nil", err)
|
||||
}
|
||||
|
||||
if err := dbConn.Model(&model.SystemConfig{}).
|
||||
Where("key = ?", model.ConfigKeySiteName).
|
||||
Update("value", "after_invalidate").Error; err != nil {
|
||||
t.Fatalf("Update(site_name) error = %v", err)
|
||||
}
|
||||
|
||||
refreshed, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name) refreshed error = %v", err)
|
||||
}
|
||||
if refreshed.Value != "after_invalidate" {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", refreshed.Value, "after_invalidate")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,24 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package custom provides custom business handlers
|
||||
package custom
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
// Hello is a sample handler for custom business logic
|
||||
// @Summary Sample Hello API
|
||||
// @Description A sample business API for customization
|
||||
// @Tags custom
|
||||
// @Produce json
|
||||
// @Success 200 {object} response.Any{data=string} "成功"
|
||||
// @Router /api/v1/custom/hello [get]
|
||||
func Hello(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK("Hello from custom business module!"))
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package health 提供健康检查端点
|
||||
package health
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
// Health 健康检查
|
||||
// @Summary 健康检查
|
||||
// @Description 检查服务是否正常运行,可用于负载均衡存活探测
|
||||
// @Tags health
|
||||
// @Produce json
|
||||
// @Success 200 {object} response.Any{data=string} "服务正常"
|
||||
// @Router /api/health [get]
|
||||
func Health(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
@@ -0,0 +1,36 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package oauth 提供 OAuth/OIDC 认证与会话管理
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// LogForAudit 将登录鉴权审计日志写入 Logger
|
||||
func LogForAudit(ctx context.Context, user *model.User, c *gin.Context) {
|
||||
auditLog := loginRequiredAuditLog{
|
||||
UserID: user.ID,
|
||||
Username: user.Username,
|
||||
ClientIP: c.ClientIP(),
|
||||
Method: c.Request.Method,
|
||||
Path: c.Request.URL.Path,
|
||||
RequestURI: c.Request.RequestURI,
|
||||
UserAgent: c.Request.UserAgent(),
|
||||
Referer: c.Request.Referer(),
|
||||
}
|
||||
auditJSON, err := json.Marshal(auditLog)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "[LoginRequiredAudit] marshal failed: %v", err)
|
||||
logger.DebugF(ctx, "[LoginRequiredAudit] %s %d %s", c.ClientIP(), user.ID, user.Username)
|
||||
} else {
|
||||
logger.DebugF(ctx, "[LoginRequiredAudit] %s", auditJSON)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
func isOIDCLoginEnabled(ctx context.Context) bool {
|
||||
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
return enabled
|
||||
}
|
||||
|
||||
func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSource, error) {
|
||||
name := strings.TrimSpace(strings.ToLower(sourceName))
|
||||
if name == "" {
|
||||
sources, err := model.GetActiveAuthSources(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(sources) == 0 {
|
||||
return nil, errors.New(errNoActiveAuthSource)
|
||||
}
|
||||
return &sources[0], nil
|
||||
}
|
||||
return model.GetAuthSourceByName(ctx, name)
|
||||
}
|
||||
|
||||
func activeLoginSources(ctx context.Context) []AuthSourceView {
|
||||
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
|
||||
if err == nil && !enabled {
|
||||
return nil
|
||||
}
|
||||
|
||||
dbSources, err := model.GetActiveAuthSources(ctx)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
sources := make([]AuthSourceView, 0, len(dbSources))
|
||||
for _, source := range dbSources {
|
||||
sources = append(sources, AuthSourceView{
|
||||
ID: source.ID,
|
||||
Name: source.Name,
|
||||
Type: source.Type,
|
||||
DisplayName: source.DisplayName,
|
||||
IsActive: source.IsActive,
|
||||
IconURL: source.IconURL,
|
||||
ClientSecretConfigured: source.ClientSecretConfigured,
|
||||
})
|
||||
}
|
||||
return sources
|
||||
}
|
||||
|
||||
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
|
||||
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress)
|
||||
if err != nil || strings.TrimSpace(sc.Value) == "" {
|
||||
return "", errors.New(errServerAddressMissing)
|
||||
}
|
||||
return strings.TrimRight(sc.Value, "/") + "/login", nil
|
||||
}
|
||||
|
||||
func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
|
||||
if source == nil {
|
||||
return nil, nil, errors.New(errAuthSourceRequired)
|
||||
}
|
||||
|
||||
if source.OpenIDDiscoveryURL == "" {
|
||||
return nil, nil, errors.New(errDiscoveryURLRequired)
|
||||
}
|
||||
|
||||
// Clean the issuer URL (trim /.well-known/openid-configuration if configured by mistake)
|
||||
issuer := strings.TrimSuffix(strings.TrimSpace(source.OpenIDDiscoveryURL), "/")
|
||||
issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration")
|
||||
issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server")
|
||||
|
||||
// 使用进程级缓存获取 provider,避免每次调用都向 issuer 发起
|
||||
// /.well-known/openid-configuration HTTP 请求。
|
||||
provider, err := globalOIDCProviderCache.get(ctx, issuer)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
verifier := provider.Verifier(&oidc.Config{ClientID: source.ClientID})
|
||||
scopes := strings.Fields(source.Scopes)
|
||||
if len(scopes) == 0 {
|
||||
scopes = []string{oidc.ScopeOpenID, "profile", "email"}
|
||||
}
|
||||
if !containsScope(scopes, oidc.ScopeOpenID) {
|
||||
scopes = append([]string{oidc.ScopeOpenID}, scopes...)
|
||||
}
|
||||
|
||||
return &oauth2.Config{
|
||||
ClientID: source.ClientID,
|
||||
ClientSecret: source.ClientSecret,
|
||||
RedirectURL: redirectURL,
|
||||
Scopes: scopes,
|
||||
Endpoint: provider.Endpoint(),
|
||||
}, verifier, nil
|
||||
}
|
||||
|
||||
func containsScope(scopes []string, scope string) bool {
|
||||
for _, item := range scopes {
|
||||
if item == scope {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Session 用户信息字段 Key
|
||||
const (
|
||||
UserNameKey = "username"
|
||||
UserIDKey = "user_id"
|
||||
UserObjKey = "user_obj"
|
||||
TokenAuthKey = "token_auth" // 标记当前请求是否通过 Access Token 鉴权
|
||||
TokenAdminKey = "token_admin" // Access Token 本身是否具有管理员权限
|
||||
SessionTokenKey = "oauth_session_token" //nolint:gosec // false positive: this is a session key, not hardcoded credentials
|
||||
PasswordHashKey = "password_hash"
|
||||
)
|
||||
|
||||
// OAuth State 缓存 Key 格式与过期时间
|
||||
const (
|
||||
OAuthStateCacheKeyFormat = "oauth:state:%s"
|
||||
OAuthStateCacheKeyExpiration = 10 * time.Minute
|
||||
)
|
||||
|
||||
// OAuth 授权用途常量
|
||||
const (
|
||||
OAuthPurposeLogin = "login"
|
||||
OAuthPurposeBind = "bind"
|
||||
)
|
||||
|
||||
type oauthStatePayload struct {
|
||||
SourceName string `json:"source_name"`
|
||||
Purpose string `json:"purpose"`
|
||||
UserID uint64 `json:"user_id,omitempty"`
|
||||
SessionHash string `json:"session_hash"`
|
||||
}
|
||||
|
||||
func encodeOAuthStatePayload(payload oauthStatePayload) (string, error) {
|
||||
data, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(data), nil
|
||||
}
|
||||
|
||||
func decodeOAuthStatePayload(value string) (oauthStatePayload, error) {
|
||||
var payload oauthStatePayload
|
||||
if err := json.Unmarshal([]byte(value), &payload); err != nil {
|
||||
return oauthStatePayload{}, err
|
||||
}
|
||||
return payload, nil
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package oauth
|
||||
|
||||
// OAuth 认证相关错误消息
|
||||
const (
|
||||
errInvalidState = "非法登录请求"
|
||||
errIDTokenVerifyFailed = "ID Token 验证失败" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errIDTokenVerifyFailedFormat = "%s: %w"
|
||||
errNonceMismatch = "nonce 不匹配,可能存在重放攻击"
|
||||
errNoActiveAuthSource = "未配置可用认证源"
|
||||
errServerAddressMissing = "服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试"
|
||||
errAuthSourceRequired = "认证源不能为空"
|
||||
errDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
|
||||
errUsernameGenerateFailed = "无法生成可用用户名"
|
||||
errUsernameFromSourceFailed = "无法从认证源获取用户名"
|
||||
errAuthSourceDisabled = "认证源未启用"
|
||||
errInvalidExternalAccountBindingID = "绑定记录 ID 无效"
|
||||
ErrTokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
)
|
||||
@@ -0,0 +1,22 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package oauth
|
||||
|
||||
import "github.com/gin-gonic/gin"
|
||||
|
||||
// GetFromContext 从 Gin 请求上下文获取指定类型的值。
|
||||
func GetFromContext[T any](c *gin.Context, key string) (T, bool) {
|
||||
value, exists := c.Get(key)
|
||||
if !exists {
|
||||
var zero T
|
||||
return zero, false
|
||||
}
|
||||
typed, ok := value.(T)
|
||||
return typed, ok
|
||||
}
|
||||
|
||||
// SetToContext 设置值到 Gin 请求上下文。
|
||||
func SetToContext[T any](c *gin.Context, key string, value T) {
|
||||
c.Set(key, value)
|
||||
}
|
||||
@@ -0,0 +1,173 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
// GetLoginURL 获取登录授权地址
|
||||
// @Summary 获取登录授权地址
|
||||
// @Description 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用第一个启用的认证源。
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Param source query string false "认证源名称,为空使用第一个启用的认证源"
|
||||
// @Success 200 {object} response.Any{data=oauth.OAuthAuthorizeResponse} "授权 URL"
|
||||
// @Failure 400 {object} response.Any "认证源不存在或未配置"
|
||||
// @Failure 500 {object} response.Any "Redis 异常 or 构造 URL 失败"
|
||||
// @Router /api/v1/oauth/login [get]
|
||||
func GetLoginURL(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
if !isOIDCLoginEnabled(ctx) {
|
||||
response.AbortBadRequest(c, errAuthSourceDisabled)
|
||||
return
|
||||
}
|
||||
|
||||
source, err := resolveAuthSource(ctx, c.Query("source"))
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if !source.IsActive {
|
||||
response.AbortBadRequest(c, errAuthSourceDisabled)
|
||||
return
|
||||
}
|
||||
|
||||
session := sessions.Default(c)
|
||||
token, isNew := ensureSessionToken(session)
|
||||
if isNew {
|
||||
if err := session.Save(); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
userID := GetUserIDFromSession(session)
|
||||
sessionHash := hashSessionToken(token)
|
||||
|
||||
state := uuid.NewString()
|
||||
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
|
||||
SourceName: source.Name,
|
||||
Purpose: OAuthPurposeLogin,
|
||||
UserID: userID,
|
||||
SessionHash: sessionHash,
|
||||
})
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
|
||||
}
|
||||
|
||||
func buildAuthorizeURL(ctx context.Context, source *model.AuthSource, state string) (string, error) {
|
||||
redirectURL, err := getFrontendLoginRedirectURL(ctx)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if verifier != nil {
|
||||
return authConfig.AuthCodeURL(state, oidc.Nonce(state)), nil
|
||||
}
|
||||
return authConfig.AuthCodeURL(state), nil
|
||||
}
|
||||
|
||||
// Authorize 发起指定认证源授权
|
||||
// @Summary 发起指定认证源授权
|
||||
// @Description 根据指定认证源名称发起 OAuth 授权,支持 purpose 参数用于区分登录和账号绑定场景。认证源必须已启用。
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Param source path string true "认证源名称"
|
||||
// @Param purpose query string false "授权目的:login(登录)或 bind(绑定账号),默认 login"
|
||||
// @Success 200 {object} response.Any{data=oauth.OAuthAuthorizeResponse} "授权 URL"
|
||||
// @Failure 400 {object} response.Any "认证源不存在或未启用"
|
||||
// @Failure 500 {object} response.Any "Redis 异常或构造 URL 失败"
|
||||
// @Router /api/v1/oauth/{source}/authorize [get]
|
||||
func Authorize(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
if !isOIDCLoginEnabled(ctx) {
|
||||
response.AbortBadRequest(c, errAuthSourceDisabled)
|
||||
return
|
||||
}
|
||||
|
||||
source, err := resolveAuthSource(ctx, c.Param("source"))
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if !source.IsActive {
|
||||
response.AbortBadRequest(c, errAuthSourceDisabled)
|
||||
return
|
||||
}
|
||||
purpose := strings.ToLower(strings.TrimSpace(c.Query("purpose")))
|
||||
if purpose != OAuthPurposeBind {
|
||||
purpose = OAuthPurposeLogin
|
||||
}
|
||||
|
||||
session := sessions.Default(c)
|
||||
userID := GetUserIDFromSession(session)
|
||||
if purpose == OAuthPurposeBind && userID == 0 {
|
||||
response.AbortUnauthorized(c, common.UnAuthorized)
|
||||
return
|
||||
}
|
||||
|
||||
token, isNew := ensureSessionToken(session)
|
||||
if isNew {
|
||||
if err := session.Save(); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
sessionHash := hashSessionToken(token)
|
||||
|
||||
state := uuid.NewString()
|
||||
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
|
||||
SourceName: source.Name,
|
||||
Purpose: purpose,
|
||||
UserID: userID,
|
||||
SessionHash: sessionHash,
|
||||
})
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
|
||||
}
|
||||
@@ -0,0 +1,227 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/listener"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// Callback OAuth 回调处理
|
||||
// @Summary OAuth 回调处理
|
||||
// @Description 接收前端传回的 state 和 code,完成 OAuth/OIDC 认证并建立会话。支持登录(login)和账号绑定(bind)两种场景。
|
||||
// @Tags oauth
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body oauth.CallbackRequest true "回调请求参数"
|
||||
// @Success 200 {object} response.Any{data=oauth.OAuthCallbackResult} "登录或绑定成功"
|
||||
// @Failure 400 {object} response.Any "state 无效、参数错误或认证源错误"
|
||||
// @Failure 401 {object} response.Any "绑定场景未登录"
|
||||
// @Failure 500 {object} response.Any "OAuth 认证失败或内部错误"
|
||||
// @Router /api/v1/oauth/callback [post]
|
||||
func Callback(c *gin.Context) {
|
||||
var req CallbackRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
stateKey := db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State))
|
||||
payloadRaw, err := db.Redis.Get(ctx, stateKey).Result()
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, errInvalidState)
|
||||
return
|
||||
}
|
||||
_ = db.Redis.Del(ctx, stateKey)
|
||||
|
||||
payload, err := decodeOAuthStatePayload(payloadRaw)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
session := sessions.Default(c)
|
||||
currentUserID := GetUserIDFromSession(session)
|
||||
|
||||
if payload.Purpose == OAuthPurposeBind && currentUserID == 0 {
|
||||
response.AbortUnauthorized(c, common.UnAuthorized)
|
||||
return
|
||||
}
|
||||
|
||||
token, ok := session.Get(SessionTokenKey).(string)
|
||||
if !ok || token == "" {
|
||||
response.AbortBadRequest(c, "invalid session context")
|
||||
return
|
||||
}
|
||||
|
||||
if hashSessionToken(token) != payload.SessionHash {
|
||||
response.AbortBadRequest(c, "session mismatch for oauth state")
|
||||
return
|
||||
}
|
||||
|
||||
if payload.Purpose == OAuthPurposeBind && currentUserID != payload.UserID {
|
||||
response.AbortBadRequest(c, "user context mismatch for oauth binding")
|
||||
return
|
||||
}
|
||||
|
||||
if !isOIDCLoginEnabled(ctx) {
|
||||
response.AbortBadRequest(c, errAuthSourceDisabled)
|
||||
return
|
||||
}
|
||||
|
||||
source, err := resolveAuthSource(ctx, payload.SourceName)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if !source.IsActive {
|
||||
response.AbortBadRequest(c, errAuthSourceDisabled)
|
||||
return
|
||||
}
|
||||
|
||||
redirectURL, err := getFrontendLoginRedirectURL(ctx)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
userInfo, err := buildOAuthUserInfo(ctx, source, req.Code, req.State, redirectURL)
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
if err := normalizeOAuthUserInfo(userInfo); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
if userInfo.Sub == "" {
|
||||
userInfo.Sub = userInfo.Username
|
||||
}
|
||||
|
||||
if payload.Purpose == OAuthPurposeBind {
|
||||
handleCallbackBind(ctx, c, source, userInfo)
|
||||
return
|
||||
}
|
||||
|
||||
handleCallbackLogin(ctx, c, source, userInfo)
|
||||
}
|
||||
|
||||
// handleCallbackBind 处理 OAuth 回调中的帐号绑定流程
|
||||
func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) {
|
||||
userID := GetUserIDFromContext(c)
|
||||
if userID == 0 {
|
||||
response.AbortUnauthorized(c, common.UnAuthorized)
|
||||
return
|
||||
}
|
||||
var user model.User
|
||||
if err := db.DB(ctx).First(&user, "id = ?", userID).Error; err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
if err := model.BindExternalAccount(ctx, &model.ExternalAccount{
|
||||
AuthSourceID: source.ID,
|
||||
UserID: user.ID,
|
||||
ExternalID: userInfo.Sub,
|
||||
ExternalUsername: userInfo.Username,
|
||||
Email: userInfo.Email,
|
||||
}); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
user.LastLoginAt = time.Now()
|
||||
_ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error
|
||||
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "bound")))
|
||||
}
|
||||
|
||||
// handleCallbackLogin 处理 OAuth 回调中的登录流程(查找已有帐号或自动注册)
|
||||
func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) {
|
||||
var user model.User
|
||||
|
||||
account, err := model.FindExternalAccount(ctx, source.ID, userInfo.Sub)
|
||||
switch {
|
||||
case err == nil:
|
||||
if err := db.DB(ctx).First(&user, "id = ?", account.UserID).Error; err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
case errors.Is(err, gorm.ErrRecordNotFound):
|
||||
newUser, ok := handleCallbackRegister(ctx, c, source, userInfo)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
user = newUser
|
||||
default:
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
user.LastLoginAt = time.Now()
|
||||
_ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error
|
||||
if err := setLoginSession(ctx, c, &user); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
logger.InfoF(ctx, "[LoginAudit] successful OAuth login via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
|
||||
|
||||
listener.EmitAdminLoggedIn(ctx, &user, c.ClientIP())
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "logged_in")))
|
||||
}
|
||||
|
||||
// handleCallbackRegister 处理 OAuth 回调中的自动注册流程
|
||||
// 若注册被禁用则保存 pending 信息并返回 false;若注册成功则返回新用户;若出错则返回 false
|
||||
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) (model.User, bool) {
|
||||
registrationEnabled, regErr := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
|
||||
if regErr != nil {
|
||||
registrationEnabled = true
|
||||
}
|
||||
|
||||
if !registrationEnabled {
|
||||
c.JSON(http.StatusOK, response.OK(buildCallbackResult(nil, "need_bind")))
|
||||
return model.User{}, false
|
||||
}
|
||||
|
||||
username, uniqueErr := uniqueUsername(ctx, userInfo.Username)
|
||||
if uniqueErr != nil {
|
||||
response.AbortInternal(c, uniqueErr.Error())
|
||||
return model.User{}, false
|
||||
}
|
||||
userInfo.Username = username
|
||||
|
||||
var user model.User
|
||||
if err := user.CreateUser(ctx, db.DB(ctx), userInfo); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return model.User{}, false
|
||||
}
|
||||
if err := model.BindExternalAccount(ctx, &model.ExternalAccount{
|
||||
AuthSourceID: source.ID,
|
||||
UserID: user.ID,
|
||||
ExternalID: userInfo.Sub,
|
||||
ExternalUsername: userInfo.Username,
|
||||
Email: userInfo.Email,
|
||||
}); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return model.User{}, false
|
||||
}
|
||||
logger.InfoF(ctx, "[LoginAudit] successful OAuth registration via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
|
||||
|
||||
return user, true
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// ListExternalAccounts 获取当前用户的外部帐号绑定列表
|
||||
// @Summary 获取外部帐号列表
|
||||
// @Description 返回当前登录用户已绑定的所有外部 OAuth 帐号信息,需要登录
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]model.ExternalAccountView} "外部帐号列表"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/oauth/external-accounts [get]
|
||||
func ListExternalAccounts(c *gin.Context) {
|
||||
userID := GetUserIDFromContext(c)
|
||||
accounts, err := model.ListExternalAccountsByUserID(c.Request.Context(), userID)
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(accounts))
|
||||
}
|
||||
|
||||
// DeleteExternalAccount 解除外部帐号绑定
|
||||
// @Summary 解除外部帐号绑定
|
||||
// @Description 解除当前登录用户与指定外部帐号的绑定关系,需要登录
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path uint64 true "外部帐号绑定记录 ID"
|
||||
// @Success 200 {object} response.Any{data=string} "解除绑定成功"
|
||||
// @Failure 400 {object} response.Any "ID 无效或解除失败"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Router /api/v1/oauth/external-accounts/{id}/delete [post]
|
||||
func DeleteExternalAccount(c *gin.Context) {
|
||||
userID := GetUserIDFromContext(c)
|
||||
if userID == 0 {
|
||||
response.AbortUnauthorized(c, common.UnAuthorized)
|
||||
return
|
||||
}
|
||||
rawID := strings.TrimSpace(c.Param("id"))
|
||||
id, err := strconv.ParseUint(rawID, 10, 64)
|
||||
if err != nil || id == 0 {
|
||||
response.AbortBadRequest(c, errInvalidExternalAccountBindingID)
|
||||
return
|
||||
}
|
||||
if err := model.DeleteExternalAccountForUser(c.Request.Context(), id, userID); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// GetLoginSources 获取可用登录源列表
|
||||
// @Summary 获取可用登录源
|
||||
// @Description 返回当前系统已启用的所有 OAuth 登录源,前端展示登录按钮列表时调用
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Success 200 {object} response.Any{data=[]oauth.AuthSourceView} "登录源列表"
|
||||
// @Router /api/v1/oauth/sources [get]
|
||||
func GetLoginSources(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(activeLoginSources(c.Request.Context())))
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
|
||||
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type loginRequiredAuditLog struct {
|
||||
UserID uint64 `json:"user_id"`
|
||||
Username string `json:"username"`
|
||||
ClientIP string `json:"client_ip"`
|
||||
Method string `json:"method"`
|
||||
Path string `json:"path"`
|
||||
RequestURI string `json:"request_uri"`
|
||||
UserAgent string `json:"user_agent"`
|
||||
Referer string `json:"referer"`
|
||||
}
|
||||
|
||||
func getUserByToken(ctx context.Context, tokenStr string) (*model.User, *model.AccessToken, error) {
|
||||
tokenHash := model.HashToken(tokenStr)
|
||||
var tokenRecord model.AccessToken
|
||||
if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&tokenRecord).Error; err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
var user model.User
|
||||
if err := db.DB(ctx).Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&user).Error; err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
return &user, &tokenRecord, nil
|
||||
}
|
||||
|
||||
// GetUserFromRequest 校验 Access Token 或 Session 并返回用户对象,如果未登录或用户失效则返回 error
|
||||
func GetUserFromRequest(c *gin.Context) (*model.User, error) {
|
||||
ctx := c.Request.Context()
|
||||
|
||||
// check token in headers
|
||||
tokenStr := c.GetHeader("X-Access-Token")
|
||||
if tokenStr == "" {
|
||||
authHeader := c.GetHeader("Authorization")
|
||||
if len(authHeader) > 7 && authHeader[:7] == "Bearer " {
|
||||
tokenStr = authHeader[7:]
|
||||
}
|
||||
}
|
||||
|
||||
// 优先使用 Access Token 鉴权
|
||||
if tokenStr != "" {
|
||||
if user, tokenRecord, err := getUserByToken(ctx, tokenStr); err == nil {
|
||||
// 强行阻止 system 用户任何会话/Token 鉴权通过
|
||||
if user.Username == "system" {
|
||||
return nil, errors.New("system user is not allowed to login")
|
||||
}
|
||||
SetToContext(c, TokenAuthKey, true)
|
||||
SetToContext(c, TokenAdminKey, tokenRecord.IsAdmin)
|
||||
return user, nil
|
||||
}
|
||||
}
|
||||
|
||||
// 降级使用 Session 鉴权
|
||||
userID := GetUserIDFromContext(c)
|
||||
if userID <= 0 {
|
||||
return nil, errors.New("unauthorized")
|
||||
}
|
||||
|
||||
var user model.User
|
||||
// load user from db to make sure is active
|
||||
tx := db.DB(ctx).Where("id = ? AND is_active = ?", userID, true).First(&user)
|
||||
if tx.Error != nil {
|
||||
return nil, tx.Error
|
||||
}
|
||||
|
||||
// 密码哈希校验:当用户存在本地密码时,要求 Session 中的密码哈希必须与当前数据库中一致
|
||||
if user.Password != "" {
|
||||
session := sessions.Default(c)
|
||||
sessionHash, _ := session.Get(PasswordHashKey).(string)
|
||||
if sessionHash != user.Password {
|
||||
return nil, errors.New("session expired due to password change")
|
||||
}
|
||||
}
|
||||
|
||||
// set keys in context for session auth
|
||||
SetToContext(c, TokenAuthKey, false)
|
||||
SetToContext(c, TokenAdminKey, false)
|
||||
|
||||
// 强行阻止 system 用户任何会话/Token 鉴权通过
|
||||
if user.Username == "system" {
|
||||
return nil, errors.New("system user is not allowed to login")
|
||||
}
|
||||
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
// LoginRequired 返回登录鉴权中间件,校验 Access Token 或 Session
|
||||
func LoginRequired() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// init trace
|
||||
ctx, span := otel_trace.Start(c.Request.Context(), "LoginRequired")
|
||||
defer span.End()
|
||||
|
||||
user, err := GetUserFromRequest(c)
|
||||
if err != nil {
|
||||
response.AbortUnauthorized(c, common.UnAuthorized)
|
||||
return
|
||||
}
|
||||
|
||||
// log
|
||||
LogForAudit(ctx, user, c)
|
||||
|
||||
// set user info
|
||||
SetToContext(c, UserObjKey, user)
|
||||
|
||||
// next
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点
|
||||
func DisallowTokenAuth() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if tokenAuth, _ := GetFromContext[bool](c, TokenAuthKey); tokenAuth {
|
||||
response.AbortForbidden(c, ErrTokenAuthNotAllowed)
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,36 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package oauth
|
||||
|
||||
// AuthSourceView 登录源展示信息
|
||||
type AuthSourceView struct {
|
||||
ID uint64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
DisplayName string `json:"display_name"`
|
||||
IsActive bool `json:"is_active"`
|
||||
IconURL string `json:"icon_url"`
|
||||
ClientSecretConfigured bool `json:"client_secret_configured"`
|
||||
}
|
||||
|
||||
// OAuthAuthorizeResponse 授权 URL 响应
|
||||
//
|
||||
//nolint:revive // OAuth 前缀保持包内语义清晰
|
||||
type OAuthAuthorizeResponse struct {
|
||||
AuthorizeURL string `json:"authorize_url"`
|
||||
}
|
||||
|
||||
// OAuthCallbackResult 回调处理结果
|
||||
//
|
||||
//nolint:revive // OAuth 前缀保持包内语义清晰
|
||||
type OAuthCallbackResult struct {
|
||||
Status string `json:"status"`
|
||||
User *BasicUserInfo `json:"user,omitempty"`
|
||||
}
|
||||
|
||||
// CallbackRequest OAuth 回调请求参数
|
||||
type CallbackRequest struct {
|
||||
State string `json:"state" binding:"required"`
|
||||
Code string `json:"code" binding:"required"`
|
||||
}
|
||||
@@ -0,0 +1,141 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
func uniqueUsername(ctx context.Context, base string) (string, error) {
|
||||
base = strings.TrimSpace(base)
|
||||
if base == "" {
|
||||
base = "user"
|
||||
}
|
||||
|
||||
var existingUsernames []string
|
||||
if err := db.DB(ctx).Model(&model.User{}).
|
||||
Where("username = ? OR username LIKE ?", base, base+"-%").
|
||||
Pluck("username", &existingUsernames).Error; err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// 将现有的用户名放入 map 中,以便 O(1) 查找
|
||||
exists := make(map[string]bool, len(existingUsernames))
|
||||
for _, u := range existingUsernames {
|
||||
exists[strings.ToLower(u)] = true
|
||||
}
|
||||
|
||||
// 检查 base 是否被占用
|
||||
if !exists[strings.ToLower(base)] {
|
||||
return base, nil
|
||||
}
|
||||
|
||||
// 顺序查找第一个可用的带后缀用户名
|
||||
for i := 1; i <= 1000; i++ {
|
||||
candidate := fmt.Sprintf("%s-%d", base, i)
|
||||
if !exists[strings.ToLower(candidate)] {
|
||||
return candidate, nil
|
||||
}
|
||||
}
|
||||
|
||||
return "", errors.New(errUsernameGenerateFailed)
|
||||
}
|
||||
|
||||
func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code string, nonce string, redirectURL string) (*model.OAuthUserInfo, error) {
|
||||
authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
token, err := authConfig.Exchange(ctx, code)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
userInfo := &model.OAuthUserInfo{Active: true}
|
||||
if verifier != nil {
|
||||
if verifyErr := verifyIDToken(ctx, verifier, token, nonce, userInfo); verifyErr != nil {
|
||||
return nil, verifyErr
|
||||
}
|
||||
}
|
||||
|
||||
if userInfo.Username == "" && userInfo.PreferredUsername != "" {
|
||||
userInfo.Username = userInfo.PreferredUsername
|
||||
}
|
||||
if userInfo.Username == "" && userInfo.Email != "" {
|
||||
userInfo.Username = strings.Split(userInfo.Email, "@")[0]
|
||||
}
|
||||
if userInfo.Username == "" && userInfo.Sub != "" {
|
||||
userInfo.Username = userInfo.Sub
|
||||
}
|
||||
if userInfo.Name == "" {
|
||||
userInfo.Name = userInfo.Username
|
||||
}
|
||||
|
||||
return userInfo, nil
|
||||
}
|
||||
|
||||
// verifyIDToken 验证 OIDC ID Token 并将 Claims 解析到 userInfo
|
||||
func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *model.OAuthUserInfo) error {
|
||||
rawIDToken, ok := token.Extra("id_token").(string)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
idToken, verifyErr := verifier.Verify(ctx, rawIDToken)
|
||||
if verifyErr != nil {
|
||||
return fmt.Errorf(errIDTokenVerifyFailedFormat, errIDTokenVerifyFailed, verifyErr)
|
||||
}
|
||||
if nonce != "" && idToken.Nonce != nonce {
|
||||
return errors.New(errNonceMismatch)
|
||||
}
|
||||
if claimsErr := idToken.Claims(userInfo); claimsErr != nil {
|
||||
return claimsErr
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error {
|
||||
userInfo.Username = strings.TrimSpace(userInfo.Username)
|
||||
userInfo.PreferredUsername = strings.TrimSpace(userInfo.PreferredUsername)
|
||||
userInfo.Email = strings.TrimSpace(userInfo.Email)
|
||||
userInfo.Name = strings.TrimSpace(userInfo.Name)
|
||||
userInfo.AvatarURL = strings.TrimSpace(userInfo.AvatarURL)
|
||||
|
||||
if userInfo.Username == "" && userInfo.PreferredUsername != "" {
|
||||
userInfo.Username = userInfo.PreferredUsername
|
||||
}
|
||||
if userInfo.Username == "" && userInfo.Email != "" {
|
||||
userInfo.Username = strings.Split(userInfo.Email, "@")[0]
|
||||
}
|
||||
if userInfo.Username == "" && userInfo.Sub != "" {
|
||||
userInfo.Username = userInfo.Sub
|
||||
}
|
||||
if userInfo.Username == "" {
|
||||
return errors.New(errUsernameFromSourceFailed)
|
||||
}
|
||||
if userInfo.Name == "" {
|
||||
userInfo.Name = userInfo.Username
|
||||
}
|
||||
if !userInfo.Active {
|
||||
userInfo.Active = true
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildCallbackResult(user *model.User, status string) OAuthCallbackResult {
|
||||
result := OAuthCallbackResult{Status: status}
|
||||
if user != nil {
|
||||
info := BuildBasicUserInfo(user, false)
|
||||
result.User = &info
|
||||
}
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"sync"
|
||||
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"golang.org/x/oauth2"
|
||||
"golang.org/x/sync/singleflight"
|
||||
)
|
||||
|
||||
// oidcProviderCache 进程级 OIDC provider 缓存。
|
||||
//
|
||||
// oidc.NewProvider 每次调用都会向远端 issuer 的
|
||||
// /.well-known/openid-configuration 发起 HTTP 请求拉取元数据。
|
||||
// 由于 provider 元数据极少变动,将其缓存后可消除登录发起与回调时的
|
||||
// 重复外部 HTTP 往返。
|
||||
//
|
||||
// 并发安全性:
|
||||
// - mu + entries 防止并发读写 map。
|
||||
// - sfGroup 保证同一 issuer 同时只有一次在途的 NewProvider 调用
|
||||
// (singleflight),后续等待者复用同一结果,彻底消除 thundering herd。
|
||||
type oidcProviderCache struct {
|
||||
mu sync.RWMutex
|
||||
entries map[string]*oidc.Provider // key: normalized issuer URL
|
||||
sfGroup singleflight.Group
|
||||
}
|
||||
|
||||
// globalOIDCProviderCache 是包级单例缓存,与进程同生命周期。
|
||||
var globalOIDCProviderCache = &oidcProviderCache{
|
||||
entries: make(map[string]*oidc.Provider),
|
||||
}
|
||||
|
||||
// discoveryContext 从请求 ctx 提取 HTTP 客户端,并绑定到不可取消的 Background ctx。
|
||||
// 这样既能在测试中注入 mock 客户端,又避免请求取消导致 provider 拉取失败。
|
||||
func discoveryContext(ctx context.Context) context.Context {
|
||||
bg := context.Background()
|
||||
if client, ok := ctx.Value(oauth2.HTTPClient).(*http.Client); ok && client != nil {
|
||||
bg = oidc.ClientContext(bg, client)
|
||||
}
|
||||
return bg
|
||||
}
|
||||
|
||||
// get 返回缓存的 provider;若无则通过 oidc.NewProvider 获取并写入缓存。
|
||||
// 同一 issuer 并发调用时,singleflight 保证只有一次实际 HTTP 请求。
|
||||
func (c *oidcProviderCache) get(ctx context.Context, issuer string) (*oidc.Provider, error) {
|
||||
// 快路径:已有缓存则直接返回。
|
||||
c.mu.RLock()
|
||||
if p, ok := c.entries[issuer]; ok {
|
||||
c.mu.RUnlock()
|
||||
return p, nil
|
||||
}
|
||||
c.mu.RUnlock()
|
||||
|
||||
// 慢路径:通过 singleflight 合并并发的首次请求。
|
||||
discCtx := discoveryContext(ctx)
|
||||
v, err, _ := c.sfGroup.Do(issuer, func() (any, error) {
|
||||
// 双检:singleflight 内再次检查,前一个并发组可能已写入缓存。
|
||||
c.mu.RLock()
|
||||
if p, ok := c.entries[issuer]; ok {
|
||||
c.mu.RUnlock()
|
||||
return p, nil
|
||||
}
|
||||
c.mu.RUnlock()
|
||||
|
||||
p, err := oidc.NewProvider(discCtx, issuer)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
c.mu.Lock()
|
||||
c.entries[issuer] = p
|
||||
c.mu.Unlock()
|
||||
return p, nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return v.(*oidc.Provider), nil //nolint:forcetypeassert // singleflight value 由同函数写入,类型确定
|
||||
}
|
||||
|
||||
// invalidate 从缓存中移除指定 issuer 对应的 provider。
|
||||
// 在认证源的 Discovery URL 被修改时调用,强制下次请求重新拉取元数据。
|
||||
func (c *oidcProviderCache) invalidate(issuer string) {
|
||||
c.mu.Lock()
|
||||
delete(c.entries, issuer)
|
||||
c.mu.Unlock()
|
||||
}
|
||||
|
||||
// InvalidateOIDCProviderCache 从进程级缓存中清除指定 issuer 的 provider 条目。
|
||||
// 当管理员更新认证源的 Discovery URL 后调用,以确保下次登录时重新拉取最新元数据。
|
||||
// issuer 值应为去掉 /.well-known/openid-configuration 后缀的规范化 URL。
|
||||
func InvalidateOIDCProviderCache(issuer string) {
|
||||
globalOIDCProviderCache.invalidate(issuer)
|
||||
}
|
||||
@@ -0,0 +1,106 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
// BasicUserInfo 用户基本信息结构体
|
||||
type BasicUserInfo struct {
|
||||
ID uint64 `json:"id"`
|
||||
Username string `json:"username"`
|
||||
Nickname string `json:"nickname"`
|
||||
Email string `json:"email"`
|
||||
AvatarURL string `json:"avatar_url"`
|
||||
IsAdmin bool `json:"is_admin"`
|
||||
NeedChangePassword bool `json:"need_change_password"`
|
||||
Bio string `json:"bio"`
|
||||
Phone string `json:"phone"`
|
||||
Gender string `json:"gender"`
|
||||
Website string `json:"website"`
|
||||
Location string `json:"location"`
|
||||
}
|
||||
|
||||
// BuildBasicUserInfo 将 User 模型转换为 BasicUserInfo
|
||||
func BuildBasicUserInfo(user *model.User, needChange bool) BasicUserInfo {
|
||||
return BasicUserInfo{
|
||||
ID: user.ID,
|
||||
Username: user.Username,
|
||||
Nickname: user.Nickname,
|
||||
Email: user.Email,
|
||||
AvatarURL: user.AvatarURL,
|
||||
IsAdmin: user.IsAdmin,
|
||||
NeedChangePassword: needChange,
|
||||
Bio: user.Bio,
|
||||
Phone: user.Phone,
|
||||
Gender: user.Gender,
|
||||
Website: user.Website,
|
||||
Location: user.Location,
|
||||
}
|
||||
}
|
||||
|
||||
// UserInfo 获取当前登录用户信息
|
||||
// @Summary 获取当前登录用户信息
|
||||
// @Description 返回当前登录用户的基本信息及余额数据,需要登录。包括用户 ID、用户名、信任等级、各类余额信息等。
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=oauth.BasicUserInfo} "用户信息"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Router /api/v1/oauth/user-info [get]
|
||||
// @Router /api/v1/user-info [get]
|
||||
// @Router /api/v1/user/self [get]
|
||||
func UserInfo(c *gin.Context) {
|
||||
user, _ := GetFromContext[*model.User](c, UserObjKey)
|
||||
session := sessions.Default(c)
|
||||
needChange := session.Get("need_change_password") == true
|
||||
|
||||
c.JSON(
|
||||
http.StatusOK,
|
||||
response.OK(BuildBasicUserInfo(user, needChange)),
|
||||
)
|
||||
}
|
||||
|
||||
// GetLoginURL 获取登录地址
|
||||
// @Summary 获取登录地址
|
||||
// @Description 生成 OAuth 登录 URL,前端跳转至该地址完成授权。返回的 URL 中包含 state 参数用于 CSRF 防护。
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Success 200 {object} response.Any{data=string} "OAuth 登录 URL"
|
||||
// @Failure 500 {object} response.Any "Redis 异常或内部错误"
|
||||
// @Router /api/v1/oauth/login [get]
|
||||
|
||||
// Logout 退出登录
|
||||
// @Summary 退出登录
|
||||
// @Description 清除当前用户的登录会话,完成退出。清除 Cookie 中的 Session 数据。
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=string} "退出成功"
|
||||
// @Failure 500 {object} response.Any "Session 清除失败"
|
||||
// @Router /api/v1/oauth/logout [get]
|
||||
func Logout(c *gin.Context) {
|
||||
session := sessions.Default(c)
|
||||
userID := session.Get(UserIDKey)
|
||||
username := session.Get(UserNameKey)
|
||||
if userID != nil {
|
||||
logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP())
|
||||
}
|
||||
session.Options(GetSessionOptions(-1))
|
||||
session.Clear()
|
||||
if err := session.Save(); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package oauth provides authentication and OAuth integration.
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/gin-contrib/sessions"
|
||||
)
|
||||
|
||||
// GetSessionOptions 根据配置构建 Session 选项
|
||||
func GetSessionOptions(maxAge int) sessions.Options {
|
||||
return sessions.Options{
|
||||
Path: "/",
|
||||
Domain: config.Config.App.SessionDomain,
|
||||
MaxAge: maxAge,
|
||||
HttpOnly: config.Config.App.SessionHTTPOnly,
|
||||
Secure: config.Config.App.SessionSecure,
|
||||
SameSite: http.SameSiteLaxMode,
|
||||
}
|
||||
}
|
||||
|
||||
// StripCookieMaxAgeAndExpires 从 Set-Cookie 响应头中移除 Max-Age 和 Expires,从而使其成为浏览器会话 Cookie
|
||||
func StripCookieMaxAgeAndExpires(header http.Header, cookieName string) {
|
||||
headers := header["Set-Cookie"]
|
||||
if len(headers) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
newHeaders := make([]string, 0, len(headers))
|
||||
for _, h := range headers {
|
||||
if strings.HasPrefix(h, cookieName+"=") {
|
||||
parts := strings.Split(h, ";")
|
||||
newParts := make([]string, 0, len(parts))
|
||||
for _, p := range parts {
|
||||
trimmed := strings.TrimSpace(p)
|
||||
lower := strings.ToLower(trimmed)
|
||||
if strings.HasPrefix(lower, "max-age=") || strings.HasPrefix(lower, "expires=") {
|
||||
continue
|
||||
}
|
||||
newParts = append(newParts, p)
|
||||
}
|
||||
newHeaders = append(newHeaders, strings.Join(newParts, ";"))
|
||||
} else {
|
||||
newHeaders = append(newHeaders, h)
|
||||
}
|
||||
}
|
||||
header["Set-Cookie"] = newHeaders
|
||||
}
|
||||
@@ -0,0 +1,83 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
// GetUserIDFromSession 从 Session 中提取用户 ID
|
||||
func GetUserIDFromSession(s sessions.Session) uint64 {
|
||||
userID, ok := s.Get(UserIDKey).(uint64)
|
||||
if !ok {
|
||||
return 0
|
||||
}
|
||||
return userID
|
||||
}
|
||||
|
||||
// GetUserIDFromContext 从 Gin Context 的 Session 中提取用户 ID
|
||||
func GetUserIDFromContext(c *gin.Context) uint64 {
|
||||
session := sessions.Default(c)
|
||||
return GetUserIDFromSession(session)
|
||||
}
|
||||
|
||||
func ensureSessionToken(s sessions.Session) (string, bool) {
|
||||
token, ok := s.Get(SessionTokenKey).(string)
|
||||
if !ok || token == "" {
|
||||
token = uuid.NewString()
|
||||
s.Set(SessionTokenKey, token)
|
||||
return token, true
|
||||
}
|
||||
return token, false
|
||||
}
|
||||
|
||||
func hashSessionToken(token string) string {
|
||||
h := sha256.New()
|
||||
h.Write([]byte(token))
|
||||
return hex.EncodeToString(h.Sum(nil))
|
||||
}
|
||||
|
||||
func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) error {
|
||||
session := sessions.Default(c)
|
||||
session.Set(UserIDKey, user.ID)
|
||||
session.Set(UserNameKey, user.Username)
|
||||
session.Set(PasswordHashKey, user.Password)
|
||||
|
||||
// 根据系统配置动态设置 Session 过期时间
|
||||
maxAge := config.Config.App.SessionAge
|
||||
isSessionCookie := false
|
||||
|
||||
ttlHours, err := repository.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours)
|
||||
if err == nil {
|
||||
switch {
|
||||
case ttlHours == -1:
|
||||
// 永不过期,设置为 10 年
|
||||
maxAge = 10 * 365 * 24 * 3600
|
||||
case ttlHours > 0:
|
||||
maxAge = ttlHours * 3600
|
||||
case ttlHours == 0:
|
||||
isSessionCookie = true
|
||||
}
|
||||
}
|
||||
session.Options(GetSessionOptions(maxAge))
|
||||
|
||||
if err := session.Save(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if isSessionCookie {
|
||||
StripCookieMaxAgeAndExpires(c.Writer.Header(), config.Config.App.SessionCookieName)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,95 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"net"
|
||||
"strings"
|
||||
|
||||
pkggeoip "github.com/rain-kl/openflare/pkg/geoip"
|
||||
)
|
||||
|
||||
var accessLogGeoProviderFactory = func() (pkggeoip.GeoIPService, error) {
|
||||
return pkggeoip.NewMaxMindGeoIPService()
|
||||
}
|
||||
|
||||
type accessLogRegionResolver struct {
|
||||
provider pkggeoip.GeoIPService
|
||||
cache map[string]string
|
||||
}
|
||||
|
||||
func newAccessLogRegionResolver() (*accessLogRegionResolver, error) {
|
||||
provider, err := accessLogGeoProviderFactory()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &accessLogRegionResolver{
|
||||
provider: provider,
|
||||
cache: make(map[string]string),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (r *accessLogRegionResolver) Close() {
|
||||
if r == nil || r.provider == nil {
|
||||
return
|
||||
}
|
||||
if err := r.provider.Close(); err != nil {
|
||||
slog.Warn("close access log geo provider failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *accessLogRegionResolver) Resolve(rawIP string) string {
|
||||
if r == nil || r.provider == nil {
|
||||
return ""
|
||||
}
|
||||
normalizedIP := normalizeAccessLogIP(rawIP)
|
||||
if normalizedIP == "" {
|
||||
return ""
|
||||
}
|
||||
if cached, ok := r.cache[normalizedIP]; ok {
|
||||
return cached
|
||||
}
|
||||
|
||||
info, err := r.provider.GetGeoInfo(net.ParseIP(normalizedIP))
|
||||
if err != nil || info == nil {
|
||||
r.cache[normalizedIP] = ""
|
||||
return ""
|
||||
}
|
||||
|
||||
region := strings.TrimSpace(info.Name)
|
||||
if region == "" {
|
||||
region = strings.TrimSpace(info.ISOCode)
|
||||
}
|
||||
r.cache[normalizedIP] = region
|
||||
return region
|
||||
}
|
||||
|
||||
func normalizeAccessLogIP(raw string) string {
|
||||
trimmed := strings.TrimSpace(raw)
|
||||
if trimmed == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
if ip := net.ParseIP(trimmed); ip != nil {
|
||||
return ip.String()
|
||||
}
|
||||
|
||||
trimmed = strings.TrimPrefix(trimmed, "[")
|
||||
trimmed = strings.TrimSuffix(trimmed, "]")
|
||||
if ip := net.ParseIP(trimmed); ip != nil {
|
||||
return ip.String()
|
||||
}
|
||||
|
||||
host, _, err := net.SplitHostPort(strings.TrimSpace(raw))
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
host = strings.TrimPrefix(host, "[")
|
||||
host = strings.TrimSuffix(host, "]")
|
||||
if ip := net.ParseIP(host); ip != nil {
|
||||
return ip.String()
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,143 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
agentTokenPositiveCacheTTL = 2 * time.Minute
|
||||
agentTokenNegativeCacheTTL = 10 * time.Minute
|
||||
)
|
||||
|
||||
type cachedAgentNode struct {
|
||||
node *model.OpenFlareNode
|
||||
expiresAt time.Time
|
||||
}
|
||||
|
||||
type accessTokenAuthCache struct {
|
||||
mu sync.RWMutex
|
||||
positive map[string]cachedAgentNode
|
||||
negative map[string]time.Time
|
||||
now func() time.Time
|
||||
loadNodeByToken func(context.Context, string) (*model.OpenFlareNode, error)
|
||||
}
|
||||
|
||||
var tokenCache = newAccessTokenAuthCache()
|
||||
|
||||
func newAccessTokenAuthCache() *accessTokenAuthCache {
|
||||
return &accessTokenAuthCache{
|
||||
positive: make(map[string]cachedAgentNode),
|
||||
negative: make(map[string]time.Time),
|
||||
now: time.Now,
|
||||
loadNodeByToken: func(ctx context.Context, token string) (*model.OpenFlareNode, error) {
|
||||
return model.GetOpenFlareNodeByAccessToken(ctx, token)
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (c *accessTokenAuthCache) authenticate(ctx context.Context, token string) (*model.OpenFlareNode, error) {
|
||||
now := c.now()
|
||||
if node, ok := c.getNode(token, now); ok {
|
||||
return node, nil
|
||||
}
|
||||
if c.isMissing(token, now) {
|
||||
return nil, gorm.ErrRecordNotFound
|
||||
}
|
||||
|
||||
node, err := c.loadNodeByToken(ctx, token)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
c.storeMissing(token, now.Add(agentTokenNegativeCacheTTL))
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
c.storeNode(token, node)
|
||||
return cloneNode(node), nil
|
||||
}
|
||||
|
||||
func (c *accessTokenAuthCache) getNode(token string, now time.Time) (*model.OpenFlareNode, bool) {
|
||||
c.mu.RLock()
|
||||
entry, ok := c.positive[token]
|
||||
c.mu.RUnlock()
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
if now.After(entry.expiresAt) {
|
||||
c.mu.Lock()
|
||||
delete(c.positive, token)
|
||||
c.mu.Unlock()
|
||||
return nil, false
|
||||
}
|
||||
return cloneNode(entry.node), true
|
||||
}
|
||||
|
||||
func (c *accessTokenAuthCache) isMissing(token string, now time.Time) bool {
|
||||
c.mu.RLock()
|
||||
expiresAt, ok := c.negative[token]
|
||||
c.mu.RUnlock()
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
if now.After(expiresAt) {
|
||||
c.mu.Lock()
|
||||
delete(c.negative, token)
|
||||
c.mu.Unlock()
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (c *accessTokenAuthCache) storeNode(token string, node *model.OpenFlareNode) {
|
||||
if token == "" || node == nil {
|
||||
return
|
||||
}
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
delete(c.negative, token)
|
||||
c.positive[token] = cachedAgentNode{
|
||||
node: cloneNode(node),
|
||||
expiresAt: c.now().Add(agentTokenPositiveCacheTTL),
|
||||
}
|
||||
}
|
||||
|
||||
func (c *accessTokenAuthCache) storeMissing(token string, expiresAt time.Time) {
|
||||
if token == "" {
|
||||
return
|
||||
}
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
delete(c.positive, token)
|
||||
c.negative[token] = expiresAt
|
||||
}
|
||||
|
||||
func (c *accessTokenAuthCache) reset() {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.positive = make(map[string]cachedAgentNode)
|
||||
c.negative = make(map[string]time.Time)
|
||||
}
|
||||
|
||||
// ResetAuthCacheForTest clears the in-memory access token cache for integration tests.
|
||||
func ResetAuthCacheForTest() {
|
||||
tokenCache.reset()
|
||||
}
|
||||
|
||||
// AuthenticateAccessToken validates X-Agent-Token against of_nodes.access_token.
|
||||
func AuthenticateAccessToken(ctx context.Context, token string) (*model.OpenFlareNode, error) {
|
||||
token = strings.TrimSpace(token)
|
||||
if token == "" {
|
||||
return nil, errors.New(errMissingAgentToken)
|
||||
}
|
||||
return tokenCache.authenticate(ctx, token)
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
openrestyrender "github.com/rain-kl/openflare/pkg/render/openresty"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type configVersionRecord struct {
|
||||
ID uint `gorm:"primaryKey"`
|
||||
Version string `gorm:"column:version"`
|
||||
SnapshotJSON string `gorm:"column:snapshot_json"`
|
||||
SupportFilesJSON string `gorm:"column:support_files_json"`
|
||||
Checksum string `gorm:"column:checksum"`
|
||||
IsActive bool `gorm:"column:is_active"`
|
||||
CreatedAt time.Time `gorm:"column:created_at"`
|
||||
}
|
||||
|
||||
func (configVersionRecord) TableName() string {
|
||||
return "of_config_versions"
|
||||
}
|
||||
|
||||
func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) {
|
||||
version, err := loadActiveConfigVersion(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &ActiveConfigMeta{
|
||||
Version: version.Version,
|
||||
Checksum: version.Checksum,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func getActiveConfigForAgent(ctx context.Context) (*ConfigResponse, error) {
|
||||
version, err := loadActiveConfigVersion(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var supportFiles []SupportFile
|
||||
if strings.TrimSpace(version.SupportFilesJSON) != "" {
|
||||
if err = json.Unmarshal([]byte(version.SupportFilesJSON), &supportFiles); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
return &ConfigResponse{
|
||||
Version: version.Version,
|
||||
Checksum: version.Checksum,
|
||||
SourceConfigJSON: version.SnapshotJSON,
|
||||
SupportFiles: sourceSupportFiles(supportFiles),
|
||||
CreatedAt: version.CreatedAt,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func loadActiveConfigVersion(ctx context.Context) (*configVersionRecord, error) {
|
||||
conn := db.DB(ctx)
|
||||
if conn == nil {
|
||||
return nil, errors.New("database not initialized")
|
||||
}
|
||||
version := &configVersionRecord{}
|
||||
err := conn.Where("is_active = ?", true).Order("id desc").First(version).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return version, nil
|
||||
}
|
||||
|
||||
func sourceSupportFiles(files []SupportFile) []SupportFile {
|
||||
if len(files) == 0 {
|
||||
return nil
|
||||
}
|
||||
result := make([]SupportFile, 0, len(files))
|
||||
for _, file := range files {
|
||||
if isRuntimeGeneratedSupportFile(file.Path) {
|
||||
continue
|
||||
}
|
||||
result = append(result, file)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func isRuntimeGeneratedSupportFile(path string) bool {
|
||||
switch strings.TrimSpace(path) {
|
||||
case "pow_config.json", "waf_config.json", openrestyrender.SourceConfigFileName:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func isActiveConfigNotFound(err error) bool {
|
||||
return errors.Is(err, gorm.ErrRecordNotFound)
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
openrestyrender "github.com/rain-kl/openflare/pkg/render/openresty"
|
||||
)
|
||||
|
||||
func TestIsRuntimeGeneratedSupportFile(t *testing.T) {
|
||||
tests := []struct {
|
||||
path string
|
||||
want bool
|
||||
}{
|
||||
{path: "pow_config.json", want: true},
|
||||
{path: "waf_config.json", want: true},
|
||||
{path: openrestyrender.SourceConfigFileName, want: true},
|
||||
{path: "runtime/custom.json", want: false},
|
||||
{path: "certs/example.pem", want: false},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
if got := isRuntimeGeneratedSupportFile(tc.path); got != tc.want {
|
||||
t.Fatalf("isRuntimeGeneratedSupportFile(%q) = %v, want %v", tc.path, got, tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSourceSupportFilesFiltersRuntimeGeneratedFiles(t *testing.T) {
|
||||
files := []SupportFile{
|
||||
{Path: "certs/example.pem", Content: "pem"},
|
||||
{Path: "pow_config.json", Content: "{}"},
|
||||
{Path: "waf_config.json", Content: "{}"},
|
||||
{Path: openrestyrender.SourceConfigFileName, Content: "{}"},
|
||||
{Path: "routes/extra.json", Content: "{}"},
|
||||
}
|
||||
|
||||
filtered := sourceSupportFiles(files)
|
||||
if len(filtered) != 2 {
|
||||
t.Fatalf("expected 2 support files, got %d: %+v", len(filtered), filtered)
|
||||
}
|
||||
if filtered[0].Path != "certs/example.pem" || filtered[1].Path != "routes/extra.json" {
|
||||
t.Fatalf("unexpected filtered files: %+v", filtered)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
const (
|
||||
errMissingAgentToken = "缺少 Agent Token"
|
||||
errInvalidAgentToken = "无权进行此操作,Agent Token 无效"
|
||||
errInvalidDiscoveryToken = "无权进行此操作,注册 Token 无效"
|
||||
errNodeMissingFromContext = "Node object missing from context"
|
||||
errNoActiveConfig = "当前没有激活版本"
|
||||
errNodeNotFound = "节点不存在"
|
||||
errNodeIDRequired = "node_id 不能为空"
|
||||
errVersionRequired = "version 不能为空"
|
||||
errInvalidApplyResult = "result 仅支持 success、warning 或 failed"
|
||||
errIPRequired = "ip 不能为空"
|
||||
errIPInvalid = "ip 格式无效"
|
||||
errAgentVersionRequired = "version 不能为空"
|
||||
errNodeIDConflict = "节点标识生成冲突,请重试"
|
||||
)
|
||||
@@ -0,0 +1,308 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"net"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
ofgeoip "github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const (
|
||||
openrestyStatusHealthy = "healthy"
|
||||
openrestyStatusUnhealthy = "unhealthy"
|
||||
openrestyStatusUnknown = "unknown"
|
||||
releaseChannelStable = "stable"
|
||||
)
|
||||
|
||||
func newRandomToken() (string, error) {
|
||||
buf := make([]byte, 16)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(buf), nil
|
||||
}
|
||||
|
||||
func newServerNodeID() (string, error) {
|
||||
token, err := newRandomToken()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return "node-" + token, nil
|
||||
}
|
||||
|
||||
func normalizeOpenrestyStatus(status string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(status)) {
|
||||
case openrestyStatusHealthy:
|
||||
return openrestyStatusHealthy
|
||||
case openrestyStatusUnhealthy:
|
||||
return openrestyStatusUnhealthy
|
||||
default:
|
||||
return openrestyStatusUnknown
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeNodePayload(payload NodePayload) NodePayload {
|
||||
payload.Name = strings.TrimSpace(payload.Name)
|
||||
payload.IP = strings.TrimSpace(payload.IP)
|
||||
payload.Version = strings.TrimSpace(payload.Version)
|
||||
payload.ExtVersion = strings.TrimSpace(payload.ExtVersion)
|
||||
payload.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
|
||||
payload.LastError = truncateForDatabase(payload.LastError, 16000)
|
||||
payload.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
|
||||
payload.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, 16000)
|
||||
return payload
|
||||
}
|
||||
|
||||
func validateNodePayload(payload NodePayload) error {
|
||||
if payload.IP == "" {
|
||||
return errPayload(errIPRequired)
|
||||
}
|
||||
if net.ParseIP(payload.IP) == nil {
|
||||
return errPayload(errIPInvalid)
|
||||
}
|
||||
if payload.Version == "" {
|
||||
return errPayload(errAgentVersionRequired)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type payloadError string
|
||||
|
||||
func (e payloadError) Error() string { return string(e) }
|
||||
|
||||
func errPayload(message string) error { return payloadError(message) }
|
||||
|
||||
func applyNodeRuntime(node *model.OpenFlareNode, payload NodePayload, preserveName bool) {
|
||||
if !preserveName || strings.TrimSpace(node.Name) == "" {
|
||||
if strings.TrimSpace(payload.Name) != "" {
|
||||
node.Name = strings.TrimSpace(payload.Name)
|
||||
}
|
||||
}
|
||||
if !node.IPManualOverride {
|
||||
node.IP = strings.TrimSpace(payload.IP)
|
||||
}
|
||||
node.Version = strings.TrimSpace(payload.Version)
|
||||
node.ExtVersion = strings.TrimSpace(payload.ExtVersion)
|
||||
node.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
|
||||
node.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, 16000)
|
||||
node.Status = nodeStatusOnline
|
||||
node.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
|
||||
now := time.Now()
|
||||
node.LastSeenAt = &now
|
||||
node.LastError = truncateForDatabase(payload.LastError, 16000)
|
||||
if !node.GeoManualOverride {
|
||||
applyGeoInfoFromIP(node, node.IP)
|
||||
}
|
||||
}
|
||||
|
||||
func applyGeoInfoFromIP(node *model.OpenFlareNode, rawIP string) {
|
||||
if node == nil {
|
||||
return
|
||||
}
|
||||
node.GeoName = ""
|
||||
node.GeoLatitude = nil
|
||||
node.GeoLongitude = nil
|
||||
ip := net.ParseIP(strings.TrimSpace(rawIP))
|
||||
if ip == nil {
|
||||
return
|
||||
}
|
||||
info, err := ofgeoip.GeoInfoFromIP(ip)
|
||||
if err != nil || info == nil {
|
||||
return
|
||||
}
|
||||
if strings.TrimSpace(info.Name) != "" {
|
||||
node.GeoName = strings.TrimSpace(info.Name)
|
||||
}
|
||||
if info.Latitude != nil && info.Longitude != nil {
|
||||
node.GeoLatitude = cloneCoordinate(info.Latitude)
|
||||
node.GeoLongitude = cloneCoordinate(info.Longitude)
|
||||
}
|
||||
}
|
||||
|
||||
func cloneCoordinate(value *float64) *float64 {
|
||||
if value == nil {
|
||||
return nil
|
||||
}
|
||||
cloned := *value
|
||||
return &cloned
|
||||
}
|
||||
|
||||
func truncateForDatabase(value string, max int) string {
|
||||
if max <= 0 {
|
||||
return ""
|
||||
}
|
||||
runes := []rune(strings.TrimSpace(value))
|
||||
if len(runes) <= max {
|
||||
return string(runes)
|
||||
}
|
||||
return string(runes[:max])
|
||||
}
|
||||
|
||||
func resolveReportedNodeIP(reportedIP string, remoteAddr string) string {
|
||||
reported := normalizeIP(reportedIP)
|
||||
remote := normalizeRemoteAddr(remoteAddr)
|
||||
if reported == "" {
|
||||
return remote
|
||||
}
|
||||
if isPublicNodeIP(reported) {
|
||||
return reported
|
||||
}
|
||||
if isPublicNodeIP(remote) {
|
||||
return remote
|
||||
}
|
||||
return reported
|
||||
}
|
||||
|
||||
func normalizeIP(raw string) string {
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" {
|
||||
return ""
|
||||
}
|
||||
host := raw
|
||||
if strings.Contains(raw, ":") {
|
||||
if h, _, err := net.SplitHostPort(raw); err == nil {
|
||||
host = h
|
||||
}
|
||||
}
|
||||
host = strings.TrimPrefix(host, "[")
|
||||
host = strings.TrimSuffix(host, "]")
|
||||
if ip := net.ParseIP(host); ip != nil {
|
||||
return ip.String()
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func normalizeRemoteAddr(remoteAddr string) string {
|
||||
remoteAddr = strings.TrimSpace(remoteAddr)
|
||||
if remoteAddr == "" {
|
||||
return ""
|
||||
}
|
||||
host, _, err := net.SplitHostPort(remoteAddr)
|
||||
if err != nil {
|
||||
return normalizeIP(remoteAddr)
|
||||
}
|
||||
return normalizeIP(host)
|
||||
}
|
||||
|
||||
func isPublicNodeIP(raw string) bool {
|
||||
ip := net.ParseIP(strings.TrimSpace(raw))
|
||||
if ip == nil {
|
||||
return false
|
||||
}
|
||||
if ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsUnspecified() {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func buildAgentSettings(node *model.OpenFlareNode, updateNow bool, updateChannel string, updateTag string, restartOpenrestyNow bool) *Settings {
|
||||
autoUpdate := false
|
||||
if node != nil {
|
||||
autoUpdate = node.AutoUpdateEnabled
|
||||
}
|
||||
if strings.TrimSpace(updateChannel) == "" {
|
||||
updateChannel = releaseChannelStable
|
||||
}
|
||||
return &Settings{
|
||||
HeartbeatInterval: model.AgentHeartbeatInterval,
|
||||
WebsocketUpgradeEnabled: model.AgentWebsocketUpgradeEnabled,
|
||||
AutoUpdate: autoUpdate,
|
||||
UpdateRepo: model.AgentUpdateRepo,
|
||||
UpdateNow: updateNow,
|
||||
UpdateChannel: updateChannel,
|
||||
UpdateTag: strings.TrimSpace(updateTag),
|
||||
RestartOpenrestyNow: restartOpenrestyNow,
|
||||
}
|
||||
}
|
||||
|
||||
func collectHeartbeatChanges(previous *model.OpenFlareNode, current *model.OpenFlareNode) map[string]any {
|
||||
if previous == nil || current == nil {
|
||||
return map[string]any{}
|
||||
}
|
||||
changes := make(map[string]any)
|
||||
appendIfChanged := func(key string, before any, after any) {
|
||||
if before != after {
|
||||
changes[key] = after
|
||||
}
|
||||
}
|
||||
appendIfChanged("name", previous.Name, current.Name)
|
||||
appendIfChanged("ip", previous.IP, current.IP)
|
||||
appendIfChanged("geo_name", previous.GeoName, current.GeoName)
|
||||
appendIfChanged("version", previous.Version, current.Version)
|
||||
appendIfChanged("ext_version", previous.ExtVersion, current.ExtVersion)
|
||||
appendIfChanged("openresty_status", previous.OpenrestyStatus, current.OpenrestyStatus)
|
||||
appendIfChanged("openresty_message", previous.OpenrestyMessage, current.OpenrestyMessage)
|
||||
appendIfChanged("status", previous.Status, current.Status)
|
||||
appendIfChanged("current_version", previous.CurrentVersion, current.CurrentVersion)
|
||||
appendIfChanged("last_error", previous.LastError, current.LastError)
|
||||
appendIfChanged("update_requested", previous.UpdateRequested, current.UpdateRequested)
|
||||
appendIfChanged("update_channel", previous.UpdateChannel, current.UpdateChannel)
|
||||
appendIfChanged("update_tag", previous.UpdateTag, current.UpdateTag)
|
||||
appendIfChanged("restart_openresty_requested", previous.RestartOpenrestyRequested, current.RestartOpenrestyRequested)
|
||||
if !coordinatesEqual(previous.GeoLatitude, current.GeoLatitude) {
|
||||
changes["geo_latitude"] = current.GeoLatitude
|
||||
}
|
||||
if !coordinatesEqual(previous.GeoLongitude, current.GeoLongitude) {
|
||||
changes["geo_longitude"] = current.GeoLongitude
|
||||
}
|
||||
if !lastSeenAtEqual(previous.LastSeenAt, current.LastSeenAt) {
|
||||
changes["last_seen_at"] = current.LastSeenAt
|
||||
}
|
||||
return changes
|
||||
}
|
||||
|
||||
func coordinatesEqual(before *float64, after *float64) bool {
|
||||
if before == nil || after == nil {
|
||||
return before == after
|
||||
}
|
||||
return *before == *after
|
||||
}
|
||||
|
||||
func lastSeenAtEqual(before *time.Time, after *time.Time) bool {
|
||||
if before == nil || after == nil {
|
||||
return before == after
|
||||
}
|
||||
return before.Equal(*after)
|
||||
}
|
||||
|
||||
func normalizeApplyLogPayload(payload ApplyLogPayload) ApplyLogPayload {
|
||||
payload.NodeID = strings.TrimSpace(payload.NodeID)
|
||||
payload.Version = strings.TrimSpace(payload.Version)
|
||||
payload.Result = strings.ToLower(strings.TrimSpace(payload.Result))
|
||||
payload.Message = truncateForDatabase(strings.TrimSpace(payload.Message), 16000)
|
||||
payload.Checksum = strings.TrimSpace(payload.Checksum)
|
||||
payload.MainConfigChecksum = strings.TrimSpace(payload.MainConfigChecksum)
|
||||
payload.RouteConfigChecksum = strings.TrimSpace(payload.RouteConfigChecksum)
|
||||
return payload
|
||||
}
|
||||
|
||||
func isUniqueConstraintError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(strings.ToLower(err.Error()), "unique")
|
||||
}
|
||||
|
||||
// RefreshAccessTokenCache updates the in-memory node cache after heartbeat mutations.
|
||||
func RefreshAccessTokenCache(ctx context.Context, node *model.OpenFlareNode) {
|
||||
if node == nil {
|
||||
return
|
||||
}
|
||||
tokenCache.storeNode(node.AccessToken, cloneNode(node))
|
||||
}
|
||||
|
||||
func cloneNode(node *model.OpenFlareNode) *model.OpenFlareNode {
|
||||
if node == nil {
|
||||
return nil
|
||||
}
|
||||
cloned := *node
|
||||
return &cloned
|
||||
}
|
||||
@@ -0,0 +1,131 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
pkggeoip "github.com/rain-kl/openflare/pkg/geoip"
|
||||
)
|
||||
|
||||
type fakeGeoIPProvider struct {
|
||||
info *pkggeoip.GeoInfo
|
||||
}
|
||||
|
||||
func (f *fakeGeoIPProvider) Name() string { return "fake-geoip" }
|
||||
|
||||
func (f *fakeGeoIPProvider) GetGeoInfo(ip net.IP) (*pkggeoip.GeoInfo, error) {
|
||||
return f.info, nil
|
||||
}
|
||||
|
||||
func (f *fakeGeoIPProvider) UpdateDatabase() error { return nil }
|
||||
|
||||
func (f *fakeGeoIPProvider) Close() error { return nil }
|
||||
|
||||
func withFakeGeoIPProvider(t *testing.T, info *pkggeoip.GeoInfo) {
|
||||
t.Helper()
|
||||
previous := pkggeoip.CurrentProvider
|
||||
pkggeoip.CurrentProvider = &fakeGeoIPProvider{info: info}
|
||||
t.Cleanup(func() {
|
||||
pkggeoip.CurrentProvider = previous
|
||||
})
|
||||
}
|
||||
|
||||
func geoipFloat(value float64) *float64 {
|
||||
return &value
|
||||
}
|
||||
|
||||
func TestApplyGeoInfoFromIP(t *testing.T) {
|
||||
latitude := 31.2304
|
||||
longitude := 121.4737
|
||||
withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{
|
||||
Name: "Shanghai",
|
||||
Latitude: geoipFloat(latitude),
|
||||
Longitude: geoipFloat(longitude),
|
||||
})
|
||||
|
||||
node := &model.OpenFlareNode{IP: "203.0.113.10"}
|
||||
applyGeoInfoFromIP(node, node.IP)
|
||||
|
||||
if node.GeoName != "Shanghai" {
|
||||
t.Fatalf("expected geo_name Shanghai, got %q", node.GeoName)
|
||||
}
|
||||
if node.GeoLatitude == nil || *node.GeoLatitude != latitude {
|
||||
t.Fatalf("unexpected geo_latitude: %+v", node.GeoLatitude)
|
||||
}
|
||||
if node.GeoLongitude == nil || *node.GeoLongitude != longitude {
|
||||
t.Fatalf("unexpected geo_longitude: %+v", node.GeoLongitude)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyGeoInfoFromIPSkipsInvalidIP(t *testing.T) {
|
||||
withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{Name: "Should Not Apply"})
|
||||
|
||||
node := &model.OpenFlareNode{
|
||||
IP: "203.0.113.10",
|
||||
GeoName: "Existing",
|
||||
GeoLatitude: geoipFloat(1),
|
||||
GeoLongitude: geoipFloat(2),
|
||||
}
|
||||
applyGeoInfoFromIP(node, "not-an-ip")
|
||||
|
||||
if node.GeoName != "" || node.GeoLatitude != nil || node.GeoLongitude != nil {
|
||||
t.Fatalf("expected geo fields to be cleared on invalid IP, got %+v", node)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyNodeRuntimeRespectsGeoManualOverride(t *testing.T) {
|
||||
withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{
|
||||
Name: "Shanghai",
|
||||
Latitude: geoipFloat(31.2304),
|
||||
Longitude: geoipFloat(121.4737),
|
||||
})
|
||||
|
||||
node := &model.OpenFlareNode{
|
||||
GeoManualOverride: true,
|
||||
GeoName: "Manual",
|
||||
GeoLatitude: geoipFloat(10),
|
||||
GeoLongitude: geoipFloat(20),
|
||||
}
|
||||
applyNodeRuntime(node, NodePayload{
|
||||
IP: "203.0.113.10",
|
||||
Version: "1.0.0",
|
||||
}, true)
|
||||
|
||||
if node.GeoName != "Manual" {
|
||||
t.Fatalf("expected manual geo_name to be preserved, got %q", node.GeoName)
|
||||
}
|
||||
if node.GeoLatitude == nil || *node.GeoLatitude != 10 {
|
||||
t.Fatalf("expected manual geo_latitude to be preserved, got %+v", node.GeoLatitude)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectHeartbeatChangesTracksGeoFields(t *testing.T) {
|
||||
before := &model.OpenFlareNode{
|
||||
IP: "10.0.0.1",
|
||||
GeoName: "Old Region",
|
||||
}
|
||||
after := &model.OpenFlareNode{
|
||||
IP: "203.0.113.10",
|
||||
GeoName: "New Region",
|
||||
GeoLatitude: geoipFloat(31.2304),
|
||||
GeoLongitude: geoipFloat(121.4737),
|
||||
}
|
||||
|
||||
changes := collectHeartbeatChanges(before, after)
|
||||
if changes["ip"] != after.IP {
|
||||
t.Fatalf("expected ip change, got %+v", changes)
|
||||
}
|
||||
if changes["geo_name"] != after.GeoName {
|
||||
t.Fatalf("expected geo_name change, got %+v", changes)
|
||||
}
|
||||
if changes["geo_latitude"] != after.GeoLatitude {
|
||||
t.Fatalf("expected geo_latitude change, got %+v", changes)
|
||||
}
|
||||
if changes["geo_longitude"] != after.GeoLongitude {
|
||||
t.Fatalf("expected geo_longitude change, got %+v", changes)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,224 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/node"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// RegisterWithAccessToken registers an agent on a reserved node token.
|
||||
func RegisterWithAccessToken(ctx context.Context, authNode *model.OpenFlareNode, payload NodePayload) (*RegistrationResponse, error) {
|
||||
payload = normalizeNodePayload(payload)
|
||||
if authNode == nil {
|
||||
return nil, errors.New(errNodeNotFound)
|
||||
}
|
||||
if err := validateNodePayload(payload); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
applyNodeRuntime(authNode, payload, true)
|
||||
if err := model.SaveOpenFlareNode(ctx, authNode); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
RefreshAccessTokenCache(ctx, authNode)
|
||||
return &RegistrationResponse{
|
||||
NodeID: authNode.NodeID,
|
||||
AccessToken: authNode.AccessToken,
|
||||
Name: authNode.Name,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// RegisterWithDiscovery registers a new node using the global discovery token.
|
||||
func RegisterWithDiscovery(ctx context.Context, payload NodePayload) (*RegistrationResponse, error) {
|
||||
payload = normalizeNodePayload(payload)
|
||||
if err := validateNodePayload(payload); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
nodeID, err := newServerNodeID()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
accessToken, err := newRandomToken()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
nodeName := payload.Name
|
||||
if nodeName == "" {
|
||||
nodeName = nodeID
|
||||
}
|
||||
|
||||
record := &model.OpenFlareNode{
|
||||
NodeID: nodeID,
|
||||
Name: nodeName,
|
||||
AccessToken: accessToken,
|
||||
Status: nodeStatusOnline,
|
||||
NodeType: "edge_node",
|
||||
CapabilitiesJSON: "[]",
|
||||
UpdateChannel: releaseChannelStable,
|
||||
}
|
||||
applyNodeRuntime(record, payload, false)
|
||||
|
||||
if err = model.CreateOpenFlareNode(ctx, record); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New(errNodeIDConflict)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
RefreshAccessTokenCache(ctx, record)
|
||||
return &RegistrationResponse{
|
||||
NodeID: record.NodeID,
|
||||
AccessToken: record.AccessToken,
|
||||
Name: record.Name,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// HeartbeatNode updates runtime state and returns agent settings.
|
||||
func HeartbeatNode(ctx context.Context, authNode *model.OpenFlareNode, payload NodePayload) (*HeartbeatResponse, error) {
|
||||
if authNode == nil {
|
||||
return nil, errors.New(errNodeNotFound)
|
||||
}
|
||||
payload.NodeID = authNode.NodeID
|
||||
payload = normalizeNodePayload(payload)
|
||||
if err := validateNodePayload(payload); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
previous := *authNode
|
||||
updateNow := authNode.UpdateRequested
|
||||
restartOpenrestyNow := authNode.RestartOpenrestyRequested
|
||||
updateChannel := strings.TrimSpace(authNode.UpdateChannel)
|
||||
updateTag := strings.TrimSpace(authNode.UpdateTag)
|
||||
|
||||
applyNodeRuntime(authNode, payload, true)
|
||||
authNode.UpdateRequested = false
|
||||
authNode.UpdateChannel = releaseChannelStable
|
||||
authNode.UpdateTag = ""
|
||||
authNode.RestartOpenrestyRequested = false
|
||||
|
||||
changes := collectHeartbeatChanges(&previous, authNode)
|
||||
if len(changes) > 0 {
|
||||
fields := make([]string, 0, len(changes))
|
||||
for field := range changes {
|
||||
fields = append(fields, field)
|
||||
}
|
||||
if err := model.UpdateOpenFlareNodeFields(ctx, authNode, fields...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
RefreshAccessTokenCache(ctx, authNode)
|
||||
|
||||
reportedAt := time.Now()
|
||||
if authNode.LastSeenAt != nil {
|
||||
reportedAt = *authNode.LastSeenAt
|
||||
}
|
||||
PersistHeartbeatObservability(ctx, authNode.NodeID, payload, reportedAt)
|
||||
|
||||
activeConfig, err := getActiveConfigMeta(ctx)
|
||||
if err != nil && !isActiveConfigNotFound(err) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
wafIPGroups, err := ChangedWAFIPGroupsForAgent(ctx, nil, payload.WAFIPGroupChecksums)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &HeartbeatResponse{
|
||||
Node: authNode,
|
||||
AgentSettings: buildAgentSettings(authNode, updateNow, updateChannel, updateTag, restartOpenrestyNow),
|
||||
ActiveConfig: activeConfig,
|
||||
WAFIPGroups: wafIPGroups,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetActiveConfig returns the active configuration for an agent.
|
||||
func GetActiveConfig(ctx context.Context) (*ConfigResponse, error) {
|
||||
config, err := getActiveConfigForAgent(ctx)
|
||||
if err != nil {
|
||||
if isActiveConfigNotFound(err) {
|
||||
return nil, errors.New(errNoActiveConfig)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return config, nil
|
||||
}
|
||||
|
||||
// SyncWAFIPGroups returns WAF IP groups whose checksums differ from the agent state.
|
||||
func SyncWAFIPGroups(ctx context.Context, input WAFIPGroupSyncInput) (*WAFIPGroupSyncResult, error) {
|
||||
groups, err := ChangedWAFIPGroupsForAgent(ctx, input.IDs, input.Checksums)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &WAFIPGroupSyncResult{Groups: groups}, nil
|
||||
}
|
||||
|
||||
// ReportApplyLog records an agent apply result.
|
||||
func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFlareApplyLog, error) {
|
||||
now := time.Now()
|
||||
payload = normalizeApplyLogPayload(payload)
|
||||
if payload.NodeID == "" {
|
||||
return nil, errors.New(errNodeIDRequired)
|
||||
}
|
||||
if payload.Version == "" {
|
||||
return nil, errors.New(errVersionRequired)
|
||||
}
|
||||
if payload.Result != applyResultOK && payload.Result != applyResultWarn && payload.Result != applyResultFailed {
|
||||
return nil, errors.New(errInvalidApplyResult)
|
||||
}
|
||||
|
||||
log := &model.OpenFlareApplyLog{
|
||||
NodeID: payload.NodeID,
|
||||
Version: payload.Version,
|
||||
Result: payload.Result,
|
||||
Message: payload.Message,
|
||||
Checksum: payload.Checksum,
|
||||
MainConfigChecksum: payload.MainConfigChecksum,
|
||||
RouteConfigChecksum: payload.RouteConfigChecksum,
|
||||
SupportFileCount: payload.SupportFileCount,
|
||||
CreatedAt: now,
|
||||
}
|
||||
|
||||
conn := db.DB(ctx)
|
||||
if conn == nil {
|
||||
return nil, errors.New("database not initialized")
|
||||
}
|
||||
|
||||
err := conn.Transaction(func(tx *gorm.DB) error {
|
||||
record := &model.OpenFlareNode{}
|
||||
if err := tx.Where("node_id = ?", payload.NodeID).First(record).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
record.Status = nodeStatusOnline
|
||||
record.LastSeenAt = &now
|
||||
if payload.Result == applyResultOK {
|
||||
record.CurrentVersion = payload.Version
|
||||
record.LastError = ""
|
||||
} else {
|
||||
record.LastError = payload.Message
|
||||
}
|
||||
if err := tx.Create(log).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(record).Select("status", "last_seen_at", "current_version", "last_error").Updates(record).Error
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return log, nil
|
||||
}
|
||||
|
||||
// ValidateDiscoveryToken delegates to the node package discovery token helper.
|
||||
func ValidateDiscoveryToken(ctx context.Context, token string) error {
|
||||
return node.ValidateDiscoveryToken(ctx, token)
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
const (
|
||||
agentTokenHeader = "X-Agent-Token"
|
||||
agentNodeContextKey = "agent_node"
|
||||
)
|
||||
|
||||
// AgentAuth validates X-Agent-Token against of_nodes.access_token.
|
||||
func AgentAuth() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
token := strings.TrimSpace(c.GetHeader(agentTokenHeader))
|
||||
node, err := AuthenticateAccessToken(c.Request.Context(), token)
|
||||
if err != nil {
|
||||
response.AbortUnauthorized(c, errInvalidAgentToken)
|
||||
return
|
||||
}
|
||||
c.Set(agentNodeContextKey, node)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// AgentRegisterAuth accepts either a node access token or the global discovery token.
|
||||
func AgentRegisterAuth() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
token := strings.TrimSpace(c.GetHeader(agentTokenHeader))
|
||||
if node, err := AuthenticateAccessToken(c.Request.Context(), token); err == nil {
|
||||
c.Set(agentNodeContextKey, node)
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
if err := ValidateDiscoveryToken(c.Request.Context(), token); err != nil {
|
||||
response.AbortUnauthorized(c, errInvalidDiscoveryToken)
|
||||
return
|
||||
}
|
||||
c.Set("discovery_enabled", true)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// AgentNodeFromContext returns the authenticated agent node.
|
||||
func AgentNodeFromContext(c *gin.Context) (*model.OpenFlareNode, bool) {
|
||||
value, ok := c.Get(agentNodeContextKey)
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
node, ok := value.(*model.OpenFlareNode)
|
||||
return node, ok
|
||||
}
|
||||
@@ -0,0 +1,199 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/option"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupAgentAuthTestDB(t *testing.T) func() {
|
||||
t.Helper()
|
||||
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, sqliteDB.AutoMigrate(
|
||||
&model.OpenFlareNode{},
|
||||
&model.OpenFlareOption{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
option.ResetInitializationForTest()
|
||||
tokenCache.reset()
|
||||
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
option.ResetInitializationForTest()
|
||||
tokenCache.reset()
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthenticateAccessToken(t *testing.T) {
|
||||
cleanup := setupAgentAuthTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
now := time.Now()
|
||||
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
|
||||
NodeID: "node-auth-1",
|
||||
Name: "edge",
|
||||
AccessToken: "valid-agent-token",
|
||||
Status: nodeStatusOnline,
|
||||
LastSeenAt: &now,
|
||||
NodeType: "edge_node",
|
||||
}).Error)
|
||||
|
||||
t.Run("valid token", func(t *testing.T) {
|
||||
node, err := AuthenticateAccessToken(ctx, "valid-agent-token")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "node-auth-1", node.NodeID)
|
||||
})
|
||||
|
||||
t.Run("cached token", func(t *testing.T) {
|
||||
originalLoader := tokenCache.loadNodeByToken
|
||||
t.Cleanup(func() {
|
||||
tokenCache.loadNodeByToken = originalLoader
|
||||
})
|
||||
tokenCache.loadNodeByToken = func(context.Context, string) (*model.OpenFlareNode, error) {
|
||||
t.Fatal("db should not be queried for cached token")
|
||||
return nil, nil
|
||||
}
|
||||
node, err := AuthenticateAccessToken(ctx, "valid-agent-token")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "node-auth-1", node.NodeID)
|
||||
})
|
||||
|
||||
t.Run("missing token", func(t *testing.T) {
|
||||
_, err := AuthenticateAccessToken(ctx, "")
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), errMissingAgentToken)
|
||||
})
|
||||
|
||||
t.Run("invalid token", func(t *testing.T) {
|
||||
_, err := AuthenticateAccessToken(ctx, "invalid-token")
|
||||
require.Error(t, err)
|
||||
})
|
||||
}
|
||||
|
||||
func TestAgentAuthMiddleware(t *testing.T) {
|
||||
cleanup := setupAgentAuthTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
now := time.Now()
|
||||
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
|
||||
NodeID: "node-mw-1",
|
||||
Name: "edge",
|
||||
AccessToken: "middleware-token",
|
||||
Status: nodeStatusOnline,
|
||||
LastSeenAt: &now,
|
||||
NodeType: "edge_node",
|
||||
}).Error)
|
||||
|
||||
router := testhelper.NewTestGinEngine()
|
||||
router.GET("/protected", AgentAuth(), func(c *gin.Context) {
|
||||
node, ok := AgentNodeFromContext(c)
|
||||
if !ok {
|
||||
c.Status(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(gin.H{"node_id": node.NodeID}))
|
||||
})
|
||||
|
||||
t.Run("authorized request", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
|
||||
req.Header.Set(agentTokenHeader, "middleware-token")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, resp.Code)
|
||||
var apiResp response.Any
|
||||
require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &apiResp))
|
||||
assert.Empty(t, apiResp.ErrorMsg)
|
||||
})
|
||||
|
||||
t.Run("unauthorized request", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
|
||||
req.Header.Set(agentTokenHeader, "bad-token")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assert.Equal(t, http.StatusUnauthorized, resp.Code)
|
||||
})
|
||||
}
|
||||
|
||||
func TestAgentRegisterAuthMiddleware(t *testing.T) {
|
||||
cleanup := setupAgentAuthTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
now := time.Now()
|
||||
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
|
||||
NodeID: "node-register-1",
|
||||
Name: "edge",
|
||||
AccessToken: "existing-node-token",
|
||||
Status: nodeStatusOnline,
|
||||
LastSeenAt: &now,
|
||||
NodeType: "edge_node",
|
||||
}).Error)
|
||||
require.NoError(t, model.UpdateOpenFlareOption(ctx, "AgentDiscoveryToken", "discovery-token"))
|
||||
|
||||
router := testhelper.NewTestGinEngine()
|
||||
router.POST("/register", AgentRegisterAuth(), func(c *gin.Context) {
|
||||
if node, ok := AgentNodeFromContext(c); ok {
|
||||
c.JSON(http.StatusOK, response.OK(gin.H{"mode": "node", "node_id": node.NodeID}))
|
||||
return
|
||||
}
|
||||
if _, ok := c.Get("discovery_enabled"); ok {
|
||||
c.JSON(http.StatusOK, response.OK(gin.H{"mode": "discovery"}))
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusInternalServerError)
|
||||
})
|
||||
|
||||
t.Run("existing node token", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/register", nil)
|
||||
req.Header.Set(agentTokenHeader, "existing-node-token")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, resp.Code)
|
||||
var apiResp response.Any
|
||||
require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &apiResp))
|
||||
data, ok := apiResp.Data.(map[string]any)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "node", data["mode"])
|
||||
})
|
||||
|
||||
t.Run("discovery token", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/register", nil)
|
||||
req.Header.Set(agentTokenHeader, "discovery-token")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, resp.Code)
|
||||
var apiResp response.Any
|
||||
require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &apiResp))
|
||||
data, ok := apiResp.Data.(map[string]any)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "discovery", data["mode"])
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,503 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"go.uber.org/zap"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
const (
|
||||
healthEventStatusActive = "active"
|
||||
healthEventStatusResolved = "resolved"
|
||||
healthSeverityInfo = "info"
|
||||
healthSeverityWarning = "warning"
|
||||
healthSeverityCritical = "critical"
|
||||
nodeAccessLogRetentionDays = 90
|
||||
nodeAccessLogRetentionWindow = nodeAccessLogRetentionDays * 24 * time.Hour
|
||||
accessLogPathMaxLength = 100
|
||||
)
|
||||
|
||||
// NodeSystemProfile is the agent-reported system profile.
|
||||
type NodeSystemProfile struct {
|
||||
Hostname string `json:"hostname"`
|
||||
OSName string `json:"os_name"`
|
||||
OSVersion string `json:"os_version"`
|
||||
KernelVersion string `json:"kernel_version"`
|
||||
Architecture string `json:"architecture"`
|
||||
CPUModel string `json:"cpu_model"`
|
||||
CPUCores int `json:"cpu_cores"`
|
||||
TotalMemoryBytes int64 `json:"total_memory_bytes"`
|
||||
TotalDiskBytes int64 `json:"total_disk_bytes"`
|
||||
UptimeSeconds int64 `json:"uptime_seconds"`
|
||||
ReportedAtUnix int64 `json:"reported_at_unix"`
|
||||
}
|
||||
|
||||
// NodeMetricSnapshot is the agent-reported capacity snapshot.
|
||||
type NodeMetricSnapshot struct {
|
||||
CapturedAtUnix int64 `json:"captured_at_unix"`
|
||||
CPUUsagePercent float64 `json:"cpu_usage_percent"`
|
||||
MemoryUsedBytes int64 `json:"memory_used_bytes"`
|
||||
MemoryTotalBytes int64 `json:"memory_total_bytes"`
|
||||
StorageUsedBytes int64 `json:"storage_used_bytes"`
|
||||
StorageTotalBytes int64 `json:"storage_total_bytes"`
|
||||
DiskReadBytes int64 `json:"disk_read_bytes"`
|
||||
DiskWriteBytes int64 `json:"disk_write_bytes"`
|
||||
NetworkRxBytes int64 `json:"network_rx_bytes"`
|
||||
NetworkTxBytes int64 `json:"network_tx_bytes"`
|
||||
}
|
||||
|
||||
// NodeOpenrestyObservation is the agent-reported openresty network observation.
|
||||
type NodeOpenrestyObservation struct {
|
||||
CapturedAtUnix int64 `json:"captured_at_unix"`
|
||||
OpenrestyRxBytes int64 `json:"openresty_rx_bytes"`
|
||||
OpenrestyTxBytes int64 `json:"openresty_tx_bytes"`
|
||||
OpenrestyConnections int64 `json:"openresty_connections"`
|
||||
}
|
||||
|
||||
// NodeTrafficReport is the agent-reported traffic window.
|
||||
type NodeTrafficReport struct {
|
||||
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
|
||||
WindowEndedAtUnix int64 `json:"window_ended_at_unix"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
ErrorCount int64 `json:"error_count"`
|
||||
UniqueVisitorCount int64 `json:"unique_visitor_count"`
|
||||
StatusCodes map[string]int64 `json:"status_codes"`
|
||||
TopDomains map[string]int64 `json:"top_domains"`
|
||||
SourceCountries map[string]int64 `json:"source_countries"`
|
||||
}
|
||||
|
||||
// NodeAccessLog is a single access log row from the agent.
|
||||
type NodeAccessLog struct {
|
||||
LoggedAtUnix int64 `json:"logged_at_unix"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
Host string `json:"host"`
|
||||
Path string `json:"path"`
|
||||
StatusCode int `json:"status_code"`
|
||||
}
|
||||
|
||||
// BufferedObservabilityRecord is a buffered observability window from the agent.
|
||||
type BufferedObservabilityRecord struct {
|
||||
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
|
||||
Snapshot *NodeMetricSnapshot `json:"snapshot,omitempty"`
|
||||
OpenrestyObservation *NodeOpenrestyObservation `json:"openresty_observation,omitempty"`
|
||||
TrafficReport *NodeTrafficReport `json:"traffic_report,omitempty"`
|
||||
AccessLogs []NodeAccessLog `json:"access_logs,omitempty"`
|
||||
}
|
||||
|
||||
// NodeHealthEvent is an agent-reported health event.
|
||||
type NodeHealthEvent struct {
|
||||
EventType string `json:"event_type"`
|
||||
Severity string `json:"severity"`
|
||||
Message string `json:"message"`
|
||||
TriggeredAtUnix int64 `json:"triggered_at_unix"`
|
||||
Metadata map[string]string `json:"metadata"`
|
||||
}
|
||||
|
||||
// PersistHeartbeatObservability stores profile, snapshots, traffic, access logs, and health events.
|
||||
func PersistHeartbeatObservability(ctx context.Context, nodeID string, payload NodePayload, reportedAt time.Time) {
|
||||
if strings.TrimSpace(nodeID) == "" {
|
||||
return
|
||||
}
|
||||
if payload.Profile == nil &&
|
||||
payload.Snapshot == nil &&
|
||||
payload.TrafficReport == nil &&
|
||||
len(payload.AccessLogs) == 0 &&
|
||||
len(payload.BufferedObservability) == 0 &&
|
||||
payload.HealthEvents == nil {
|
||||
return
|
||||
}
|
||||
|
||||
conn := db.DB(ctx)
|
||||
if conn == nil {
|
||||
return
|
||||
}
|
||||
|
||||
if err := conn.Transaction(func(tx *gorm.DB) error {
|
||||
if err := persistNodeSystemProfile(tx, nodeID, payload.Profile, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistBufferedObservability(tx, nodeID, payload.BufferedObservability, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeMetricSnapshot(tx, nodeID, payload.Snapshot, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeOpenrestyObservation(tx, nodeID, payload.OpenrestyObservation, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeTrafficReport(tx, nodeID, payload.TrafficReport, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeAccessLogs(tx, nodeID, payload.AccessLogs, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if payload.HealthEvents != nil {
|
||||
if err := reconcileNodeHealthEvents(tx, nodeID, payload.HealthEvents, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
zap.L().Error("persist heartbeat observability failed", zap.String("node_id", nodeID), zap.Error(err))
|
||||
}
|
||||
}
|
||||
|
||||
func persistBufferedObservability(tx *gorm.DB, nodeID string, records []BufferedObservabilityRecord, reportedAt time.Time) error {
|
||||
for _, record := range records {
|
||||
if err := persistNodeMetricSnapshot(tx, nodeID, record.Snapshot, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeOpenrestyObservation(tx, nodeID, record.OpenrestyObservation, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeTrafficReport(tx, nodeID, record.TrafficReport, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeAccessLogs(tx, nodeID, record.AccessLogs, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func persistNodeSystemProfile(tx *gorm.DB, nodeID string, profile *NodeSystemProfile, reportedAt time.Time) error {
|
||||
if profile == nil {
|
||||
return nil
|
||||
}
|
||||
record := &model.OpenFlareNodeSystemProfile{
|
||||
NodeID: nodeID,
|
||||
Hostname: strings.TrimSpace(profile.Hostname),
|
||||
OSName: strings.TrimSpace(profile.OSName),
|
||||
OSVersion: strings.TrimSpace(profile.OSVersion),
|
||||
KernelVersion: strings.TrimSpace(profile.KernelVersion),
|
||||
Architecture: strings.TrimSpace(profile.Architecture),
|
||||
CPUModel: strings.TrimSpace(profile.CPUModel),
|
||||
CPUCores: profile.CPUCores,
|
||||
TotalMemoryBytes: profile.TotalMemoryBytes,
|
||||
TotalDiskBytes: profile.TotalDiskBytes,
|
||||
UptimeSeconds: profile.UptimeSeconds,
|
||||
ReportedAt: timeFromUnix(profile.ReportedAtUnix, reportedAt),
|
||||
}
|
||||
return tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "node_id"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{
|
||||
"hostname",
|
||||
"os_name",
|
||||
"os_version",
|
||||
"kernel_version",
|
||||
"architecture",
|
||||
"cpu_model",
|
||||
"cpu_cores",
|
||||
"total_memory_bytes",
|
||||
"total_disk_bytes",
|
||||
"uptime_seconds",
|
||||
"reported_at",
|
||||
"updated_at",
|
||||
}),
|
||||
}).Create(record).Error
|
||||
}
|
||||
|
||||
func persistNodeMetricSnapshot(tx *gorm.DB, nodeID string, snapshot *NodeMetricSnapshot, reportedAt time.Time) error {
|
||||
if snapshot == nil {
|
||||
return nil
|
||||
}
|
||||
record := &model.OpenFlareMetricSnapshot{
|
||||
NodeID: nodeID,
|
||||
CapturedAt: timeFromUnix(snapshot.CapturedAtUnix, reportedAt),
|
||||
CPUUsagePercent: snapshot.CPUUsagePercent,
|
||||
MemoryUsedBytes: snapshot.MemoryUsedBytes,
|
||||
MemoryTotalBytes: snapshot.MemoryTotalBytes,
|
||||
StorageUsedBytes: snapshot.StorageUsedBytes,
|
||||
StorageTotalBytes: snapshot.StorageTotalBytes,
|
||||
DiskReadBytes: snapshot.DiskReadBytes,
|
||||
DiskWriteBytes: snapshot.DiskWriteBytes,
|
||||
NetworkRxBytes: snapshot.NetworkRxBytes,
|
||||
NetworkTxBytes: snapshot.NetworkTxBytes,
|
||||
}
|
||||
exists, err := metricSnapshotExists(tx, nodeID, record.CapturedAt)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
return nil
|
||||
}
|
||||
return tx.Create(record).Error
|
||||
}
|
||||
|
||||
func persistNodeOpenrestyObservation(tx *gorm.DB, nodeID string, obs *NodeOpenrestyObservation, reportedAt time.Time) error {
|
||||
if obs == nil {
|
||||
return nil
|
||||
}
|
||||
record := &model.OpenFlareNodeObservationOpenresty{
|
||||
NodeID: nodeID,
|
||||
CapturedAt: timeFromUnix(obs.CapturedAtUnix, reportedAt),
|
||||
OpenrestyRxBytes: obs.OpenrestyRxBytes,
|
||||
OpenrestyTxBytes: obs.OpenrestyTxBytes,
|
||||
OpenrestyConnections: obs.OpenrestyConnections,
|
||||
}
|
||||
return tx.Create(record).Error
|
||||
}
|
||||
|
||||
func persistNodeTrafficReport(tx *gorm.DB, nodeID string, report *NodeTrafficReport, reportedAt time.Time) error {
|
||||
if report == nil {
|
||||
return nil
|
||||
}
|
||||
if report.WindowEndedAtUnix > 0 && report.WindowStartedAtUnix > report.WindowEndedAtUnix {
|
||||
return errors.New("traffic report window_started_at_unix 不能大于 window_ended_at_unix")
|
||||
}
|
||||
record := &model.OpenFlareRequestReport{
|
||||
NodeID: nodeID,
|
||||
WindowStartedAt: timeFromUnix(report.WindowStartedAtUnix, reportedAt),
|
||||
WindowEndedAt: timeFromUnix(report.WindowEndedAtUnix, reportedAt),
|
||||
RequestCount: report.RequestCount,
|
||||
ErrorCount: report.ErrorCount,
|
||||
UniqueVisitorCount: report.UniqueVisitorCount,
|
||||
StatusCodesJSON: marshalJSON(report.StatusCodes),
|
||||
TopDomainsJSON: marshalJSON(report.TopDomains),
|
||||
SourceCountriesJSON: marshalJSON(report.SourceCountries),
|
||||
}
|
||||
exists, err := requestReportExists(tx, nodeID, record.WindowStartedAt, record.WindowEndedAt)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
return nil
|
||||
}
|
||||
return tx.Create(record).Error
|
||||
}
|
||||
|
||||
func persistNodeAccessLogs(tx *gorm.DB, nodeID string, logs []NodeAccessLog, reportedAt time.Time) error {
|
||||
if len(logs) == 0 {
|
||||
return nil
|
||||
}
|
||||
resolver, err := newAccessLogRegionResolver()
|
||||
if err != nil {
|
||||
slog.Warn("initialize access log geo resolver failed", "node_id", nodeID, "error", err)
|
||||
}
|
||||
if resolver != nil {
|
||||
defer resolver.Close()
|
||||
}
|
||||
for _, item := range logs {
|
||||
record := &model.OpenFlareAccessLog{
|
||||
NodeID: nodeID,
|
||||
LoggedAt: timeFromUnix(item.LoggedAtUnix, reportedAt),
|
||||
RemoteAddr: strings.TrimSpace(item.RemoteAddr),
|
||||
Region: "",
|
||||
Host: strings.TrimSpace(item.Host),
|
||||
Path: truncateForDatabase(strings.TrimSpace(item.Path), accessLogPathMaxLength),
|
||||
StatusCode: item.StatusCode,
|
||||
}
|
||||
if resolver != nil {
|
||||
record.Region = resolver.Resolve(record.RemoteAddr)
|
||||
}
|
||||
exists, err := accessLogExists(tx, record)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
continue
|
||||
}
|
||||
if err := tx.Create(record).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
_, err = deleteAccessLogsByNodeBefore(tx, nodeID, reportedAt.Add(-nodeAccessLogRetentionWindow))
|
||||
return err
|
||||
}
|
||||
|
||||
func reconcileNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEvent, reportedAt time.Time) error {
|
||||
return ReconcileScopedNodeHealthEvents(tx, nodeID, events, reportedAt, nil)
|
||||
}
|
||||
|
||||
// ReconcileScopedNodeHealthEvents reconciles health events, optionally scoped to managed event types.
|
||||
func ReconcileScopedNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEvent, reportedAt time.Time, managedEventTypes map[string]struct{}) error {
|
||||
activeTypes := make(map[string]NodeHealthEvent, len(events))
|
||||
for _, event := range events {
|
||||
eventType := normalizeHealthEventType(event.EventType)
|
||||
if eventType == "" {
|
||||
continue
|
||||
}
|
||||
if len(managedEventTypes) > 0 {
|
||||
if _, ok := managedEventTypes[eventType]; !ok {
|
||||
continue
|
||||
}
|
||||
}
|
||||
event.EventType = eventType
|
||||
event.Severity = normalizeHealthSeverity(event.Severity)
|
||||
if event.TriggeredAtUnix <= 0 {
|
||||
event.TriggeredAtUnix = reportedAt.Unix()
|
||||
}
|
||||
activeTypes[eventType] = event
|
||||
}
|
||||
|
||||
var activeEvents []*model.OpenFlareHealthEvent
|
||||
query := tx.Where("node_id = ? AND status = ?", nodeID, healthEventStatusActive)
|
||||
if len(managedEventTypes) > 0 {
|
||||
scopedTypes := make([]string, 0, len(managedEventTypes))
|
||||
for eventType := range managedEventTypes {
|
||||
eventType = normalizeHealthEventType(eventType)
|
||||
if eventType != "" {
|
||||
scopedTypes = append(scopedTypes, eventType)
|
||||
}
|
||||
}
|
||||
if len(scopedTypes) == 0 {
|
||||
return nil
|
||||
}
|
||||
query = query.Where("event_type IN ?", scopedTypes)
|
||||
}
|
||||
if err := query.Find(&activeEvents).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
activeByType := make(map[string]*model.OpenFlareHealthEvent, len(activeEvents))
|
||||
for _, event := range activeEvents {
|
||||
activeByType[event.EventType] = event
|
||||
}
|
||||
|
||||
for eventType, event := range activeTypes {
|
||||
triggeredAt := timeFromUnix(event.TriggeredAtUnix, reportedAt)
|
||||
if existing, ok := activeByType[eventType]; ok {
|
||||
existing.Severity = event.Severity
|
||||
existing.Message = normalizeHealthEventMessage(event.Message)
|
||||
existing.LastTriggeredAt = triggeredAt
|
||||
existing.ReportedAt = reportedAt
|
||||
existing.MetadataJSON = marshalJSON(event.Metadata)
|
||||
existing.ResolvedAt = nil
|
||||
if err := tx.Save(existing).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
continue
|
||||
}
|
||||
record := &model.OpenFlareHealthEvent{
|
||||
NodeID: nodeID,
|
||||
EventType: eventType,
|
||||
Severity: event.Severity,
|
||||
Status: healthEventStatusActive,
|
||||
Message: normalizeHealthEventMessage(event.Message),
|
||||
FirstTriggeredAt: triggeredAt,
|
||||
LastTriggeredAt: triggeredAt,
|
||||
ReportedAt: reportedAt,
|
||||
MetadataJSON: marshalJSON(event.Metadata),
|
||||
}
|
||||
if err := tx.Create(record).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
for _, existing := range activeEvents {
|
||||
if _, ok := activeTypes[existing.EventType]; ok {
|
||||
continue
|
||||
}
|
||||
resolvedAt := reportedAt
|
||||
existing.Status = healthEventStatusResolved
|
||||
existing.ReportedAt = reportedAt
|
||||
existing.ResolvedAt = &resolvedAt
|
||||
if err := tx.Save(existing).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func metricSnapshotExists(tx *gorm.DB, nodeID string, capturedAt time.Time) (bool, error) {
|
||||
var count int64
|
||||
if err := tx.Model(&model.OpenFlareMetricSnapshot{}).
|
||||
Where("node_id = ? AND captured_at = ?", nodeID, capturedAt).
|
||||
Limit(1).
|
||||
Count(&count).Error; err != nil {
|
||||
return false, err
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func requestReportExists(tx *gorm.DB, nodeID string, windowStartedAt, windowEndedAt time.Time) (bool, error) {
|
||||
var count int64
|
||||
if err := tx.Model(&model.OpenFlareRequestReport{}).
|
||||
Where("node_id = ? AND window_started_at = ? AND window_ended_at = ?", nodeID, windowStartedAt, windowEndedAt).
|
||||
Limit(1).
|
||||
Count(&count).Error; err != nil {
|
||||
return false, err
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func deleteAccessLogsByNodeBefore(tx *gorm.DB, nodeID string, before time.Time) (int64, error) {
|
||||
result := tx.Where("node_id = ? AND logged_at < ?", nodeID, before).Delete(&model.OpenFlareAccessLog{})
|
||||
return result.RowsAffected, result.Error
|
||||
}
|
||||
|
||||
func accessLogExists(tx *gorm.DB, record *model.OpenFlareAccessLog) (bool, error) {
|
||||
var count int64
|
||||
if err := tx.Model(&model.OpenFlareAccessLog{}).
|
||||
Where(
|
||||
"node_id = ? AND logged_at = ? AND remote_addr = ? AND host = ? AND path = ? AND status_code = ?",
|
||||
record.NodeID,
|
||||
record.LoggedAt,
|
||||
record.RemoteAddr,
|
||||
record.Host,
|
||||
record.Path,
|
||||
record.StatusCode,
|
||||
).
|
||||
Limit(1).
|
||||
Count(&count).Error; err != nil {
|
||||
return false, err
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func normalizeHealthEventType(eventType string) string {
|
||||
eventType = strings.TrimSpace(strings.ToLower(eventType))
|
||||
eventType = strings.ReplaceAll(eventType, " ", "_")
|
||||
return eventType
|
||||
}
|
||||
|
||||
func normalizeHealthSeverity(severity string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(severity)) {
|
||||
case healthSeverityCritical:
|
||||
return healthSeverityCritical
|
||||
case healthSeverityInfo:
|
||||
return healthSeverityInfo
|
||||
default:
|
||||
return healthSeverityWarning
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeHealthEventMessage(message string) string {
|
||||
return truncateForDatabase(message, 4096)
|
||||
}
|
||||
|
||||
func timeFromUnix(unixSeconds int64, fallback time.Time) time.Time {
|
||||
if unixSeconds <= 0 {
|
||||
return fallback
|
||||
}
|
||||
return time.Unix(unixSeconds, 0).UTC()
|
||||
}
|
||||
|
||||
// MarshalJSON serializes a value for database JSON columns.
|
||||
func MarshalJSON(value any) string {
|
||||
return marshalJSON(value)
|
||||
}
|
||||
|
||||
func marshalJSON(value any) string {
|
||||
if value == nil {
|
||||
return ""
|
||||
}
|
||||
raw, err := json.Marshal(value)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return string(raw)
|
||||
}
|
||||
@@ -0,0 +1,212 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/pages"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// RegisterHandler registers or discovers an agent node.
|
||||
// @Summary 注册或发现 Agent 节点
|
||||
// @Description 使用节点 access token 重新注册,或使用全局 discovery token 发现新节点;请求头需携带 X-Agent-Token
|
||||
// @Tags openflare-agent
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security AgentTokenAuth
|
||||
// @Param body body agent.NodePayload true "节点上报数据"
|
||||
// @Success 200 {object} response.Any{data=agent.RegistrationResponse} "注册成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "Token 无效"
|
||||
// @Router /api/v1/agent/nodes/register [post]
|
||||
func RegisterHandler(c *gin.Context) {
|
||||
var payload NodePayload
|
||||
if !apiutil.BindJSON(c, &payload) {
|
||||
return
|
||||
}
|
||||
payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
|
||||
|
||||
var (
|
||||
result *RegistrationResponse
|
||||
err error
|
||||
)
|
||||
if authNode, ok := AgentNodeFromContext(c); ok {
|
||||
result, err = RegisterWithAccessToken(c.Request.Context(), authNode, payload)
|
||||
} else {
|
||||
result, err = RegisterWithDiscovery(c.Request.Context(), payload)
|
||||
}
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
// HeartbeatHandler records agent heartbeat state.
|
||||
// @Summary Agent 心跳上报
|
||||
// @Description 上报节点状态、指标与健康事件,返回远程控制配置与活跃配置元信息
|
||||
// @Tags openflare-agent
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security AgentTokenAuth
|
||||
// @Param body body agent.NodePayload true "心跳数据"
|
||||
// @Success 200 {object} response.Any{data=agent.HeartbeatResponse} "心跳成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "Token 无效"
|
||||
// @Router /api/v1/agent/nodes/heartbeat [post]
|
||||
func HeartbeatHandler(c *gin.Context) {
|
||||
var payload NodePayload
|
||||
if !apiutil.BindJSON(c, &payload) {
|
||||
return
|
||||
}
|
||||
payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
|
||||
|
||||
authNode, ok := AgentNodeFromContext(c)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errInvalidAgentToken)
|
||||
return
|
||||
}
|
||||
|
||||
heartbeat, err := HeartbeatNode(c.Request.Context(), authNode, payload)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(heartbeat))
|
||||
}
|
||||
|
||||
// GetActiveConfigHandler returns the active configuration version.
|
||||
// @Summary 获取活跃配置版本
|
||||
// @Description 返回当前生效的完整配置包,供 Agent 拉取并应用
|
||||
// @Tags openflare-agent
|
||||
// @Produce json
|
||||
// @Security AgentTokenAuth
|
||||
// @Success 200 {object} response.Any{data=agent.ConfigResponse} "活跃配置"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "Token 无效"
|
||||
// @Router /api/v1/agent/config-versions/active [get]
|
||||
func GetActiveConfigHandler(c *gin.Context) {
|
||||
if _, ok := AgentNodeFromContext(c); !ok {
|
||||
response.AbortUnauthorized(c, errNodeMissingFromContext)
|
||||
return
|
||||
}
|
||||
config, err := GetActiveConfig(c.Request.Context())
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(config))
|
||||
}
|
||||
|
||||
// SyncWAFIPGroupsHandler syncs WAF IP groups for an agent.
|
||||
// @Summary 同步 WAF IP 组
|
||||
// @Description 按 ID 与校验和增量同步 WAF IP 组定义
|
||||
// @Tags openflare-agent
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security AgentTokenAuth
|
||||
// @Param body body agent.WAFIPGroupSyncInput true "同步请求"
|
||||
// @Success 200 {object} response.Any{data=agent.WAFIPGroupSyncResult} "同步结果"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "Token 无效"
|
||||
// @Router /api/v1/agent/waf/ip-groups/sync [post]
|
||||
func SyncWAFIPGroupsHandler(c *gin.Context) {
|
||||
var input WAFIPGroupSyncInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
result, err := SyncWAFIPGroups(c.Request.Context(), input)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
// ReportApplyLogHandler records an agent apply log entry.
|
||||
// @Summary 上报配置应用日志
|
||||
// @Description 记录 Agent 配置下发与应用结果
|
||||
// @Tags openflare-agent
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security AgentTokenAuth
|
||||
// @Param body body agent.ApplyLogPayload true "应用日志"
|
||||
// @Success 200 {object} response.Any{data=model.OpenFlareApplyLog} "日志记录"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "Token 无效"
|
||||
// @Router /api/v1/agent/apply-logs [post]
|
||||
func ReportApplyLogHandler(c *gin.Context) {
|
||||
var payload ApplyLogPayload
|
||||
if !apiutil.BindJSON(c, &payload) {
|
||||
return
|
||||
}
|
||||
if authNode, ok := AgentNodeFromContext(c); ok {
|
||||
payload.NodeID = authNode.NodeID
|
||||
}
|
||||
log, err := ReportApplyLog(c.Request.Context(), payload)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(log))
|
||||
}
|
||||
|
||||
// DownloadPagesPackageHandler streams the Pages deployment artifact to an authenticated agent.
|
||||
// @Summary 下载 Pages 部署包
|
||||
// @Description 流式下载指定部署的静态资源压缩包,供 Agent 边缘分发
|
||||
// @Tags openflare-agent
|
||||
// @Produce application/octet-stream
|
||||
// @Security AgentTokenAuth
|
||||
// @Param deployment_id path int true "部署 ID"
|
||||
// @Success 200 {file} binary "部署包文件"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "Token 无效"
|
||||
// @Router /api/v1/agent/pages/deployments/{deployment_id}/package [get]
|
||||
func DownloadPagesPackageHandler(c *gin.Context) {
|
||||
deploymentID, ok := pagesDeploymentIDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
packageObj, fileName, err := pages.OpenDeploymentPackage(c.Request.Context(), deploymentID)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
defer packageObj.Body.Close()
|
||||
c.Header("Content-Disposition", "attachment; filename="+fileName)
|
||||
if packageObj.ContentType != "" {
|
||||
c.Header("Content-Type", packageObj.ContentType)
|
||||
}
|
||||
c.DataFromReader(http.StatusOK, packageObj.ContentLength, packageObj.ContentType, packageObj.Body, nil)
|
||||
}
|
||||
|
||||
func pagesDeploymentIDParam(c *gin.Context) (uint, bool) {
|
||||
raw := c.Param("deployment_id")
|
||||
if raw == "" {
|
||||
response.AbortBadRequest(c, "无效的 ID")
|
||||
return 0, false
|
||||
}
|
||||
id64, err := strconv.ParseUint(raw, 10, 64)
|
||||
if err != nil || id64 == 0 {
|
||||
response.AbortBadRequest(c, "无效的 ID")
|
||||
return 0, false
|
||||
}
|
||||
return uint(id64), true
|
||||
}
|
||||
|
||||
// AgentWebSocketHandler upgrades an authenticated agent websocket connection.
|
||||
// @Summary Agent WebSocket 连接
|
||||
// @Description 升级为 WebSocket 长连接,用于实时推送配置同步、WAF IP 组等指令;需携带 X-Agent-Token
|
||||
// @Tags openflare-agent
|
||||
// @Security AgentTokenAuth
|
||||
// @Failure 401 {object} response.Any "Token 无效"
|
||||
// @Router /api/v1/agent/ws [get]
|
||||
func AgentWebSocketHandler(c *gin.Context) {
|
||||
authNode, ok := AgentNodeFromContext(c)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errInvalidAgentToken)
|
||||
return
|
||||
}
|
||||
websocket.ServeAgent(c, authNode.NodeID, HandleWSStatus)
|
||||
}
|
||||
@@ -0,0 +1,119 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const (
|
||||
nodeStatusOnline = "online"
|
||||
applyResultOK = "success"
|
||||
applyResultWarn = "warning"
|
||||
applyResultFailed = "failed"
|
||||
)
|
||||
|
||||
// NodePayload is the agent register/heartbeat payload.
|
||||
type NodePayload struct {
|
||||
NodeID string `json:"node_id"`
|
||||
Name string `json:"name"`
|
||||
IP string `json:"ip"`
|
||||
Version string `json:"version"`
|
||||
ExtVersion string `json:"ext_version"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
LastError string `json:"last_error"`
|
||||
OpenrestyStatus string `json:"openresty_status"`
|
||||
OpenrestyMessage string `json:"openresty_message"`
|
||||
Profile *NodeSystemProfile `json:"profile,omitempty"`
|
||||
Snapshot *NodeMetricSnapshot `json:"snapshot,omitempty"`
|
||||
OpenrestyObservation *NodeOpenrestyObservation `json:"openresty_observation,omitempty"`
|
||||
TrafficReport *NodeTrafficReport `json:"traffic_report,omitempty"`
|
||||
AccessLogs []NodeAccessLog `json:"access_logs,omitempty"`
|
||||
BufferedObservability []BufferedObservabilityRecord `json:"buffered_observability,omitempty"`
|
||||
HealthEvents []NodeHealthEvent `json:"health_events"`
|
||||
WAFIPGroupChecksums map[string]string `json:"waf_ip_group_checksums,omitempty"`
|
||||
}
|
||||
|
||||
// ApplyLogPayload is the agent apply log report payload.
|
||||
type ApplyLogPayload struct {
|
||||
NodeID string `json:"node_id"`
|
||||
Version string `json:"version"`
|
||||
Result string `json:"result"`
|
||||
Message string `json:"message"`
|
||||
Checksum string `json:"checksum"`
|
||||
MainConfigChecksum string `json:"main_config_checksum"`
|
||||
RouteConfigChecksum string `json:"route_config_checksum"`
|
||||
SupportFileCount int `json:"support_file_count"`
|
||||
}
|
||||
|
||||
// RegistrationResponse is returned after agent registration.
|
||||
type RegistrationResponse struct {
|
||||
NodeID string `json:"node_id"`
|
||||
AccessToken string `json:"access_token"`
|
||||
Name string `json:"name"`
|
||||
}
|
||||
|
||||
// Settings carries remote agent control flags.
|
||||
type Settings struct {
|
||||
HeartbeatInterval int `json:"heartbeat_interval"`
|
||||
WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"`
|
||||
AutoUpdate bool `json:"auto_update"`
|
||||
UpdateRepo string `json:"update_repo"`
|
||||
UpdateNow bool `json:"update_now"`
|
||||
UpdateChannel string `json:"update_channel"`
|
||||
UpdateTag string `json:"update_tag"`
|
||||
RestartOpenrestyNow bool `json:"restart_openresty_now"`
|
||||
}
|
||||
|
||||
// ActiveConfigMeta summarizes the active configuration version.
|
||||
type ActiveConfigMeta struct {
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
}
|
||||
|
||||
// SupportFile is a configuration support artifact shipped to agents.
|
||||
type SupportFile struct {
|
||||
Path string `json:"path"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
// ConfigResponse is the full active config payload for agents.
|
||||
type ConfigResponse struct {
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
SourceConfigJSON string `json:"source_config_json"`
|
||||
SupportFiles []SupportFile `json:"support_files"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
// WAFIPGroup is a WAF IP group snapshot for agents.
|
||||
type WAFIPGroup struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Enabled bool `json:"enabled"`
|
||||
IPList []string `json:"ip_list"`
|
||||
Checksum string `json:"checksum"`
|
||||
}
|
||||
|
||||
// WAFIPGroupSyncInput requests changed WAF IP groups.
|
||||
type WAFIPGroupSyncInput struct {
|
||||
IDs []uint `json:"ids"`
|
||||
Checksums map[string]string `json:"checksums"`
|
||||
}
|
||||
|
||||
// WAFIPGroupSyncResult returns synced WAF IP groups.
|
||||
type WAFIPGroupSyncResult struct {
|
||||
Groups []WAFIPGroup `json:"groups"`
|
||||
}
|
||||
|
||||
// HeartbeatResponse is the heartbeat handler result.
|
||||
type HeartbeatResponse struct {
|
||||
Node *model.OpenFlareNode `json:"node"`
|
||||
AgentSettings *Settings `json:"agent_settings"`
|
||||
ActiveConfig *ActiveConfigMeta `json:"active_config"`
|
||||
WAFIPGroups []WAFIPGroup `json:"waf_ip_groups,omitempty"`
|
||||
}
|
||||
@@ -0,0 +1,205 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
type snapshotWAFRuleGroupRef struct {
|
||||
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
|
||||
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
|
||||
}
|
||||
|
||||
type snapshotWAFSection struct {
|
||||
RuleGroups []snapshotWAFRuleGroupRef `json:"rule_groups"`
|
||||
}
|
||||
|
||||
type activeConfigSnapshot struct {
|
||||
WAF snapshotWAFSection `json:"waf"`
|
||||
}
|
||||
|
||||
// WAFIPGroupsForAgent builds agent-facing WAF IP group payloads for the given ids.
|
||||
func WAFIPGroupsForAgent(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
|
||||
return buildAgentWAFIPGroups(ctx, ids)
|
||||
}
|
||||
|
||||
// ChangedWAFIPGroupsForAgent returns WAF IP groups whose checksums differ from the agent state.
|
||||
func ChangedWAFIPGroupsForAgent(ctx context.Context, ids []uint, checksums map[string]string) ([]WAFIPGroup, error) {
|
||||
targetIDs := uniqueUintIDs(ids)
|
||||
if len(targetIDs) == 0 {
|
||||
activeIDs, err := activeConfigWAFIPGroupIDs(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targetIDs = activeIDs
|
||||
}
|
||||
if len(targetIDs) == 0 {
|
||||
return []WAFIPGroup{}, nil
|
||||
}
|
||||
groups, err := buildAgentWAFIPGroups(ctx, targetIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
changed := make([]WAFIPGroup, 0, len(groups))
|
||||
for _, group := range groups {
|
||||
if strings.TrimSpace(checksums[fmt.Sprintf("%d", group.ID)]) == group.Checksum {
|
||||
continue
|
||||
}
|
||||
changed = append(changed, group)
|
||||
}
|
||||
return changed, nil
|
||||
}
|
||||
|
||||
func buildAgentWAFIPGroups(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
|
||||
ids = uniqueUintIDs(ids)
|
||||
if len(ids) == 0 {
|
||||
return []WAFIPGroup{}, nil
|
||||
}
|
||||
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
|
||||
groups, err := model.ListOpenFlareWAFIPGroupsByIDs(ctx, ids)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
groupByID := make(map[uint]*model.OpenFlareWAFIPGroup, len(groups))
|
||||
for _, group := range groups {
|
||||
groupByID[group.ID] = group
|
||||
}
|
||||
result := make([]WAFIPGroup, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
group := groupByID[id]
|
||||
if group == nil {
|
||||
continue
|
||||
}
|
||||
agentGroup, err := buildAgentWAFIPGroup(group)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result = append(result, agentGroup)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func buildAgentWAFIPGroup(group *model.OpenFlareWAFIPGroup) (WAFIPGroup, error) {
|
||||
if group == nil {
|
||||
return WAFIPGroup{}, errors.New("IP 组不存在")
|
||||
}
|
||||
ips, err := decodeWAFIPGroupStringList(group.IPList)
|
||||
if err != nil {
|
||||
return WAFIPGroup{}, err
|
||||
}
|
||||
if !group.Enabled {
|
||||
ips = []string{}
|
||||
}
|
||||
agentGroup := WAFIPGroup{
|
||||
ID: group.ID,
|
||||
Name: group.Name,
|
||||
Type: group.Type,
|
||||
Enabled: group.Enabled,
|
||||
IPList: ips,
|
||||
}
|
||||
agentGroup.Checksum = checksumAgentWAFIPGroup(agentGroup)
|
||||
return agentGroup, nil
|
||||
}
|
||||
|
||||
func checksumAgentWAFIPGroup(group WAFIPGroup) string {
|
||||
payload := struct {
|
||||
ID uint `json:"id"`
|
||||
Enabled bool `json:"enabled"`
|
||||
IPList []string `json:"ip_list"`
|
||||
}{
|
||||
ID: group.ID,
|
||||
Enabled: group.Enabled,
|
||||
IPList: append([]string{}, group.IPList...),
|
||||
}
|
||||
sort.Strings(payload.IPList)
|
||||
data, _ := json.Marshal(payload)
|
||||
sum := sha256.Sum256(data)
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
|
||||
version, err := loadActiveConfigVersion(ctx)
|
||||
if err != nil {
|
||||
if isActiveConfigNotFound(err) {
|
||||
return []uint{}, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
snapshot, err := parseActiveConfigSnapshot(version.SnapshotJSON)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
idSet := make(map[uint]struct{})
|
||||
for _, group := range snapshot.WAF.RuleGroups {
|
||||
for _, id := range group.IPWhitelistGroups {
|
||||
if id > 0 {
|
||||
idSet[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
for _, id := range group.IPBlacklistGroups {
|
||||
if id > 0 {
|
||||
idSet[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
}
|
||||
ids := make([]uint, 0, len(idSet))
|
||||
for id := range idSet {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
func parseActiveConfigSnapshot(snapshotJSON string) (*activeConfigSnapshot, error) {
|
||||
text := strings.TrimSpace(snapshotJSON)
|
||||
if text == "" {
|
||||
return &activeConfigSnapshot{}, nil
|
||||
}
|
||||
var snapshot activeConfigSnapshot
|
||||
if err := json.Unmarshal([]byte(text), &snapshot); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if snapshot.WAF.RuleGroups == nil {
|
||||
snapshot.WAF.RuleGroups = []snapshotWAFRuleGroupRef{}
|
||||
}
|
||||
return &snapshot, nil
|
||||
}
|
||||
|
||||
func decodeWAFIPGroupStringList(raw string) ([]string, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
return []string{}, nil
|
||||
}
|
||||
var items []string
|
||||
if err := json.Unmarshal([]byte(text), &items); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func uniqueUintIDs(ids []uint) []uint {
|
||||
normalized := make([]uint, 0, len(ids))
|
||||
seen := make(map[uint]struct{}, len(ids))
|
||||
for _, id := range ids {
|
||||
if id == 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
normalized = append(normalized, id)
|
||||
}
|
||||
return normalized
|
||||
}
|
||||
@@ -0,0 +1,158 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"strconv"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupWAFIPGroupTestDB(t *testing.T) func() {
|
||||
t.Helper()
|
||||
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, sqliteDB.AutoMigrate(
|
||||
&model.OpenFlareWAFIPGroup{},
|
||||
&configVersionRecord{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
}
|
||||
}
|
||||
|
||||
func seedActiveConfigWithWAFIPGroup(t *testing.T, ctx context.Context, ipGroupID uint) {
|
||||
t.Helper()
|
||||
|
||||
snapshot := map[string]any{
|
||||
"routes": []any{},
|
||||
"waf": map[string]any{
|
||||
"rule_groups": []map[string]any{
|
||||
{
|
||||
"id": 1,
|
||||
"name": "agent refs",
|
||||
"enabled": true,
|
||||
"ip_blacklist_group_ids": []uint{ipGroupID},
|
||||
},
|
||||
},
|
||||
"bindings": []any{},
|
||||
},
|
||||
}
|
||||
snapshotJSON, err := json.Marshal(snapshot)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, db.DB(ctx).Create(&configVersionRecord{
|
||||
Version: "20260618-001",
|
||||
SnapshotJSON: string(snapshotJSON),
|
||||
Checksum: "test-checksum",
|
||||
IsActive: true,
|
||||
}).Error)
|
||||
}
|
||||
|
||||
func TestChangedWAFIPGroupsForAgentReturnsChecksumDelta(t *testing.T) {
|
||||
cleanup := setupWAFIPGroupTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
ipGroup := &model.OpenFlareWAFIPGroup{
|
||||
Name: "agent runtime group",
|
||||
Type: "manual",
|
||||
Enabled: true,
|
||||
IPList: `["203.0.113.44"]`,
|
||||
}
|
||||
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
|
||||
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
|
||||
|
||||
groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, groups, 1)
|
||||
assert.Equal(t, ipGroup.ID, groups[0].ID)
|
||||
assert.Equal(t, "203.0.113.44", groups[0].IPList[0])
|
||||
assert.NotEmpty(t, groups[0].Checksum)
|
||||
|
||||
groupKey := strconv.FormatUint(uint64(ipGroup.ID), 10)
|
||||
same, err := ChangedWAFIPGroupsForAgent(ctx, nil, map[string]string{groupKey: groups[0].Checksum})
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, same)
|
||||
|
||||
ipGroup.IPList = `["203.0.113.45"]`
|
||||
require.NoError(t, model.UpdateOpenFlareWAFIPGroup(ctx, ipGroup))
|
||||
|
||||
delta, err := ChangedWAFIPGroupsForAgent(ctx, nil, map[string]string{groupKey: groups[0].Checksum})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, delta, 1)
|
||||
assert.Equal(t, ipGroup.ID, delta[0].ID)
|
||||
assert.Equal(t, "203.0.113.45", delta[0].IPList[0])
|
||||
assert.NotEqual(t, groups[0].Checksum, delta[0].Checksum)
|
||||
}
|
||||
|
||||
func TestSyncWAFIPGroupsReturnsChangedGroups(t *testing.T) {
|
||||
cleanup := setupWAFIPGroupTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
ipGroup := &model.OpenFlareWAFIPGroup{
|
||||
Name: "sync group",
|
||||
Type: "manual",
|
||||
Enabled: true,
|
||||
IPList: `["198.51.100.10"]`,
|
||||
}
|
||||
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
|
||||
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
|
||||
|
||||
result, err := SyncWAFIPGroups(ctx, WAFIPGroupSyncInput{
|
||||
IDs: []uint{ipGroup.ID},
|
||||
Checksums: map[string]string{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, result.Groups, 1)
|
||||
assert.Equal(t, ipGroup.ID, result.Groups[0].ID)
|
||||
assert.Equal(t, "198.51.100.10", result.Groups[0].IPList[0])
|
||||
|
||||
result, err = SyncWAFIPGroups(ctx, WAFIPGroupSyncInput{
|
||||
IDs: []uint{ipGroup.ID},
|
||||
Checksums: map[string]string{
|
||||
strconv.FormatUint(uint64(ipGroup.ID), 10): result.Groups[0].Checksum,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, result.Groups)
|
||||
}
|
||||
|
||||
func TestChangedWAFIPGroupsForAgentDisabledGroupClearsIPList(t *testing.T) {
|
||||
cleanup := setupWAFIPGroupTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
ipGroup := &model.OpenFlareWAFIPGroup{
|
||||
Name: "disabled group",
|
||||
Type: "manual",
|
||||
Enabled: true,
|
||||
IPList: `["203.0.113.10"]`,
|
||||
}
|
||||
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
|
||||
ipGroup.Enabled = false
|
||||
require.NoError(t, model.UpdateOpenFlareWAFIPGroup(ctx, ipGroup))
|
||||
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
|
||||
|
||||
groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, groups, 1)
|
||||
assert.False(t, groups[0].Enabled)
|
||||
assert.Empty(t, groups[0].IPList)
|
||||
assert.NotEmpty(t, groups[0].Checksum)
|
||||
}
|
||||
@@ -0,0 +1,57 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
|
||||
ofws "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
// HandleWSStatus processes an agent websocket status payload (replaces HTTP heartbeat in WS mode).
|
||||
func HandleWSStatus(ctx context.Context, nodeID, remoteAddr string, rawPayload json.RawMessage) {
|
||||
var payload NodePayload
|
||||
if err := json.Unmarshal(rawPayload, &payload); err != nil {
|
||||
slog.Debug("agent ws status payload decode failed", "node_id", nodeID, "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
authNode, err := model.GetOpenFlareNodeByNodeID(ctx, nodeID)
|
||||
if err != nil {
|
||||
slog.Debug("agent ws status reload node failed", "node_id", nodeID, "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
payload.IP = resolveReportedNodeIP(payload.IP, remoteAddr)
|
||||
response, err := HeartbeatNode(ctx, authNode, payload)
|
||||
if err != nil {
|
||||
slog.Debug("agent ws status handling failed", "node_id", nodeID, "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
settingsSent := false
|
||||
if response.AgentSettings != nil {
|
||||
settingsSent = ofws.SendAgentSettings(nodeID, response.AgentSettings)
|
||||
}
|
||||
activeConfigSent := false
|
||||
if response.ActiveConfig != nil {
|
||||
activeConfigSent = ofws.SendAgentActiveConfig(nodeID, response.ActiveConfig)
|
||||
}
|
||||
wafIPGroupsSent := false
|
||||
if len(response.WAFIPGroups) > 0 {
|
||||
wafIPGroupsSent = ofws.SendAgentWAFIPGroups(nodeID, response.WAFIPGroups)
|
||||
}
|
||||
|
||||
slog.Debug("agent ws status processed",
|
||||
"node_id", nodeID,
|
||||
"current_version", payload.CurrentVersion,
|
||||
"openresty_status", payload.OpenrestyStatus,
|
||||
"settings_sent", settingsSent,
|
||||
"active_config_sent", activeConfigSent,
|
||||
"waf_ip_groups_sent", wafIPGroupsSent,
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package apiutil provides HTTP helpers for OpenFlare v1 custom API handlers.
|
||||
package apiutil
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
const errInvalidParams = "参数错误"
|
||||
const errInvalidID = "无效的 ID"
|
||||
|
||||
// BindJSON binds JSON body; returns false after aborting with 400.
|
||||
func BindJSON(c *gin.Context, dst any) bool {
|
||||
if err := c.ShouldBindJSON(dst); err != nil {
|
||||
response.AbortBadRequest(c, errInvalidParams)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// IDParam parses :id from the URL path.
|
||||
func IDParam(c *gin.Context) (uint, bool) {
|
||||
raw := c.Param("id")
|
||||
if raw == "" {
|
||||
response.AbortBadRequest(c, errInvalidID)
|
||||
return 0, false
|
||||
}
|
||||
id64, err := strconv.ParseUint(raw, 10, 64)
|
||||
if err != nil || id64 == 0 {
|
||||
response.AbortBadRequest(c, errInvalidID)
|
||||
return 0, false
|
||||
}
|
||||
return uint(id64), true
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package apiutil
|
||||
|
||||
import (
|
||||
"errors"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// AbortNotFoundIfMissing maps gorm.ErrRecordNotFound to 404; other errors to 400.
|
||||
func AbortNotFoundIfMissing(c *gin.Context, err error, notFoundMsg string) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, notFoundMsg)
|
||||
return true
|
||||
}
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return true
|
||||
}
|
||||
|
||||
// AbortBadRequestOnError writes a 400 for any non-nil error.
|
||||
func AbortBadRequestOnError(c *gin.Context, err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package apiutil
|
||||
|
||||
import (
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/admin"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// AdminMiddlewares returns Wavelet-standard middlewares for OpenFlare console routes.
|
||||
// OpenFlare no longer distinguishes Admin vs Root tiers; all management endpoints share
|
||||
// the same gate: user.IsAdmin for session users, token_admin for Access Token callers.
|
||||
func AdminMiddlewares() []gin.HandlerFunc {
|
||||
return []gin.HandlerFunc{oauth.LoginRequired(), admin.LoginAdminRequired()}
|
||||
}
|
||||
@@ -0,0 +1,158 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package apiutil
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/admin"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-contrib/sessions/cookie"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupAdminMiddlewareTest(t *testing.T) (*gin.Engine, *gorm.DB, func()) {
|
||||
t.Helper()
|
||||
|
||||
dbConn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, dbConn.AutoMigrate(&model.User{}, &model.AccessToken{}))
|
||||
db.SetDB(dbConn)
|
||||
|
||||
sessionCookieName := "test_admin_middleware_session"
|
||||
if config.Config.App.SessionCookieName != "" {
|
||||
sessionCookieName = config.Config.App.SessionCookieName
|
||||
}
|
||||
store := cookie.NewStore([]byte("test_admin_middleware_session_secret"))
|
||||
store.Options(oauth.GetSessionOptions(3600))
|
||||
engine := testhelper.NewTestGinEngine(sessions.Sessions(sessionCookieName, store))
|
||||
protected := engine.Group("/protected", AdminMiddlewares()...)
|
||||
protected.GET("", func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(gin.H{"ok": true}))
|
||||
})
|
||||
|
||||
cleanup := func() {
|
||||
db.SetDB(nil)
|
||||
}
|
||||
|
||||
return engine, dbConn, cleanup
|
||||
}
|
||||
|
||||
func seedUser(t *testing.T, dbConn *gorm.DB, username string, isAdmin bool) *model.User {
|
||||
t.Helper()
|
||||
|
||||
user := &model.User{
|
||||
ID: idgen.NextUint64ID(),
|
||||
Username: username,
|
||||
Nickname: username,
|
||||
Email: username + "@openflare.test",
|
||||
IsActive: true,
|
||||
IsAdmin: isAdmin,
|
||||
}
|
||||
require.NoError(t, dbConn.Create(user).Error)
|
||||
return user
|
||||
}
|
||||
|
||||
func seedAccessToken(t *testing.T, dbConn *gorm.DB, user *model.User, isAdmin bool) string {
|
||||
t.Helper()
|
||||
|
||||
token, err := model.GenerateTokenString()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, dbConn.Create(&model.AccessToken{
|
||||
UserID: user.ID,
|
||||
Name: user.Username + "-token",
|
||||
TokenHash: model.HashToken(token),
|
||||
MaskedToken: model.MaskTokenString(token),
|
||||
IsAdmin: isAdmin,
|
||||
}).Error)
|
||||
return token
|
||||
}
|
||||
|
||||
func decodeResponse(t *testing.T, rec *httptest.ResponseRecorder) response.Any {
|
||||
t.Helper()
|
||||
|
||||
var resp response.Any
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
|
||||
return resp
|
||||
}
|
||||
|
||||
func TestAdminRequiredUnauthenticated(t *testing.T) {
|
||||
engine, _, cleanup := setupAdminMiddlewareTest(t)
|
||||
defer cleanup()
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
|
||||
engine.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusUnauthorized, rec.Code)
|
||||
resp := decodeResponse(t, rec)
|
||||
assert.NotEmpty(t, resp.ErrorMsg)
|
||||
}
|
||||
|
||||
func TestAdminRequiredNonAdminToken(t *testing.T) {
|
||||
engine, dbConn, cleanup := setupAdminMiddlewareTest(t)
|
||||
defer cleanup()
|
||||
|
||||
user := seedUser(t, dbConn, "regular", false)
|
||||
token := seedAccessToken(t, dbConn, user, false)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
|
||||
req.Header.Set("X-Access-Token", token)
|
||||
engine.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusNotFound, rec.Code)
|
||||
resp := decodeResponse(t, rec)
|
||||
assert.Equal(t, admin.TokenAdminRequired, resp.ErrorMsg)
|
||||
}
|
||||
|
||||
func TestAdminRequiredAdminWithoutTokenAdmin(t *testing.T) {
|
||||
engine, dbConn, cleanup := setupAdminMiddlewareTest(t)
|
||||
defer cleanup()
|
||||
|
||||
user := seedUser(t, dbConn, "admin-no-token-admin", true)
|
||||
token := seedAccessToken(t, dbConn, user, false)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
|
||||
req.Header.Set("X-Access-Token", token)
|
||||
engine.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusNotFound, rec.Code)
|
||||
resp := decodeResponse(t, rec)
|
||||
assert.Equal(t, admin.TokenAdminRequired, resp.ErrorMsg)
|
||||
}
|
||||
|
||||
func TestAdminRequiredAdminWithTokenAdmin(t *testing.T) {
|
||||
engine, dbConn, cleanup := setupAdminMiddlewareTest(t)
|
||||
defer cleanup()
|
||||
|
||||
user := seedUser(t, dbConn, "admin", true)
|
||||
token := seedAccessToken(t, dbConn, user, true)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
|
||||
req.Header.Set("X-Access-Token", token)
|
||||
engine.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, rec.Code)
|
||||
resp := decodeResponse(t, rec)
|
||||
assert.Empty(t, resp.ErrorMsg)
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package apiutil
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// RegisterCollection registers a collection endpoint on both "" and "/" so requests
|
||||
// work with or without a trailing slash.
|
||||
func RegisterCollection(route *gin.RouterGroup, method string, handlers ...gin.HandlerFunc) {
|
||||
route.Handle(method, "/", handlers...)
|
||||
if !strings.HasSuffix(route.BasePath(), "/") {
|
||||
route.Handle(method, "", handlers...)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package apply_log
|
||||
|
||||
const (
|
||||
errRetentionDaysOutOfRange = "retention_days 必须在 1 到 3650 之间"
|
||||
)
|
||||
@@ -0,0 +1,128 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package apply_log
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultApplyLogPageSize = 20
|
||||
maxApplyLogPageSize = 200
|
||||
maxApplyLogRetentionDays = 3650
|
||||
)
|
||||
|
||||
// ListQuery filters apply logs for paginated listing.
|
||||
type ListQuery struct {
|
||||
NodeID string `json:"node_id"`
|
||||
PageNo int `json:"pageNo"`
|
||||
PageSize int `json:"pageSize"`
|
||||
}
|
||||
|
||||
// ListResult is the paginated apply log list response.
|
||||
type ListResult struct {
|
||||
Rows []*model.OpenFlareApplyLog `json:"rows"`
|
||||
Current int `json:"current"`
|
||||
Total int `json:"total"`
|
||||
TotalPage int `json:"totalPage"`
|
||||
}
|
||||
|
||||
// CleanupInput controls apply log cleanup behavior.
|
||||
type CleanupInput struct {
|
||||
DeleteAll bool `json:"delete_all"`
|
||||
RetentionDays int `json:"retention_days"`
|
||||
}
|
||||
|
||||
// CleanupResult reports apply log cleanup outcome.
|
||||
type CleanupResult struct {
|
||||
DeleteAll bool `json:"delete_all"`
|
||||
RetentionDays int `json:"retention_days"`
|
||||
DeletedCount int64 `json:"deleted_count"`
|
||||
Cutoff *time.Time `json:"cutoff,omitempty"`
|
||||
}
|
||||
|
||||
// ListPage returns paginated apply logs with optional node_id filter.
|
||||
func ListPage(ctx context.Context, input ListQuery) (*ListResult, error) {
|
||||
pageNo := normalizePageNo(input.PageNo)
|
||||
pageSize := normalizePageSize(input.PageSize)
|
||||
nodeID := strings.TrimSpace(input.NodeID)
|
||||
|
||||
rows, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{
|
||||
NodeID: nodeID,
|
||||
PageNo: pageNo,
|
||||
PageSize: pageSize,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
total, err := model.CountOpenFlareApplyLogs(ctx, nodeID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
totalPage := 0
|
||||
if total > 0 {
|
||||
totalPage = int((total + int64(pageSize) - 1) / int64(pageSize))
|
||||
}
|
||||
|
||||
return &ListResult{
|
||||
Rows: rows,
|
||||
Current: pageNo,
|
||||
Total: int(total),
|
||||
TotalPage: totalPage,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Cleanup removes old apply logs or deletes all records.
|
||||
func Cleanup(ctx context.Context, input CleanupInput) (*CleanupResult, error) {
|
||||
if input.DeleteAll {
|
||||
deleted, err := model.DeleteAllOpenFlareApplyLogs(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &CleanupResult{
|
||||
DeleteAll: true,
|
||||
DeletedCount: deleted,
|
||||
}, nil
|
||||
}
|
||||
|
||||
if input.RetentionDays <= 0 || input.RetentionDays > maxApplyLogRetentionDays {
|
||||
return nil, errors.New(errRetentionDaysOutOfRange)
|
||||
}
|
||||
|
||||
cutoff := time.Now().UTC().Add(-time.Duration(input.RetentionDays) * 24 * time.Hour)
|
||||
deleted, err := model.DeleteOpenFlareApplyLogsBefore(ctx, cutoff)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &CleanupResult{
|
||||
RetentionDays: input.RetentionDays,
|
||||
DeletedCount: deleted,
|
||||
Cutoff: &cutoff,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizePageNo(pageNo int) int {
|
||||
if pageNo <= 0 {
|
||||
return 1
|
||||
}
|
||||
return pageNo
|
||||
}
|
||||
|
||||
func normalizePageSize(pageSize int) int {
|
||||
if pageSize <= 0 {
|
||||
return defaultApplyLogPageSize
|
||||
}
|
||||
if pageSize > maxApplyLogPageSize {
|
||||
return maxApplyLogPageSize
|
||||
}
|
||||
return pageSize
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package apply_log
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupApplyLogTestDB(t *testing.T) func() {
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
err = sqliteDB.AutoMigrate(&model.OpenFlareApplyLog{})
|
||||
require.NoError(t, err)
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListPageAndCleanup(t *testing.T) {
|
||||
cleanup := setupApplyLogTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
now := time.Now().UTC()
|
||||
|
||||
logs := []model.OpenFlareApplyLog{
|
||||
{NodeID: "node-logs", Version: "v1", Result: "success", Message: "1", CreatedAt: now.Add(-10 * 24 * time.Hour)},
|
||||
{NodeID: "node-logs", Version: "v2", Result: "success", Message: "2", CreatedAt: now.Add(-5 * 24 * time.Hour)},
|
||||
{NodeID: "node-logs", Version: "v3", Result: "success", Message: "3", CreatedAt: now},
|
||||
}
|
||||
for i := range logs {
|
||||
require.NoError(t, db.DB(ctx).Create(&logs[i]).Error)
|
||||
}
|
||||
|
||||
pageResult, err := ListPage(ctx, ListQuery{
|
||||
NodeID: "node-logs",
|
||||
PageNo: 1,
|
||||
PageSize: 2,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 3, pageResult.Total)
|
||||
assert.Len(t, pageResult.Rows, 2)
|
||||
assert.Equal(t, 2, pageResult.TotalPage)
|
||||
assert.Equal(t, 1, pageResult.Current)
|
||||
|
||||
cleanupResult, err := Cleanup(ctx, CleanupInput{
|
||||
DeleteAll: false,
|
||||
RetentionDays: 7,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(1), cleanupResult.DeletedCount)
|
||||
assert.NotNil(t, cleanupResult.Cutoff)
|
||||
|
||||
remaining, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{
|
||||
NodeID: "node-logs",
|
||||
PageNo: 1,
|
||||
PageSize: 10,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, remaining, 2)
|
||||
|
||||
cleanupAll, err := Cleanup(ctx, CleanupInput{DeleteAll: true})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(2), cleanupAll.DeletedCount)
|
||||
assert.True(t, cleanupAll.DeleteAll)
|
||||
|
||||
finalLogs, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{
|
||||
NodeID: "node-logs",
|
||||
PageNo: 1,
|
||||
PageSize: 10,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, finalLogs)
|
||||
}
|
||||
|
||||
func TestCleanupInvalidRetentionDays(t *testing.T) {
|
||||
cleanup := setupApplyLogTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
_, err := Cleanup(ctx, CleanupInput{RetentionDays: 0})
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, errRetentionDaysOutOfRange, err.Error())
|
||||
|
||||
_, err = Cleanup(ctx, CleanupInput{RetentionDays: 4000})
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, errRetentionDaysOutOfRange, err.Error())
|
||||
}
|
||||
@@ -0,0 +1,77 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package apply_log
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
|
||||
// GetApplyLogs lists apply logs with pagination and optional node_id filter.
|
||||
// @Summary 获取配置下发日志
|
||||
// @Description 分页返回节点配置下发记录,支持按节点 ID 筛选,需要管理员权限
|
||||
// @Tags openflare-apply-log
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param node_id query string false "节点 ID 筛选"
|
||||
// @Param pageNo query int false "页码"
|
||||
// @Param page_no query int false "页码(别名)"
|
||||
// @Param pageSize query int false "每页数量"
|
||||
// @Param page_size query int false "每页数量(别名)"
|
||||
// @Success 200 {object} response.Any{data=apply_log.ListResult} "下发日志列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/apply-logs [get]
|
||||
func GetApplyLogs(c *gin.Context) {
|
||||
result, err := ListPage(c.Request.Context(), ListQuery{
|
||||
NodeID: c.Query("node_id"),
|
||||
PageNo: readIntQuery(c, "pageNo", "page_no"),
|
||||
PageSize: readIntQuery(c, "pageSize", "page_size"),
|
||||
})
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
// CleanupApplyLogs removes old apply logs or deletes all records.
|
||||
// @Summary 清理配置下发日志
|
||||
// @Description 按保留天数清理历史下发记录,或删除全部记录,需要管理员权限
|
||||
// @Tags openflare-apply-log
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param body body apply_log.CleanupInput true "清理参数"
|
||||
// @Success 200 {object} response.Any{data=apply_log.CleanupResult} "清理结果"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/apply-logs/cleanup [post]
|
||||
func CleanupApplyLogs(c *gin.Context) {
|
||||
var input CleanupInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
|
||||
result, err := Cleanup(c.Request.Context(), input)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
func readIntQuery(c *gin.Context, primary, secondary string) int {
|
||||
value := c.Query(primary)
|
||||
if value == "" {
|
||||
value = c.Query(secondary)
|
||||
}
|
||||
parsed, _ := strconv.Atoi(value)
|
||||
return parsed
|
||||
}
|
||||
@@ -0,0 +1,200 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package openflare
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/tasks"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/uptimekuma"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
)
|
||||
|
||||
const (
|
||||
// SSLRenewTask renews due ACME TLS certificates.
|
||||
SSLRenewTask = "openflare:ssl_renew"
|
||||
// TaskTypeSSLRenew is the admin task type for SSL renewal.
|
||||
TaskTypeSSLRenew = "of_ssl_renew"
|
||||
|
||||
// DatabaseAutoCleanupTask prunes observability tables by retention policy.
|
||||
DatabaseAutoCleanupTask = "openflare:database_auto_cleanup"
|
||||
// TaskTypeDatabaseAutoCleanup is the admin task type for observability cleanup.
|
||||
TaskTypeDatabaseAutoCleanup = "of_database_auto_cleanup"
|
||||
|
||||
// WAFIPGroupSyncTask syncs due automatic/subscription WAF IP groups.
|
||||
WAFIPGroupSyncTask = "openflare:waf_ip_group_sync"
|
||||
// TaskTypeWAFIPGroupSync is the admin task type for WAF IP group sync.
|
||||
TaskTypeWAFIPGroupSync = "of_waf_ip_group_sync"
|
||||
|
||||
// UptimeKumaSyncTask synchronizes proxy routes to Uptime Kuma monitors.
|
||||
UptimeKumaSyncTask = "openflare:uptime_kuma_sync"
|
||||
// TaskTypeUptimeKumaSync is the admin task type for Uptime Kuma sync.
|
||||
TaskTypeUptimeKumaSync = "of_uptime_kuma_sync"
|
||||
)
|
||||
|
||||
var (
|
||||
lastUptimeKumaSyncTime time.Time
|
||||
uptimeKumaSyncMutex sync.Mutex
|
||||
)
|
||||
|
||||
// SSLRenewMeta describes the SSL renewal task.
|
||||
var SSLRenewMeta = task.TaskMeta{
|
||||
Type: TaskTypeSSLRenew,
|
||||
AsynqTask: SSLRenewTask,
|
||||
Name: "OpenFlare SSL 自动续期",
|
||||
Description: "扫描即将到期的 ACME 证书并触发自动续期",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
Retryable: true,
|
||||
}
|
||||
|
||||
// DatabaseAutoCleanupMeta describes the observability auto-cleanup task.
|
||||
var DatabaseAutoCleanupMeta = task.TaskMeta{
|
||||
Type: TaskTypeDatabaseAutoCleanup,
|
||||
AsynqTask: DatabaseAutoCleanupTask,
|
||||
Name: "OpenFlare 可观测数据自动清理",
|
||||
Description: "按保留天数清理访问日志、性能快照与请求聚合数据",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
Retryable: true,
|
||||
}
|
||||
|
||||
// WAFIPGroupSyncMeta describes the WAF IP group sync task.
|
||||
var WAFIPGroupSyncMeta = task.TaskMeta{
|
||||
Type: TaskTypeWAFIPGroupSync,
|
||||
AsynqTask: WAFIPGroupSyncTask,
|
||||
Name: "OpenFlare WAF IP 组同步",
|
||||
Description: "同步到期的自动规则与订阅类型 WAF IP 组",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
Retryable: true,
|
||||
}
|
||||
|
||||
// UptimeKumaSyncMeta describes the Uptime Kuma sync task.
|
||||
var UptimeKumaSyncMeta = task.TaskMeta{
|
||||
Type: TaskTypeUptimeKumaSync,
|
||||
AsynqTask: UptimeKumaSyncTask,
|
||||
Name: "OpenFlare Uptime Kuma 同步",
|
||||
Description: "将启用的代理规则同步到 Uptime Kuma 监控",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
Retryable: true,
|
||||
}
|
||||
|
||||
// SSLRenewHandler renews due TLS certificates.
|
||||
type SSLRenewHandler struct{}
|
||||
|
||||
// Execute runs SSL certificate renewal for all due certificates.
|
||||
func (h *SSLRenewHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
|
||||
task.AppendLog(ctx, "开始扫描待续期证书")
|
||||
if err := tasks.RunSSLRenewJob(ctx); err != nil {
|
||||
task.AppendLog(ctx, "SSL 自动续期失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
msg := "SSL 自动续期任务完成"
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
// DatabaseAutoCleanupHandler prunes observability data when auto-cleanup is enabled.
|
||||
type DatabaseAutoCleanupHandler struct{}
|
||||
|
||||
// Execute runs retention-based cleanup for all observability targets.
|
||||
func (h *DatabaseAutoCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
|
||||
if !model.DatabaseAutoCleanupEnabled {
|
||||
msg := "自动清理未启用,跳过执行"
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "开始执行可观测数据自动清理,保留天数=%d", model.DatabaseAutoCleanupRetentionDays)
|
||||
summary, err := tasks.RunDatabaseAutoCleanupOnce(time.Now())
|
||||
if err != nil {
|
||||
task.AppendLog(ctx, "可观测数据自动清理失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
if summary == nil {
|
||||
msg := "自动清理未启用,跳过执行"
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
var totalDeleted int64
|
||||
for _, item := range summary.Results {
|
||||
totalDeleted += item.DeletedCount
|
||||
task.AppendLog(ctx, "清理 %s:删除 %d 条", item.TargetLabel, item.DeletedCount)
|
||||
}
|
||||
|
||||
msg := fmt.Sprintf(
|
||||
"可观测数据自动清理完成,保留 %d 天,共删除 %d 条",
|
||||
summary.RetentionDays,
|
||||
totalDeleted,
|
||||
)
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
// WAFIPGroupSyncHandler syncs due WAF IP groups to agents.
|
||||
type WAFIPGroupSyncHandler struct{}
|
||||
|
||||
// Execute syncs all due automatic/subscription WAF IP groups.
|
||||
func (h *WAFIPGroupSyncHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
|
||||
task.AppendLog(ctx, "开始同步到期的 WAF IP 组")
|
||||
if err := waf.SyncDueWAFIPGroups(ctx); err != nil {
|
||||
task.AppendLog(ctx, "WAF IP 组同步失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
msg := "WAF IP 组同步完成"
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
// UptimeKumaSyncHandler synchronizes proxy routes to Uptime Kuma.
|
||||
type UptimeKumaSyncHandler struct{}
|
||||
|
||||
// Execute runs Uptime Kuma sync when integration is enabled and the interval has elapsed.
|
||||
func (h *UptimeKumaSyncHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
|
||||
if !model.UptimeKumaEnabled {
|
||||
msg := "Uptime Kuma 集成未启用,跳过执行"
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
interval := model.UptimeKumaSyncInterval
|
||||
if interval <= 0 {
|
||||
interval = 5
|
||||
}
|
||||
if time.Since(lastUptimeKumaSyncTime) < time.Duration(interval)*time.Minute {
|
||||
msg := fmt.Sprintf("距上次同步不足 %d 分钟,跳过执行", interval)
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
if !uptimeKumaSyncMutex.TryLock() {
|
||||
msg := "Uptime Kuma 同步任务正在执行,跳过本次调度"
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
defer uptimeKumaSyncMutex.Unlock()
|
||||
|
||||
task.AppendLog(ctx, "开始同步代理规则到 Uptime Kuma")
|
||||
if err := uptimekuma.SyncToUptimeKuma(ctx); err != nil {
|
||||
task.AppendLog(ctx, "Uptime Kuma 同步失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
lastUptimeKumaSyncTime = time.Now()
|
||||
msg := "Uptime Kuma 同步完成"
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user