This commit is contained in:
ryan
2026-06-08 20:08:14 +08:00
parent 02f458856d
commit cd3d0c9f82
61 changed files with 2180 additions and 1891 deletions
+3 -2
View File
@@ -21,6 +21,7 @@ import (
"fmt"
"net/http"
"github.com/Rain-kl/Wavelet/internal/apps/admin"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/gin-gonic/gin"
@@ -217,7 +218,7 @@ func DeleteAuthSource(c *gin.Context) {
func parseSourceID(c *gin.Context) (uint64, error) {
raw := c.Param("id")
if raw == "" {
return 0, errors.New("认证源 ID 无效")
return 0, errors.New(admin.InvalidAuthSourceID)
}
source, err := model.GetAuthSourceByName(raw)
if err == nil {
@@ -225,7 +226,7 @@ func parseSourceID(c *gin.Context) (uint64, error) {
}
var id uint64
if _, scanErr := fmt.Sscanf(raw, "%d", &id); scanErr != nil || id == 0 {
return 0, errors.New("认证源 ID 无效")
return 0, errors.New(admin.InvalidAuthSourceID)
}
return id, nil
}
+4 -1
View File
@@ -18,5 +18,8 @@ limitations under the License.
package admin
const (
AdminRequired = "未经授权访问"
AdminRequired = "未经授权访问"
InvalidAuthSourceID = "认证源 ID 无效"
InvalidCursorParam = "无效的 cursor 参数"
InvalidTaskExecutionID = "无效的任务执行记录 ID"
)
+2 -1
View File
@@ -21,6 +21,7 @@ import (
"encoding/json"
"net/http"
"github.com/Rain-kl/Wavelet/internal/apps/admin"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/gin-gonic/gin"
@@ -53,7 +54,7 @@ func GetLogs(c *gin.Context) {
var cursor, limit int
if _, err := parsePositiveInt(cursorStr, &cursor); err != nil {
c.JSON(http.StatusBadRequest, util.Err("无效的 cursor 参数"))
c.JSON(http.StatusBadRequest, util.Err(admin.InvalidCursorParam))
return
}
if _, err := parsePositiveInt(limitStr, &limit); err != nil || limit <= 0 {
@@ -227,52 +227,6 @@ func UpdateSystemConfig(c *gin.Context) {
c.JSON(http.StatusOK, util.OKNil())
}
// DeleteSystemConfig 删除系统配置
// @Summary 删除系统配置
// @Description 根据配置键删除对应配置,同时从 Redis 中移除对应缓存,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param key path string true "配置键"
// @Success 200 {object} util.ResponseAny{data=string} "删除成功"
// @Failure 401 {object} util.ResponseAny "未登录"
// @Failure 403 {object} util.ResponseAny "无管理员权限"
// @Failure 404 {object} util.ResponseAny "配置不存在"
// @Failure 500 {object} util.ResponseAny "内部错误"
// @Router /api/v1/admin/system-configs/{key} [delete]
func DeleteSystemConfig(c *gin.Context) {
key := c.Param("key")
// 检查配置是否存在
var config model.SystemConfig
if err := db.DB(c.Request.Context()).Where("key = ?", key).First(&config).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
c.JSON(http.StatusNotFound, util.Err(SystemConfigNotFound))
} else {
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
}
return
}
if err := db.DB(c.Request.Context()).Transaction(func(tx *gorm.DB) error {
// 删除配置
if err := tx.Delete(&config).Error; err != nil {
return err
}
if err := db.Redis.HDel(c.Request.Context(), db.PrefixedKey(model.SystemConfigRedisHashKey), key).Err(); err != nil {
return err
}
return nil
}); err != nil {
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
return
}
c.JSON(http.StatusOK, util.OKNil())
}
// TestSMTPRequest 测试 SMTP 配置请求
type TestSMTPRequest struct {
SMTPHost string `json:"smtp_host" binding:"required,max=255"`
@@ -56,7 +56,6 @@ func setupTestRouter(authUser *model.User) *gin.Engine {
{
systemConfigRouter.GET("", GetSystemConfig)
systemConfigRouter.PUT("", UpdateSystemConfig)
systemConfigRouter.DELETE("", DeleteSystemConfig)
}
return r
@@ -263,48 +262,6 @@ func TestUpdateSystemConfig(t *testing.T) {
})
}
func TestDeleteSystemConfig(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("delete successfully", func(t *testing.T) {
req, _ := http.NewRequest("DELETE", "/api/v1/admin/system-configs/"+model.ConfigKeySiteName, 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())
}
// Verify database
var count int64
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySiteName).Count(&count)
if count != 0 {
t.Error("config still exists in DB")
}
// Verify Redis Cache removal
var redisConfig model.SystemConfig
err := db.HGetJSON(context.Background(), model.SystemConfigRedisHashKey, model.ConfigKeySiteName, &redisConfig)
if err == nil {
t.Error("config should have been deleted from Redis cache")
}
})
t.Run("delete non-existent config", func(t *testing.T) {
req, _ := http.NewRequest("DELETE", "/api/v1/admin/system-configs/invalid_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 TestTestSMTP(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
+3 -2
View File
@@ -24,6 +24,7 @@ import (
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/admin"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
taskhandlers "github.com/Rain-kl/Wavelet/internal/task/handlers"
@@ -156,7 +157,7 @@ func ListTaskExecutions(c *gin.Context) {
func GetTaskExecution(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, util.Err("无效的任务执行记录 ID"))
c.JSON(http.StatusBadRequest, util.Err(admin.InvalidTaskExecutionID))
return
}
@@ -186,7 +187,7 @@ func GetTaskExecution(c *gin.Context) {
func RetryTask(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, util.Err("无效的任务执行记录 ID"))
c.JSON(http.StatusBadRequest, util.Err(admin.InvalidTaskExecutionID))
return
}
+22
View File
@@ -0,0 +1,22 @@
/*
Copyright 2026 Arctel.net
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package cap
const (
errCapTokenMissing = "验证码验证失败,缺少验证码凭证"
errCapTokenInvalidOrExpired = "验证码校验失败或已过期,请重试"
)
+2 -2
View File
@@ -35,13 +35,13 @@ func VerifyMiddleware(mgr *caputil.Manager, scope string, enabledFunc func() boo
token := c.GetHeader("X-Cap-Token")
if token == "" {
c.AbortWithStatusJSON(http.StatusUnauthorized, util.Err("验证码验证失败,缺少验证码凭证"))
c.AbortWithStatusJSON(http.StatusUnauthorized, util.Err(errCapTokenMissing))
return
}
valid, err := mgr.VerifyToken(c.Request.Context(), token, scope)
if err != nil || !valid {
c.AbortWithStatusJSON(http.StatusUnauthorized, util.Err("验证码校验失败或已过期,请重试"))
c.AbortWithStatusJSON(http.StatusUnauthorized, util.Err(errCapTokenInvalidOrExpired))
return
}
+7 -3
View File
@@ -23,9 +23,13 @@ import (
)
const (
UserNameKey = "username"
UserIDKey = "user_id"
UserObjKey = "user_obj"
UserNameKey = "username"
UserIDKey = "user_id"
UserObjKey = "user_obj"
PendingOAuthSourceIDKey = "pending_oauth_source_id"
PendingOAuthExternalIDKey = "pending_oauth_external_id"
PendingOAuthExternalUsernameKey = "pending_oauth_external_username"
PendingOAuthEmailKey = "pending_oauth_email"
)
const (
+12 -3
View File
@@ -18,7 +18,16 @@ limitations under the License.
package oauth
const (
InvalidState = "非法登录请求"
IDTokenVerifyFailed = "ID Token 验证失败"
NonceMismatch = "nonce 不匹配,可能存在重放攻击"
InvalidState = "非法登录请求"
IDTokenVerifyFailed = "ID Token 验证失败"
IDTokenVerifyFailedFormat = "%s: %w"
NonceMismatch = "nonce 不匹配,可能存在重放攻击"
NoActiveAuthSource = "未配置可用认证源"
ServerAddressMissing = "服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试"
AuthSourceRequired = "认证源不能为空"
DiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
UsernameGenerateFailed = "无法生成可用用户名"
UsernameFromSourceFailed = "无法从认证源获取用户名"
AuthSourceDisabled = "认证源未启用"
InvalidExternalAccountBindingID = "绑定记录 ID 无效"
)
+49 -145
View File
@@ -267,7 +267,50 @@ func generateMockIDToken(issuer, sub, aud, nonce, username, email, name string)
// -----------------------------------------------------------------------------
// Test Helpers
// -----------------------------------------------------------------------------
func newMockOIDCClient(issuer, clientID, expectedState, sub, username, email, name string) *http.Client {
cleanIssuer := strings.TrimRight(issuer, "/")
return &http.Client{
Transport: &mockRoundTripper{
roundTripFunc: func(req *http.Request) (*http.Response, error) {
urlStr := req.URL.String()
if req.Method == http.MethodGet && strings.Contains(urlStr, "/.well-known/openid-configuration") {
body := fmt.Sprintf(`{
"issuer": %q,
"authorization_endpoint": %q,
"token_endpoint": %q,
"jwks_uri": %q,
"response_types_supported": ["code"],
"subject_types_supported": ["public"],
"id_token_signing_alg_values_supported": ["RS256"]
}`, cleanIssuer, cleanIssuer+"/oauth2/authorize", cleanIssuer+"/oauth2/token", cleanIssuer+"/oauth2/keys")
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
Header: make(http.Header),
}, nil
}
if req.Method == http.MethodGet && (strings.Contains(urlStr, "/keys") || strings.Contains(urlStr, "/jwks")) {
jwksJSON, _ := json.Marshal(testJWKS)
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(bytes.NewReader(jwksJSON)),
Header: make(http.Header),
}, nil
}
if req.Method == http.MethodPost && (strings.Contains(urlStr, "/token") || strings.Contains(urlStr, "/access_token")) {
idToken := generateMockIDToken(issuer, sub, clientID, expectedState, username, email, name)
body := fmt.Sprintf(`{"access_token":"mock_access_token","token_type":"Bearer","expires_in":3600,"id_token":"%s"}`, idToken)
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
Header: make(http.Header),
}, nil
}
return nil, fmt.Errorf("unexpected mock request: %s %s", req.Method, req.URL)
},
},
}
}
func setupTestDB(t *testing.T) *gorm.DB {
dbConn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
@@ -593,29 +636,7 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
state := uuid.NewString()
// 1. Mock the outgoing HTTP client for token exchange and user info fetching
httpMock := &http.Client{
Transport: &mockRoundTripper{
roundTripFunc: func(req *http.Request) (*http.Response, error) {
if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/.well-known/openid-configuration") {
return oidcDiscoveryResponse(), nil
}
if req.Method == http.MethodGet && req.URL.String() == testJWKSURL {
return jwksResponse(), nil
}
// Handle Token Exchange
if req.Method == http.MethodPost && req.URL.String() == testTokenURL {
idToken := generateMockIDToken(testIssuerURL, "88888", testClientID, state, "test_oauth_user", "oauth@linux.do", "Oauth Test User")
body := fmt.Sprintf(`{"access_token":"mock_access_token","token_type":"Bearer","expires_in":3600,"id_token":"%s"}`, idToken)
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
Header: make(http.Header),
}, nil
}
return nil, fmt.Errorf("unexpected outgoing request: %s %s", req.Method, req.URL)
},
},
}
httpMock := newMockOIDCClient(testIssuerURL, testClientID, state, "88888", "test_oauth_user", "oauth@linux.do", "Oauth Test User")
util.SetHTTPClient(httpMock)
router := setupTestRouter(dbConn, mockRedis, httpMock)
@@ -679,27 +700,7 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
mockRedis.Set(context.Background(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state2)), payloadValue, OAuthStateCacheKeyExpiration)
// Callback with same username but different external ID (99999)
httpMock2 := &http.Client{
Transport: &mockRoundTripper{
roundTripFunc: func(req *http.Request) (*http.Response, error) {
if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/.well-known/openid-configuration") {
return oidcDiscoveryResponse(), nil
}
if req.Method == http.MethodGet && req.URL.String() == testJWKSURL {
return jwksResponse(), nil
}
if req.Method == http.MethodPost && req.URL.String() == testTokenURL {
idToken := generateMockIDToken(testIssuerURL, "99999", testClientID, state2, "test_oauth_user", "another@linux.do", "Another User")
body := fmt.Sprintf(`{"access_token":"mock_access_token","token_type":"Bearer","expires_in":3600,"id_token":"%s"}`, idToken)
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
}, nil
}
return nil, fmt.Errorf("unexpected outgoing request")
},
},
}
httpMock2 := newMockOIDCClient(testIssuerURL, testClientID, state2, "99999", "test_oauth_user", "another@linux.do", "Another User")
util.SetHTTPClient(httpMock2)
// Create another router for this mock client
@@ -740,28 +741,7 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
})
mockRedis.Set(context.Background(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state4)), payloadValue4, OAuthStateCacheKeyExpiration)
httpMock4 := &http.Client{
Transport: &mockRoundTripper{
roundTripFunc: func(req *http.Request) (*http.Response, error) {
if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/.well-known/openid-configuration") {
return oidcDiscoveryResponse(), nil
}
if req.Method == http.MethodGet && req.URL.String() == testJWKSURL {
return jwksResponse(), nil
}
if req.Method == http.MethodPost && req.URL.String() == testTokenURL {
idToken := generateMockIDToken(testIssuerURL, "77777", testClientID, state4, "need_bind_user", "needbind@linux.do", "Need Bind User")
body := fmt.Sprintf(`{"access_token":"mock_access_token","token_type":"Bearer","expires_in":3600,"id_token":"%s"}`, idToken)
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
Header: make(http.Header),
}, nil
}
return nil, fmt.Errorf("unexpected request")
},
},
}
httpMock4 := newMockOIDCClient(testIssuerURL, testClientID, state4, "77777", "need_bind_user", "needbind@linux.do", "Need Bind User")
util.SetHTTPClient(httpMock4)
router4 := setupTestRouter(dbConn, mockRedis, httpMock4)
@@ -824,47 +804,7 @@ func TestCallbackBind(t *testing.T) {
mockRedis.Set(context.Background(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration)
// Mock OIDC discovery, JWKS, and Token exchange for custom source (GitHub)
httpMock := &http.Client{
Transport: &mockRoundTripper{
roundTripFunc: func(req *http.Request) (*http.Response, error) {
// Handle OIDC Discovery Document
if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/.well-known/openid-configuration") {
body := `{
"issuer": "https://github.com",
"authorization_endpoint": "https://github.com/login/oauth/authorize",
"token_endpoint": "https://github.com/login/oauth/access_token",
"jwks_uri": "https://github.com/oauth/keys",
"response_types_supported": ["code"],
"subject_types_supported": ["public"],
"id_token_signing_alg_values_supported": ["RS256"]
}`
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
}, nil
}
// Handle JWKS Key Set Fetch
if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/oauth/keys") {
jwksJSON, _ := json.Marshal(testJWKS)
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(bytes.NewReader(jwksJSON)),
}, nil
}
// Handle Token Exchange
if req.Method == http.MethodPost && strings.Contains(req.URL.String(), "/login/oauth/access_token") {
// Generate signed RS256 token matching issuer and aud
idToken := generateMockIDToken("https://github.com", "github_user_123", "gh_client", state, "github_tester", "tester@github.com", "GitHub Tester")
body := fmt.Sprintf(`{"access_token":"mock_access_token","token_type":"Bearer","id_token":"%s"}`, idToken)
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
}, nil
}
return nil, fmt.Errorf("unexpected request: %s", req.URL)
},
},
}
httpMock := newMockOIDCClient("https://github.com", "gh_client", state, "github_user_123", "github_tester", "tester@github.com", "GitHub Tester")
util.SetHTTPClient(httpMock)
router := setupTestRouter(dbConn, mockRedis, httpMock)
@@ -951,43 +891,7 @@ func TestCallbackBind(t *testing.T) {
mockRedis.Set(context.Background(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state3)), payloadValue, OAuthStateCacheKeyExpiration)
// Re-sign token for new state (since state serves as OIDC Nonce)
httpMock3 := &http.Client{
Transport: &mockRoundTripper{
roundTripFunc: func(req *http.Request) (*http.Response, error) {
if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/.well-known/openid-configuration") {
body := `{
"issuer": "https://github.com",
"authorization_endpoint": "https://github.com/login/oauth/authorize",
"token_endpoint": "https://github.com/login/oauth/access_token",
"jwks_uri": "https://github.com/oauth/keys",
"response_types_supported": ["code"],
"subject_types_supported": ["public"],
"id_token_signing_alg_values_supported": ["RS256"]
}`
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
}, nil
}
if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/oauth/keys") {
jwksJSON, _ := json.Marshal(testJWKS)
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(bytes.NewReader(jwksJSON)),
}, nil
}
if req.Method == http.MethodPost && strings.Contains(req.URL.String(), "/login/oauth/access_token") {
idToken := generateMockIDToken("https://github.com", "github_user_123", "gh_client", state3, "github_tester", "tester@github.com", "GitHub Tester")
body := fmt.Sprintf(`{"access_token":"mock_access_token","token_type":"Bearer","id_token":"%s"}`, idToken)
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
}, nil
}
return nil, fmt.Errorf("unexpected request")
},
},
}
httpMock3 := newMockOIDCClient("https://github.com", "gh_client", state3, "github_user_123", "github_tester", "tester@github.com", "GitHub Tester")
util.SetHTTPClient(httpMock3)
router3 := setupTestRouter(dbConn, mockRedis, httpMock3)
+13 -13
View File
@@ -87,7 +87,7 @@ func resolveAuthSource(sourceName string) (*model.AuthSource, error) {
return nil, err
}
if len(sources) == 0 {
return nil, errors.New("未配置可用认证源")
return nil, errors.New(NoActiveAuthSource)
}
return &sources[0], nil
}
@@ -122,18 +122,18 @@ func activeLoginSources() []AuthSourceView {
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
var sc model.SystemConfig
if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err != nil || strings.TrimSpace(sc.Value) == "" {
return "", errors.New("服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试")
return "", errors.New(ServerAddressMissing)
}
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("认证源不能为空")
return nil, nil, errors.New(AuthSourceRequired)
}
if source.OpenIDDiscoveryURL == "" {
return nil, nil, errors.New("OIDC 认证源必须配置 Discovery URL")
return nil, nil, errors.New(DiscoveryURLRequired)
}
// Clean the issuer URL (trim /.well-known/openid-configuration if configured by mistake)
@@ -194,7 +194,7 @@ func uniqueUsername(ctx context.Context, base string) (string, error) {
}
candidate = fmt.Sprintf("%s-%d", base, i+1)
}
return "", errors.New("无法生成可用用户名")
return "", errors.New(UsernameGenerateFailed)
}
func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code string, nonce string, redirectURL string) (*model.OAuthUserInfo, error) {
@@ -213,7 +213,7 @@ func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code stri
if rawIDToken, ok := token.Extra("id_token").(string); ok {
idToken, verifyErr := verifier.Verify(ctx, rawIDToken)
if verifyErr != nil {
return nil, fmt.Errorf("%s: %w", IDTokenVerifyFailed, verifyErr)
return nil, fmt.Errorf(IDTokenVerifyFailedFormat, IDTokenVerifyFailed, verifyErr)
}
if nonce != "" && idToken.Nonce != nonce {
return nil, errors.New(NonceMismatch)
@@ -257,7 +257,7 @@ func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error {
userInfo.Username = userInfo.Sub
}
if userInfo.Username == "" {
return errors.New("无法从认证源获取用户名")
return errors.New(UsernameFromSourceFailed)
}
if userInfo.Name == "" {
userInfo.Name = userInfo.Username
@@ -359,7 +359,7 @@ func Authorize(c *gin.Context) {
return
}
if !source.IsActive {
c.JSON(http.StatusBadRequest, util.Err("认证源未启用"))
c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled))
return
}
purpose := strings.ToLower(strings.TrimSpace(c.Query("purpose")))
@@ -490,10 +490,10 @@ func Callback(c *gin.Context) {
if !registrationEnabled {
// 如果不允许注册,临时记录到 session 并向前端返回 "need_bind" 状态
session := sessions.Default(c)
session.Set("pending_oauth_source_id", source.ID)
session.Set("pending_oauth_external_id", userInfo.Sub)
session.Set("pending_oauth_external_username", userInfo.Username)
session.Set("pending_oauth_email", userInfo.Email)
session.Set(PendingOAuthSourceIDKey, source.ID)
session.Set(PendingOAuthExternalIDKey, userInfo.Sub)
session.Set(PendingOAuthExternalUsernameKey, userInfo.Username)
session.Set(PendingOAuthEmailKey, userInfo.Email)
if err := session.Save(); err != nil {
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
return
@@ -577,7 +577,7 @@ func DeleteExternalAccount(c *gin.Context) {
rawID := strings.TrimSpace(c.Param("id"))
id, err := strconv.ParseUint(rawID, 10, 64)
if err != nil || id == 0 {
c.JSON(http.StatusBadRequest, util.Err("绑定记录 ID 无效"))
c.JSON(http.StatusBadRequest, util.Err(InvalidExternalAccountBindingID))
return
}
if err := model.DeleteExternalAccountForUser(id, userID); err != nil {
+19
View File
@@ -30,4 +30,23 @@ const (
ErrInvalidFilePath = "非法文件路径"
ErrSaveUploadRecordFailed = "保存上传记录失败"
ErrQueryHistoryUploadFailed = "查询历史上传记录失败"
ErrGenericFileTooLarge = "文件大小不能超过 32MB"
ErrFileContentExtensionMismatch = "文件内容与扩展名不匹配,可能包含安全风险"
ErrFileValidationFailed = "文件校验失败"
ErrInvalidMetadataJSON = "元数据 JSON 格式不合法"
ErrInvalidFileID = "无效的文件 ID"
ErrQueryUploadRecordFailed = "查询文件记录失败"
ErrInvalidBatchDownloadRequest = "参数绑定失败,请传入有效的文件 ID 数组"
ErrInvalidIDValueFormat = "无效的 ID 值: %s"
ErrRetrieveUploadRecordsFailed = "检索文件记录失败"
ErrNoValidFilesForArchive = "没有找到任何有效的文件记录进行打包"
ErrInvalidParams = "参数错误"
ErrQueryFileCountFailed = "查询文件数量失败"
ErrQueryFileListFailed = "查询文件列表失败"
ErrDeleteFileFailed = "删除文件失败"
ErrS3KeyRequired = "s3 key must not be empty"
ErrS3KeyTooLongFormat = "s3 key exceeds maximum length of %d"
ErrS3KeyStartsWithSlash = "s3 key must not start with /"
ErrS3KeyContainsNullBytes = "s3 key must not contain null bytes"
ErrQueryUnusedUploadsFailed = "查询未使用的上传文件失败: %w"
)
+16 -16
View File
@@ -92,7 +92,7 @@ func UploadFile(c *gin.Context) {
// 校验大小
if header.Size > maxUploadSize {
c.JSON(http.StatusOK, util.Err("文件大小不能超过 32MB"))
c.JSON(http.StatusOK, util.Err(ErrGenericFileTooLarge))
return
}
@@ -145,7 +145,7 @@ func UploadFile(c *gin.Context) {
}
}
if isImageExt && !strings.HasPrefix(mimeType, "image/") {
c.JSON(http.StatusOK, util.Err("文件内容与扩展名不匹配,可能包含安全风险"))
c.JSON(http.StatusOK, util.Err(ErrFileContentExtensionMismatch))
return
}
@@ -179,7 +179,7 @@ func UploadFile(c *gin.Context) {
c.JSON(http.StatusOK, util.OK(newUpload))
return
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
c.JSON(http.StatusOK, util.Err("文件校验失败"))
c.JSON(http.StatusOK, util.Err(ErrFileValidationFailed))
return
}
@@ -188,7 +188,7 @@ func UploadFile(c *gin.Context) {
var meta model.UploadMetadata
if metadataStr != "" {
if err := json.Unmarshal([]byte(metadataStr), &meta); err != nil {
c.JSON(http.StatusOK, util.Err("元数据 JSON 格式不合法"))
c.JSON(http.StatusOK, util.Err(ErrInvalidMetadataJSON))
return
}
}
@@ -281,7 +281,7 @@ func DownloadFile(c *gin.Context) {
idStr := c.Param("id")
uploadID, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
c.JSON(http.StatusOK, util.Err("无效的文件 ID"))
c.JSON(http.StatusOK, util.Err(ErrInvalidFileID))
return
}
@@ -291,7 +291,7 @@ func DownloadFile(c *gin.Context) {
c.AbortWithStatus(http.StatusNotFound)
return
}
c.JSON(http.StatusOK, util.Err("查询文件记录失败"))
c.JSON(http.StatusOK, util.Err(ErrQueryUploadRecordFailed))
return
}
@@ -339,7 +339,7 @@ func BatchDownloadFiles(c *gin.Context) {
var req batchDownloadRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusOK, util.Err("参数绑定失败,请传入有效的文件 ID 数组"))
c.JSON(http.StatusOK, util.Err(ErrInvalidBatchDownloadRequest))
return
}
@@ -348,7 +348,7 @@ func BatchDownloadFiles(c *gin.Context) {
for _, idStr := range req.IDs {
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
c.JSON(http.StatusOK, util.Err(fmt.Sprintf("无效的 ID 值: %s", idStr)))
c.JSON(http.StatusOK, util.Err(fmt.Sprintf(ErrInvalidIDValueFormat, idStr)))
return
}
ids = append(ids, id)
@@ -357,12 +357,12 @@ func BatchDownloadFiles(c *gin.Context) {
// 查库获取所有匹配且正常的文件记录
var uploads []model.Upload
if err := db.DB(ctx).Where("id IN ? AND status IN (?, ?)", ids, model.UploadStatusPending, model.UploadStatusUsed).Find(&uploads).Error; err != nil {
c.JSON(http.StatusOK, util.Err("检索文件记录失败"))
c.JSON(http.StatusOK, util.Err(ErrRetrieveUploadRecordsFailed))
return
}
if len(uploads) == 0 {
c.JSON(http.StatusOK, util.Err("没有找到任何有效的文件记录进行打包"))
c.JSON(http.StatusOK, util.Err(ErrNoValidFilesForArchive))
return
}
@@ -458,7 +458,7 @@ func ListMyFiles(c *gin.Context) {
var req listMyFilesRequest
if err := c.ShouldBindQuery(&req); err != nil {
c.JSON(http.StatusOK, util.Err("参数错误"))
c.JSON(http.StatusOK, util.Err(ErrInvalidParams))
return
}
if req.Page <= 0 {
@@ -483,14 +483,14 @@ func ListMyFiles(c *gin.Context) {
var total int64
if err := query.Count(&total).Error; err != nil {
c.JSON(http.StatusOK, util.Err("查询文件数量失败"))
c.JSON(http.StatusOK, util.Err(ErrQueryFileCountFailed))
return
}
var items []model.Upload
offset := (req.Page - 1) * req.PageSize
if err := query.Order("created_at DESC").Offset(offset).Limit(req.PageSize).Find(&items).Error; err != nil {
c.JSON(http.StatusOK, util.Err("查询文件列表失败"))
c.JSON(http.StatusOK, util.Err(ErrQueryFileListFailed))
return
}
@@ -520,7 +520,7 @@ func DeleteFile(c *gin.Context) {
idStr := c.Param("id")
uploadID, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
c.JSON(http.StatusOK, util.Err("无效的文件 ID"))
c.JSON(http.StatusOK, util.Err(ErrInvalidFileID))
return
}
@@ -530,7 +530,7 @@ func DeleteFile(c *gin.Context) {
c.AbortWithStatus(http.StatusNotFound)
return
}
c.JSON(http.StatusOK, util.Err("查询文件记录失败"))
c.JSON(http.StatusOK, util.Err(ErrQueryUploadRecordFailed))
return
}
@@ -541,7 +541,7 @@ func DeleteFile(c *gin.Context) {
}
if err := db.DB(ctx).Model(&upload).Update("status", model.UploadStatusDeleted).Error; err != nil {
c.JSON(http.StatusOK, util.Err("删除文件失败"))
c.JSON(http.StatusOK, util.Err(ErrDeleteFileFailed))
return
}
+1 -1
View File
@@ -53,7 +53,7 @@ func (h *CleanupUnusedUploadsHandler) Execute(ctx context.Context, payload []byt
Limit(batchSize).
Find(&unusedUploads).Error; err != nil {
task.AppendLog(ctx, "查询未使用的上传文件失败: %v", err)
return nil, fmt.Errorf("查询未使用的上传文件失败: %w", err)
return nil, fmt.Errorf(ErrQueryUnusedUploadsFailed, err)
}
// 没有更多数据,退出循环
+4 -4
View File
@@ -27,19 +27,19 @@ const maxS3KeyLength = 1024
// ValidateS3Key validates an S3 object key for safety.
func ValidateS3Key(key string) error {
if key == "" {
return fmt.Errorf("s3 key must not be empty")
return fmt.Errorf(ErrS3KeyRequired)
}
if len(key) > maxS3KeyLength {
return fmt.Errorf("s3 key exceeds maximum length of %d", maxS3KeyLength)
return fmt.Errorf(ErrS3KeyTooLongFormat, maxS3KeyLength)
}
if strings.HasPrefix(key, "/") {
return fmt.Errorf("s3 key must not start with /")
return fmt.Errorf(ErrS3KeyStartsWithSlash)
}
if strings.Contains(key, "\x00") {
return fmt.Errorf("s3 key must not contain null bytes")
return fmt.Errorf(ErrS3KeyContainsNullBytes)
}
return nil
+9 -9
View File
@@ -78,13 +78,13 @@ func CreateAccessToken(c *gin.Context) {
var req createTokenRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusOK, util.Err("参数绑定失败"))
c.JSON(http.StatusOK, util.Err(errBindParamsFailed))
return
}
req.Name = strings.TrimSpace(req.Name)
if req.Name == "" {
c.JSON(http.StatusOK, util.Err("令牌名称不能为空"))
c.JSON(http.StatusOK, util.Err(errTokenNameRequired))
return
}
@@ -101,14 +101,14 @@ func CreateAccessToken(c *gin.Context) {
}
if int(count) >= maxLimit {
c.JSON(http.StatusOK, util.Err("已达到访问令牌最大创建数量限制"))
c.JSON(http.StatusOK, util.Err(errAccessTokenLimitReached))
return
}
// 生成 Token
tokenStr, err := model.GenerateTokenString()
if err != nil {
c.JSON(http.StatusOK, util.Err("生成令牌失败"))
c.JSON(http.StatusOK, util.Err(errGenerateTokenFailed))
return
}
@@ -150,7 +150,7 @@ func DeleteAccessToken(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
c.JSON(http.StatusOK, util.Err("无效的令牌ID"))
c.JSON(http.StatusOK, util.Err(errInvalidTokenID))
return
}
@@ -161,7 +161,7 @@ func DeleteAccessToken(c *gin.Context) {
}
if tx.RowsAffected == 0 {
c.JSON(http.StatusOK, util.Err("令牌不存在或无权操作"))
c.JSON(http.StatusOK, util.Err(errTokenNotFoundOrForbidden))
return
}
@@ -185,20 +185,20 @@ func RotateAccessToken(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
c.JSON(http.StatusOK, util.Err("无效的令牌ID"))
c.JSON(http.StatusOK, util.Err(errInvalidTokenID))
return
}
var tokenRecord model.AccessToken
if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, currUser.ID).First(&tokenRecord).Error; err != nil {
c.JSON(http.StatusOK, util.Err("令牌不存在或无权操作"))
c.JSON(http.StatusOK, util.Err(errTokenNotFoundOrForbidden))
return
}
// 生成新的 Token
newTokenStr, err := model.GenerateTokenString()
if err != nil {
c.JSON(http.StatusOK, util.Err("生成令牌失败"))
c.JSON(http.StatusOK, util.Err(errGenerateTokenFailed))
return
}
+106 -118
View File
@@ -77,6 +77,64 @@ func generateVerificationCode() string {
return fmt.Sprintf("%06d", n.Int64()+100000)
}
func getEmailCodeKey(scene, email string) string {
return fmt.Sprintf("email_code:%s:%s", scene, email)
}
func getEmailCooldownKey(email string) string {
return fmt.Sprintf("email_code:cooldown:%s", email)
}
func sendEmailVerificationCode(ctx context.Context, email, scene, templateName string) error {
code := generateVerificationCode()
codeKey := getEmailCodeKey(scene, email)
cooldownKey := getEmailCooldownKey(email)
// 使用模板管理获取并渲染邮件标题和正文。模板缺失或渲染失败时不发送验证码。
emailSubject, emailBody, err := model.RenderTemplate(
ctx,
templateName,
map[string]any{"Code": code},
)
if err != nil {
return fmt.Errorf(errRenderEmailTemplateFailed, err)
}
// 存验证码,5分钟有效
if err := db.SetJSON(ctx, codeKey, code, 5*time.Minute); err != nil {
return fmt.Errorf(errGenerateEmailCodeFailed)
}
// 存冷却,60秒有效
_ = db.SetJSON(ctx, cooldownKey, "1", 60*time.Second)
// 构建异步邮件发送任务
payload := SendEmailPayload{
To: email,
Subject: emailSubject,
Body: emailBody,
}
payloadBytes, _ := json.Marshal(payload)
_, err = task.DispatchTask(ctx, task.TaskTypeSendEmail, payloadBytes, "system")
if err != nil {
return fmt.Errorf(errDispatchEmailTaskFailed)
}
return nil
}
func verifyEmailCode(ctx context.Context, email, scene, code string) bool {
codeKey := getEmailCodeKey(scene, email)
var storedCode string
if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil {
return false
}
if storedCode != code {
return false
}
// 验证成功,删除验证码
_ = db.Redis.Del(ctx, db.PrefixedKey(codeKey)).Err()
return true
}
func isPasswordLoginEnabled() bool {
enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyPasswordLoginEnabled)
if err != nil {
@@ -124,7 +182,7 @@ func setLoginSession(c *gin.Context, user *model.User) error {
// @Router /api/v1/user/login [post]
func Login(c *gin.Context) {
if !isPasswordLoginEnabled() {
c.JSON(http.StatusOK, util.Err("管理员关闭了密码登录"))
c.JSON(http.StatusOK, util.Err(errPasswordLoginDisabled))
return
}
var req loginRequest
@@ -134,14 +192,14 @@ func Login(c *gin.Context) {
}
req.Username = strings.TrimSpace(req.Username)
if req.Username == "" || req.Password == "" {
c.JSON(http.StatusOK, util.Err("无效的参数"))
c.JSON(http.StatusOK, util.Err(errInvalidParams))
return
}
var user model.User
ctx := c.Request.Context()
if err := db.DB(ctx).Where("username = ?", req.Username).First(&user).Error; err != nil {
c.JSON(http.StatusOK, util.Err("用户名或密码错误"))
c.JSON(http.StatusOK, util.Err(errUsernameOrPasswordWrong))
return
}
if !user.IsActive {
@@ -150,79 +208,43 @@ func Login(c *gin.Context) {
}
// 判定是否是明文密码存储
isPlaintext := !(strings.HasPrefix(user.Password, "$2a$") || strings.HasPrefix(user.Password, "$2b$") || strings.HasPrefix(user.Password, "$2y$"))
isPlaintext := !user.IsPasswordEncrypted()
if !user.CheckPassword(req.Password) {
c.JSON(http.StatusOK, util.Err("用户名或密码错误"))
c.JSON(http.StatusOK, util.Err(errUsernameOrPasswordWrong))
return
}
if isEmailLoginVerificationEnabled() {
if user.Email == "" {
c.JSON(http.StatusOK, util.Err("该账号未绑定邮箱,请联系管理员绑定邮箱后再登录"))
c.JSON(http.StatusOK, util.Err(errLoginEmailMissing))
return
}
if req.Code == "" {
// 校验 Redis 发送冷却时间
cooldownKey := fmt.Sprintf("email_code:cooldown:%s", user.Email)
cooldownKey := getEmailCooldownKey(user.Email)
var temp string
err := db.GetJSON(ctx, cooldownKey, &temp)
if err != nil {
// 没有冷却,触发验证码发送
code := generateVerificationCode()
codeKey := fmt.Sprintf("email_code:login:%s", user.Email)
// 存验证码,5分钟有效
if err := db.SetJSON(ctx, codeKey, code, 5*time.Minute); err != nil {
c.JSON(http.StatusOK, util.Err("生成验证码失败,请重试"))
return
}
// 存冷却,60秒有效
_ = db.SetJSON(ctx, cooldownKey, "1", 60*time.Second)
// 使用模板管理获取并渲染邮件标题和正文
emailSubject, emailBody := model.RenderTemplate(
ctx,
"login_email",
map[string]any{"Code": code},
"Wavelet 登录验证码",
fmt.Sprintf("<h3>Wavelet 登录验证</h3><p>您的登录验证码为:<strong>%s</strong>,5分钟内有效,请勿将验证码泄露给他人。</p>", code),
)
// 构建异步邮件发送任务
payload := SendEmailPayload{
To: user.Email,
Subject: emailSubject,
Body: emailBody,
}
payloadBytes, _ := json.Marshal(payload)
_, err = task.DispatchTask(ctx, task.TaskTypeSendEmail, payloadBytes, "system")
if err != nil {
c.JSON(http.StatusOK, util.Err("投递验证邮件发送任务失败,请重试"))
if err := sendEmailVerificationCode(ctx, user.Email, "login", "login_email"); err != nil {
c.JSON(http.StatusOK, util.Err(err.Error()))
return
}
}
// 脱敏邮箱并返回错误,提示前端需要输入验证码
maskedEmail := util.MaskEmail(user.Email)
c.JSON(http.StatusOK, util.Err("need_email_code:"+maskedEmail))
c.JSON(http.StatusOK, util.Err(errNeedEmailCodePrefix+maskedEmail))
return
}
// 校验验证码
codeKey := fmt.Sprintf("email_code:login:%s", user.Email)
var storedCode string
if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil {
c.JSON(http.StatusOK, util.Err("验证码错误或已过期"))
if !verifyEmailCode(ctx, user.Email, "login", req.Code) {
c.JSON(http.StatusOK, util.Err(errEmailCodeInvalidOrExpired))
return
}
if storedCode != req.Code {
c.JSON(http.StatusOK, util.Err("验证码错误或已过期"))
return
}
// 验证成功,删除验证码
_ = db.Redis.Del(ctx, db.PrefixedKey(codeKey)).Err()
}
session := sessions.Default(c)
@@ -232,7 +254,7 @@ func Login(c *gin.Context) {
if isPlaintext {
if err := user.SetEncryptedPassword(req.Password); err == nil {
if err := db.DB(ctx).Model(&user).Update("password", user.Password).Error; err != nil {
c.JSON(http.StatusOK, util.Err("升级密码安全算法失败,请重试"))
c.JSON(http.StatusOK, util.Err(errPasswordUpgradeFailed))
return
}
needChangePassword = true
@@ -247,15 +269,15 @@ func Login(c *gin.Context) {
return
}
if err := setLoginSession(c, &user); err != nil {
c.JSON(http.StatusOK, util.Err("无法保存会话信息,请重试"))
c.JSON(http.StatusOK, util.Err(errSaveSessionFailed))
return
}
// 检查是否有未完成的 OAuth/OIDC 绑定
pendingSourceID := session.Get("pending_oauth_source_id")
pendingExternalID := session.Get("pending_oauth_external_id")
pendingExternalUsername := session.Get("pending_oauth_external_username")
pendingEmail := session.Get("pending_oauth_email")
// 检查是否有未完成 of OAuth/OIDC 绑定
pendingSourceID := session.Get(oauth.PendingOAuthSourceIDKey)
pendingExternalID := session.Get(oauth.PendingOAuthExternalIDKey)
pendingExternalUsername := session.Get(oauth.PendingOAuthExternalUsernameKey)
pendingEmail := session.Get(oauth.PendingOAuthEmailKey)
if pendingSourceID != nil && pendingExternalID != nil {
var sourceID uint64
@@ -281,10 +303,10 @@ func Login(c *gin.Context) {
})
}
// 清除 pending 信息
session.Delete("pending_oauth_source_id")
session.Delete("pending_oauth_external_id")
session.Delete("pending_oauth_external_username")
session.Delete("pending_oauth_email")
session.Delete(oauth.PendingOAuthSourceIDKey)
session.Delete(oauth.PendingOAuthExternalIDKey)
session.Delete(oauth.PendingOAuthExternalUsernameKey)
session.Delete(oauth.PendingOAuthEmailKey)
_ = session.Save()
}
@@ -304,7 +326,7 @@ func Login(c *gin.Context) {
// @Router /api/v1/user/register [post]
func Register(c *gin.Context) {
if !isRegistrationEnabled() || !isPasswordRegisterEnabled() {
c.JSON(http.StatusOK, util.Err("管理员关闭了注册"))
c.JSON(http.StatusOK, util.Err(errRegistrationDisabled))
return
}
@@ -322,11 +344,11 @@ func Register(c *gin.Context) {
req.Code = strings.TrimSpace(req.Code)
if req.Username == "" || req.Password == "" {
c.JSON(http.StatusOK, util.Err("无效的参数"))
c.JSON(http.StatusOK, util.Err(errInvalidParams))
return
}
if len(req.Password) < 8 {
c.JSON(http.StatusOK, util.Err("密码长度不能少于 8 位"))
c.JSON(http.StatusOK, util.Err(errPasswordTooShort))
return
}
@@ -335,23 +357,14 @@ func Register(c *gin.Context) {
// 邮箱注册验证校验
if isEmailRegisterVerificationEnabled() {
if req.Email == "" || req.Code == "" {
c.JSON(http.StatusOK, util.Err("邮箱或验证码未填写"))
c.JSON(http.StatusOK, util.Err(errEmailOrCodeRequired))
return
}
codeKey := fmt.Sprintf("email_code:register:%s", req.Email)
var storedCode string
if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil {
c.JSON(http.StatusOK, util.Err("验证码错误或已过期"))
if !verifyEmailCode(ctx, req.Email, "register", req.Code) {
c.JSON(http.StatusOK, util.Err(errEmailCodeInvalidOrExpired))
return
}
if storedCode != req.Code {
c.JSON(http.StatusOK, util.Err("验证码错误或已过期"))
return
}
// 验证通过,删除 Redis 中的验证码
_ = db.Redis.Del(ctx, db.PrefixedKey(codeKey)).Err()
}
user := model.User{
@@ -380,7 +393,7 @@ func Register(c *gin.Context) {
}
if err := setLoginSession(c, &user); err != nil {
c.JSON(http.StatusOK, util.Err("无法保存会话信息,请重试"))
c.JSON(http.StatusOK, util.Err(errSaveSessionFailed))
return
}
@@ -434,36 +447,36 @@ func ChangePassword(c *gin.Context) {
req.NewPassword = strings.TrimSpace(req.NewPassword)
if req.OldPassword == "" || req.NewPassword == "" {
c.JSON(http.StatusOK, util.Err("无效的参数"))
c.JSON(http.StatusOK, util.Err(errInvalidParams))
return
}
if len(req.NewPassword) < 8 {
c.JSON(http.StatusOK, util.Err("新密码长度不能少于 8 位"))
c.JSON(http.StatusOK, util.Err(errNewPasswordTooShort))
return
}
userObj, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
if userObj == nil {
c.JSON(http.StatusUnauthorized, util.Err("请先登录"))
c.JSON(http.StatusUnauthorized, util.Err(errLoginRequired))
return
}
ctx := c.Request.Context()
var dbUser model.User
if err := db.DB(ctx).Where("id = ?", userObj.ID).First(&dbUser).Error; err != nil {
c.JSON(http.StatusOK, util.Err("未找到该用户"))
c.JSON(http.StatusOK, util.Err(errUserNotFound))
return
}
// 校验旧密码
if !dbUser.CheckPassword(req.OldPassword) {
c.JSON(http.StatusOK, util.Err("原密码不正确"))
c.JSON(http.StatusOK, util.Err(errOldPasswordIncorrect))
return
}
// 加密并更新为新密码
if err := dbUser.SetEncryptedPassword(req.NewPassword); err != nil {
c.JSON(http.StatusOK, util.Err("密码加密失败,请重试"))
c.JSON(http.StatusOK, util.Err(errPasswordEncryptFailed))
return
}
@@ -499,12 +512,12 @@ func SendEmailCode(c *gin.Context) {
req.Email = strings.TrimSpace(req.Email)
if req.Email == "" {
c.JSON(http.StatusOK, util.Err("邮箱地址不能为空"))
c.JSON(http.StatusOK, util.Err(errEmailRequired))
return
}
if req.Scene != "register" {
c.JSON(http.StatusOK, util.Err("不支持的验证场景"))
c.JSON(http.StatusOK, util.Err(errUnsupportedEmailScene))
return
}
@@ -517,47 +530,22 @@ func SendEmailCode(c *gin.Context) {
return
}
if count > 0 {
c.JSON(http.StatusOK, util.Err("该邮箱已被注册"))
c.JSON(http.StatusOK, util.Err(errEmailAlreadyRegistered))
return
}
// 2. 校验 Redis 发送冷却时间
cooldownKey := fmt.Sprintf("email_code:cooldown:%s", req.Email)
cooldownKey := getEmailCooldownKey(req.Email)
var temp string
err := db.GetJSON(ctx, cooldownKey, &temp)
if err == nil {
c.JSON(http.StatusOK, util.Err("验证码发送频繁,请稍后再试"))
c.JSON(http.StatusOK, util.Err(errEmailCodeCooldown))
return
}
// 3. 生成并缓存验证码
code := generateVerificationCode()
codeKey := fmt.Sprintf("email_code:register:%s", req.Email)
if err := db.SetJSON(ctx, codeKey, code, 5*time.Minute); err != nil {
c.JSON(http.StatusOK, util.Err("生成验证码失败,请重试"))
return
}
_ = db.SetJSON(ctx, cooldownKey, "1", 60*time.Second)
// 使用模板管理获取并渲染邮件标题和正文
emailSubject, emailBody := model.RenderTemplate(
ctx,
"register_email",
map[string]any{"Code": code},
"Wavelet 注册验证码",
fmt.Sprintf("<h3>Wavelet 注册验证</h3><p>您的注册验证码为:<strong>%s</strong>,5分钟内有效,请勿泄露给他人。</p>", code),
)
// 4. 投递异步邮件发送任务
payload := SendEmailPayload{
To: req.Email,
Subject: emailSubject,
Body: emailBody,
}
payloadBytes, _ := json.Marshal(payload)
_, err = task.DispatchTask(ctx, task.TaskTypeSendEmail, payloadBytes, "system")
if err != nil {
c.JSON(http.StatusOK, util.Err("投递验证邮件发送任务失败,请重试"))
// 3. 发送验证码
if err := sendEmailVerificationCode(ctx, req.Email, "register", "register_email"); err != nil {
c.JSON(http.StatusOK, util.Err(err.Error()))
return
}
@@ -595,14 +583,14 @@ func UpdateProfile(c *gin.Context) {
userObj, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
if userObj == nil {
c.JSON(http.StatusUnauthorized, util.Err("请先登录"))
c.JSON(http.StatusUnauthorized, util.Err(errLoginRequired))
return
}
ctx := c.Request.Context()
var dbUser model.User
if err := db.DB(ctx).Where("id = ?", userObj.ID).First(&dbUser).Error; err != nil {
c.JSON(http.StatusOK, util.Err("未找到该用户"))
c.JSON(http.StatusOK, util.Err(errUserNotFound))
return
}
@@ -610,7 +598,7 @@ func UpdateProfile(c *gin.Context) {
req.Email = strings.TrimSpace(req.Email)
if req.Email != "" && req.Email != dbUser.Email {
if !strings.Contains(req.Email, "@") || !strings.Contains(req.Email, ".") {
c.JSON(http.StatusOK, util.Err("邮箱格式不正确"))
c.JSON(http.StatusOK, util.Err(errEmailFormatInvalid))
return
}
@@ -620,7 +608,7 @@ func UpdateProfile(c *gin.Context) {
return
}
if count > 0 {
c.JSON(http.StatusOK, util.Err("该邮箱已被其他账号绑定"))
c.JSON(http.StatusOK, util.Err(errEmailAlreadyBound))
return
}
}
+39 -1
View File
@@ -17,4 +17,42 @@ limitations under the License.
package user
const ()
const (
errBindParamsFailed = "参数绑定失败"
errInvalidParams = "无效的参数"
errPasswordLoginDisabled = "管理员关闭了密码登录"
errUsernameOrPasswordWrong = "用户名或密码错误"
errLoginEmailMissing = "该账号未绑定邮箱,请联系管理员绑定邮箱后再登录"
errNeedEmailCodePrefix = "need_email_code:"
errEmailCodeInvalidOrExpired = "验证码错误或已过期"
errPasswordUpgradeFailed = "升级密码安全算法失败,请重试"
errSaveSessionFailed = "无法保存会话信息,请重试"
errRegistrationDisabled = "管理员关闭了注册"
errPasswordTooShort = "密码长度不能少于 8 位"
errEmailOrCodeRequired = "邮箱或验证码未填写"
errNewPasswordTooShort = "新密码长度不能少于 8 位"
errLoginRequired = "请先登录"
errUserNotFound = "未找到该用户"
errOldPasswordIncorrect = "原密码不正确"
errPasswordEncryptFailed = "密码加密失败,请重试"
errEmailRequired = "邮箱地址不能为空"
errUnsupportedEmailScene = "不支持的验证场景"
errEmailAlreadyRegistered = "该邮箱已被注册"
errEmailCodeCooldown = "验证码发送频繁,请稍后再试"
errEmailFormatInvalid = "邮箱格式不正确"
errEmailAlreadyBound = "该邮箱已被其他账号绑定"
errRenderEmailTemplateFailed = "渲染验证邮件模板失败:%w"
errGenerateEmailCodeFailed = "生成验证码失败,请重试"
errDispatchEmailTaskFailed = "投递验证邮件发送任务失败,请重试"
errTokenNameRequired = "令牌名称不能为空"
errAccessTokenLimitReached = "已达到访问令牌最大创建数量限制"
errGenerateTokenFailed = "生成令牌失败"
errInvalidTokenID = "无效的令牌ID"
errTokenNotFoundOrForbidden = "令牌不存在或无权操作"
errTaskPayloadRequired = "任务参数不能为空"
errInvalidJSONFormat = "无效的 JSON 格式: %w"
errEmailTaskFieldsRequired = "to、subject、body 不能为空"
errParseEmailPayloadFailed = "解析邮件发送参数失败: %w"
errSMTPConfigIncomplete = "系统 SMTP 邮件服务配置不完整"
errSendMailFailed = "发送邮件失败: %w"
)
+6 -6
View File
@@ -44,12 +44,12 @@ type SendEmailHandler struct{}
// 校验并标准化邮件发送参数,框架在 Admin 下发时自动调用
func (h *SendEmailHandler) ValidatePayload(payload []byte) ([]byte, error) {
if len(payload) == 0 {
return nil, errors.New("任务参数不能为空")
return nil, errors.New(errTaskPayloadRequired)
}
var req SendEmailPayload
if err := json.Unmarshal(payload, &req); err != nil {
return nil, fmt.Errorf("无效的 JSON 格式: %w", err)
return nil, fmt.Errorf(errInvalidJSONFormat, err)
}
req.To = strings.TrimSpace(req.To)
@@ -57,7 +57,7 @@ func (h *SendEmailHandler) ValidatePayload(payload []byte) ([]byte, error) {
req.Body = strings.TrimSpace(req.Body)
if req.To == "" || req.Subject == "" || req.Body == "" {
return nil, errors.New("to、subject、body 不能为空")
return nil, errors.New(errEmailTaskFieldsRequired)
}
return json.Marshal(req)
@@ -68,7 +68,7 @@ func (h *SendEmailHandler) Execute(ctx context.Context, payload []byte) (*task.T
var req SendEmailPayload
if err := json.Unmarshal(payload, &req); err != nil {
task.AppendLog(ctx, "解析邮件发送参数失败: %v", err)
return nil, fmt.Errorf("解析邮件发送参数失败: %w", err)
return nil, fmt.Errorf(errParseEmailPayloadFailed, err)
}
task.AppendLog(ctx, "开始准备发送邮件到: %s, 主题: %s", req.To, req.Subject)
@@ -94,7 +94,7 @@ func (h *SendEmailHandler) Execute(ctx context.Context, payload []byte) (*task.T
}
if smtpHost == "" || smtpPortVal == "" || smtpUsername == "" {
err := errors.New("系统 SMTP 邮件服务配置不完整")
err := errors.New(errSMTPConfigIncomplete)
task.AppendLog(ctx, "发送失败: %v", err)
return nil, err
}
@@ -117,7 +117,7 @@ func (h *SendEmailHandler) Execute(ctx context.Context, payload []byte) (*task.T
err = mail.SendMailHTML(cfg, req.To, req.Subject, req.Body)
if err != nil {
task.AppendLog(ctx, "邮件发送失败: %v", err)
return nil, fmt.Errorf("发送邮件失败: %w", err)
return nil, fmt.Errorf(errSendMailFailed, err)
}
msg := fmt.Sprintf("邮件成功发送至: %s", req.To)
+24
View File
@@ -0,0 +1,24 @@
/*
Copyright 2026 Arctel.net
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package db
const (
errRedisHashSetFailed = "failed to set redis hash: %w"
errUnmarshalDataFailed = "failed to unmarshal data: %w"
errMarshalDataFailed = "failed to marshal data: %w"
errRedisKeySetFailed = "failed to set redis key: %w"
)
+5 -5
View File
@@ -123,7 +123,7 @@ func HSetJSON[T any](ctx context.Context, hashKey, fieldKey string, data T) erro
}
if err := Redis.HSet(ctx, PrefixedKey(hashKey), fieldKey, jsonData).Err(); err != nil {
return fmt.Errorf("failed to set redis hash: %w", err)
return fmt.Errorf(errRedisHashSetFailed, err)
}
return nil
@@ -141,7 +141,7 @@ func HGetJSON[T any](ctx context.Context, hashKey, fieldKey string, data *T) err
}
if err := json.Unmarshal([]byte(val), data); err != nil {
return fmt.Errorf("failed to unmarshal data: %w", err)
return fmt.Errorf(errUnmarshalDataFailed, err)
}
return nil
@@ -158,7 +158,7 @@ func GetJSON[T any](ctx context.Context, key string, data *T) error {
}
if err := json.Unmarshal(val, data); err != nil {
return fmt.Errorf("failed to unmarshal data: %w", err)
return fmt.Errorf(errUnmarshalDataFailed, err)
}
return nil
@@ -172,11 +172,11 @@ func GetJSON[T any](ctx context.Context, key string, data *T) error {
func SetJSON[T any](ctx context.Context, key string, data T, expiration time.Duration) error {
jsonData, err := json.Marshal(data)
if err != nil {
return fmt.Errorf("failed to marshal data: %w", err)
return fmt.Errorf(errMarshalDataFailed, err)
}
if err := Redis.Set(ctx, PrefixedKey(key), jsonData, expiration).Err(); err != nil {
return fmt.Errorf("failed to set redis key: %w", err)
return fmt.Errorf(errRedisKeySetFailed, err)
}
return nil
+21
View File
@@ -0,0 +1,21 @@
/*
Copyright 2026 Arctel.net
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package logger
const (
errCreateLogFileDirFailed = "[Logger] create log file dir err: %w"
)
+1 -1
View File
@@ -55,7 +55,7 @@ func initWriter() (zapcore.WriteSyncer, error) {
logPath := logConfig.FilePath
logDir := filepath.Dir(logPath)
if err := os.MkdirAll(logDir, 0750); err != nil {
return nil, fmt.Errorf("[Logger] create log file dir err: %w", err)
return nil, fmt.Errorf(errCreateLogFileDirFailed, err)
}
// 配置日志轮转
+13 -13
View File
@@ -91,19 +91,19 @@ func (source *AuthSource) Normalize() {
func (source *AuthSource) Validate() error {
source.Normalize()
if source.Name == "" {
return errors.New("认证源名称不能为空")
return errors.New(errAuthSourceNameRequired)
}
if !authSourceNamePattern.MatchString(source.Name) {
return errors.New("认证源名称只能包含字母、数字、短横线或下划线,且必须以字母或数字开头")
return errors.New(errAuthSourceNameInvalid)
}
if source.Type != AuthSourceTypeOIDC {
return errors.New("认证源类型仅支持 oidc")
return errors.New(errAuthSourceTypeUnsupported)
}
if source.OpenIDDiscoveryURL == "" {
return errors.New("OIDC 认证源必须配置 Discovery URL")
return errors.New(errAuthSourceDiscoveryURLRequired)
}
if source.IsActive && (source.ClientID == "" || source.ClientSecret == "") {
return errors.New("启用认证源前必须配置 Client ID 和 Client Secret")
return errors.New(errAuthSourceClientCredentialsRequired)
}
return nil
}
@@ -137,7 +137,7 @@ func GetActiveAuthSources() ([]AuthSource, error) {
func GetAuthSourceByID(id uint64) (*AuthSource, error) {
if id == 0 {
return nil, errors.New("认证源 ID 不能为空")
return nil, errors.New(errAuthSourceIDRequired)
}
var source AuthSource
if err := db.DB(context.Background()).First(&source, "id = ?", id).Error; err != nil {
@@ -150,7 +150,7 @@ func GetAuthSourceByID(id uint64) (*AuthSource, error) {
func GetAuthSourceByName(name string) (*AuthSource, error) {
name = strings.TrimSpace(name)
if name == "" {
return nil, errors.New("认证源名称不能为空")
return nil, errors.New(errAuthSourceNameRequired)
}
var source AuthSource
if err := db.DB(context.Background()).First(&source, "name = ?", name).Error; err != nil {
@@ -169,7 +169,7 @@ func CreateAuthSource(source *AuthSource) error {
func UpdateAuthSource(source *AuthSource, keepSecret bool) error {
if source.ID == 0 {
return errors.New("认证源 ID 不能为空")
return errors.New(errAuthSourceIDRequired)
}
var current AuthSource
if err := db.DB(context.Background()).First(&current, "id = ?", source.ID).Error; err != nil {
@@ -208,7 +208,7 @@ func ToggleAuthSource(id uint64, isActive bool) error {
func DeleteAuthSource(id uint64) error {
if id == 0 {
return errors.New("认证源 ID 不能为空")
return errors.New(errAuthSourceIDRequired)
}
return db.DB(context.Background()).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("auth_source_id = ?", id).Delete(&ExternalAccount{}).Error; err != nil {
@@ -228,7 +228,7 @@ func FindExternalAccount(sourceID uint64, externalID string) (*ExternalAccount,
func BindExternalAccount(account *ExternalAccount) error {
if account.UserID == 0 || strings.TrimSpace(account.ExternalID) == "" {
return errors.New("外部账号绑定信息不完整")
return errors.New(errExternalAccountBindingIncomplete)
}
account.ExternalID = strings.TrimSpace(account.ExternalID)
account.ExternalUsername = strings.TrimSpace(account.ExternalUsername)
@@ -239,7 +239,7 @@ func BindExternalAccount(account *ExternalAccount) error {
err := tx.Where("auth_source_id = ? AND external_id = ?", account.AuthSourceID, account.ExternalID).First(&current).Error
if err == nil {
if current.UserID != account.UserID {
return errors.New("该外部账号已绑定到其他用户")
return errors.New(errExternalAccountAlreadyBoundToAnother)
}
return tx.Model(&current).Updates(map[string]any{
"external_username": account.ExternalUsername,
@@ -255,7 +255,7 @@ func BindExternalAccount(account *ExternalAccount) error {
func ListExternalAccountsByUserID(userID uint64) ([]ExternalAccountView, error) {
if userID == 0 {
return nil, errors.New("用户 ID 不能为空")
return nil, errors.New(errUserIDRequired)
}
var accounts []ExternalAccount
if err := db.DB(context.Background()).Where("user_id = ?", userID).Order("id asc").Find(&accounts).Error; err != nil {
@@ -296,7 +296,7 @@ func ListExternalAccountsByUserID(userID uint64) ([]ExternalAccountView, error)
func DeleteExternalAccountForUser(id uint64, userID uint64) error {
if id == 0 || userID == 0 {
return errors.New("绑定记录 ID 不能为空")
return errors.New(errExternalAccountBindingIDRequired)
}
return db.DB(context.Background()).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error
}
+43
View File
@@ -0,0 +1,43 @@
/*
Copyright 2026 Arctel.net
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package model
const (
errRegistrationDisabled = "注册已关闭"
errDatabaseNotInitialized = "database not initialized"
errUsernameExists = "用户名已存在"
errEmailAlreadyBound = "该邮箱已被其他账号绑定"
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
errTemplateKeyRequired = "模板标识符不能为空"
errTemplateNameRequired = "模板名称不能为空"
errTemplateContentRequired = "模板内容不能为空"
errTemplateUnavailable = "模板 %s 不存在或不可用: %w"
errTemplateRenderFailed = "模板 %s 渲染失败: %w"
errAuthSourceNameRequired = "认证源名称不能为空"
errAuthSourceNameInvalid = "认证源名称只能包含字母、数字、短横线或下划线,且必须以字母或数字开头"
errAuthSourceTypeUnsupported = "认证源类型仅支持 oidc"
errAuthSourceDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
errAuthSourceClientCredentialsRequired = "启用认证源前必须配置 Client ID 和 Client Secret"
errAuthSourceIDRequired = "认证源 ID 不能为空"
errExternalAccountBindingIncomplete = "外部账号绑定信息不完整"
errExternalAccountAlreadyBoundToAnother = "该外部账号已绑定到其他用户"
errUserIDRequired = "用户 ID 不能为空"
errExternalAccountBindingIDRequired = "绑定记录 ID 不能为空"
)
+5 -5
View File
@@ -85,7 +85,7 @@ func (sc *SystemConfig) GetByKey(ctx context.Context, key string) error {
// 查数据库
database := db.DB(ctx)
if database == nil {
return errors.New("database not initialized")
return errors.New(errDatabaseNotInitialized)
}
if err := database.Where("key = ?", key).First(sc).Error; err != nil {
@@ -109,7 +109,7 @@ func GetIntByKey(ctx context.Context, key string) (int, error) {
value, err := strconv.Atoi(sc.Value)
if err != nil {
return 0, fmt.Errorf("配置 %s 的值 '%s' 无法转换为整数: %w", key, sc.Value, err)
return 0, fmt.Errorf(errConfigIntParseFailed, key, sc.Value, err)
}
return value, nil
@@ -125,7 +125,7 @@ func GetDecimalByKey(ctx context.Context, key string, precision int32) (decimal.
value, err := decimal.NewFromString(sc.Value)
if err != nil {
return decimal.Zero, fmt.Errorf("配置 %s 的值 '%s' 无法转换为decimal: %w", key, sc.Value, err)
return decimal.Zero, fmt.Errorf(errConfigDecimalParseFailed, key, sc.Value, err)
}
// 裁剪到指定小数位数
@@ -141,7 +141,7 @@ func GetBoolByKey(ctx context.Context, key string) (bool, error) {
value, err := strconv.ParseBool(sc.Value)
if err != nil {
return false, fmt.Errorf("配置 %s 的值 '%s' 无法转换为布尔值: %w", key, sc.Value, err)
return false, fmt.Errorf(errConfigBoolParseFailed, key, sc.Value, err)
}
return value, nil
@@ -160,7 +160,7 @@ func GetMenuDisplayConfig(ctx context.Context) (map[string]bool, error) {
}
if err := json.Unmarshal([]byte(sc.Value), &config); err != nil {
return nil, fmt.Errorf("解析目录显示配置失败: %w", err)
return nil, fmt.Errorf(errParseMenuDisplayConfigFailed, err)
}
return config, nil
+12 -21
View File
@@ -20,6 +20,7 @@ import (
"bytes"
"context"
"errors"
"fmt"
"strings"
"text/template"
"time"
@@ -55,13 +56,13 @@ func (t *Template) Normalize() {
func (t *Template) Validate() error {
t.Normalize()
if t.Key == "" {
return errors.New("模板标识符不能为空")
return errors.New(errTemplateKeyRequired)
}
if t.Name == "" {
return errors.New("模板名称不能为空")
return errors.New(errTemplateNameRequired)
}
if t.Content == "" {
return errors.New("模板内容不能为空")
return errors.New(errTemplateContentRequired)
}
return nil
}
@@ -94,26 +95,16 @@ func (t *Template) Render(data any) (string, string, error) {
return subject, bodyBuf.String(), nil
}
// RenderTemplate 渲染模板的高级包装。如果读取或渲染失败,将使用 fallbackSubject 和 fallbackBody 进行解析和返回。
func RenderTemplate(ctx context.Context, key string, data any, fallbackSubject, fallbackBody string) (string, string) {
// RenderTemplate 渲染指定模板。模板不存在或渲染失败时返回错误,由调用方决定如何处理。
func RenderTemplate(ctx context.Context, key string, data any) (string, string, error) {
var t Template
if err := db.DB(ctx).Where("key = ?", key).First(&t).Error; err == nil {
subject, body, err := t.Render(data)
if err == nil {
return subject, body
}
if err := db.DB(ctx).Where("key = ?", key).First(&t).Error; err != nil {
return "", "", fmt.Errorf(errTemplateUnavailable, key, err)
}
// 降级使用传入的默认模板内容渲染
tFallback := Template{
Key: key + "_fallback",
Subject: fallbackSubject,
Content: fallbackBody,
subject, body, err := t.Render(data)
if err != nil {
return "", "", fmt.Errorf(errTemplateRenderFailed, key, err)
}
subject, body, err := tFallback.Render(data)
if err == nil {
return subject, body
}
return fallbackSubject, fallbackBody
return subject, body, nil
}
+9 -6
View File
@@ -92,12 +92,15 @@ func (u *User) SetEncryptedPassword(password string) error {
return nil
}
func (u *User) IsPasswordEncrypted() bool {
return strings.HasPrefix(u.Password, "$2a$") || strings.HasPrefix(u.Password, "$2b$") || strings.HasPrefix(u.Password, "$2y$")
}
func (u *User) CheckPassword(password string) bool {
if u.Password == "" || password == "" {
return false
}
isBcrypt := strings.HasPrefix(u.Password, "$2a$") || strings.HasPrefix(u.Password, "$2b$") || strings.HasPrefix(u.Password, "$2y$")
if isBcrypt {
if u.IsPasswordEncrypted() {
return util.CheckPasswordHash(u.Password, password)
}
return u.Password == password
@@ -132,7 +135,7 @@ func (u *User) CheckActive() error {
func (u *User) CreateUser(ctx context.Context, tx *gorm.DB, oauthInfo *OAuthUserInfo) error {
enabled, err := GetBoolByKey(ctx, ConfigKeyRegistrationEnabled)
if err == nil && !enabled {
return errors.New("注册已关闭")
return errors.New(errRegistrationDisabled)
}
now := time.Now()
@@ -158,7 +161,7 @@ func (u *User) CreateUser(ctx context.Context, tx *gorm.DB, oauthInfo *OAuthUser
func (u *User) RegisterUser(ctx context.Context, tx *gorm.DB) error {
enabled, err := GetBoolByKey(ctx, ConfigKeyRegistrationEnabled)
if err == nil && !enabled {
return errors.New("注册已关闭")
return errors.New(errRegistrationDisabled)
}
// 检查用户名冲突
@@ -167,7 +170,7 @@ func (u *User) RegisterUser(ctx context.Context, tx *gorm.DB) error {
return err
}
if count > 0 {
return errors.New("用户名已存在")
return errors.New(errUsernameExists)
}
// 检查邮箱冲突
@@ -177,7 +180,7 @@ func (u *User) RegisterUser(ctx context.Context, tx *gorm.DB) error {
return err
}
if emailCount > 0 {
return errors.New("该邮箱已被其他账号绑定")
return errors.New(errEmailAlreadyBound)
}
}
+1 -2
View File
@@ -1,8 +1,7 @@
//go:build embed_frontend
/*
Copyright 2025 linux.do
Modified by Arctel.net, 2026
Copyright 2026 Arctel.net
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
-1
View File
@@ -215,7 +215,6 @@ func Serve() {
{
systemConfigRouter.GET("", system_config.GetSystemConfig)
systemConfigRouter.PUT("", system_config.UpdateSystemConfig)
systemConfigRouter.DELETE("", system_config.DeleteSystemConfig)
}
// Templates
+13 -4
View File
@@ -1,6 +1,5 @@
/*
Copyright 2025 linux.do
Modified by Arctel.net, 2026
Copyright 2026 Arctel.net
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
@@ -20,11 +19,21 @@ package storage
type ErrS3InitializationFailed struct{}
func (e ErrS3InitializationFailed) Error() string {
return "S3存储初始化失败"
return errS3InitializationFailed
}
type LocalCacheError struct{}
func (e LocalCacheError) Error() string {
return "本地缓存错误"
return errLocalCache
}
const (
errS3InitializationFailed = "S3存储初始化失败"
errLocalCache = "本地缓存错误"
errS3PutObjectFailed = "s3 put object failed: %w"
errS3GetObjectFailed = "s3 get object failed: %w"
errCDNRequestFailed = "cdn request failed: %w"
errCDNStatusFailed = "cdn returned status %d"
errS3DeleteObjectFailed = "s3 delete object failed: %w"
)
+5 -5
View File
@@ -146,7 +146,7 @@ func putObjectDefault(ctx context.Context, key string, body io.Reader, size int6
_, err := client.PutObject(ctx, input)
if err != nil {
span.SetStatus(codes.Error, fmt.Sprintf("S3 put object failed: %v", err))
return fmt.Errorf("s3 put object failed: %w", err)
return fmt.Errorf(errS3PutObjectFailed, err)
}
return nil
}
@@ -181,7 +181,7 @@ func getObjectDefault(ctx context.Context, key string) (*ObjectInfo, error) {
})
if err != nil {
span.SetStatus(codes.Error, fmt.Sprintf("S3 get object failed: %v", err))
return nil, fmt.Errorf("s3 get object failed: %w", err)
return nil, fmt.Errorf(errS3GetObjectFailed, err)
}
contentType := "application/octet-stream"
@@ -223,13 +223,13 @@ func GetObjectViaProxy(ctx context.Context, key string) (*ObjectInfo, error) {
resp, err := util.Request(ctx, http.MethodGet, url, nil, nil, nil)
if err != nil {
span.SetStatus(codes.Error, fmt.Sprintf("cdn request failed: %v", err))
return nil, fmt.Errorf("cdn request failed: %w", err)
return nil, fmt.Errorf(errCDNRequestFailed, err)
}
if resp.StatusCode != http.StatusOK {
resp.Body.Close()
span.SetStatus(codes.Error, fmt.Sprintf("cdn returned status %d", resp.StatusCode))
return nil, fmt.Errorf("cdn returned status %d", resp.StatusCode)
return nil, fmt.Errorf(errCDNStatusFailed, resp.StatusCode)
}
contentType := resp.Header.Get("Content-Type")
@@ -265,7 +265,7 @@ func deleteObjectDefault(ctx context.Context, key string) error {
})
if err != nil {
span.SetStatus(codes.Error, fmt.Sprintf("S3 delete object failed: %v", err))
return fmt.Errorf("s3 delete object failed: %w", err)
return fmt.Errorf(errS3DeleteObjectFailed, err)
}
return nil
}
+30
View File
@@ -0,0 +1,30 @@
/*
Copyright 2026 Arctel.net
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package task
const (
errUnknownTaskType = "未知的任务类型: %s"
errCreateTaskExecutionFailed = "创建任务执行记录失败: %w"
errTaskEnqueueFailed = "任务入队失败: %w"
errTaskExecutionNotFound = "任务执行记录不存在: %w"
errRetryOnlyFailedTask = "只有失败的任务才能重试,当前状态: %s"
errTaskNotRetryable = "该任务不支持重试"
errTaskMaxRetryExceeded = "已达到最大重试次数 %d"
errCreateRetryExecutionFailed = "创建重试任务执行记录失败: %w"
errRetryTaskEnqueueFailed = "重试任务入队失败: %w"
errUnregisteredTaskHandler = "未注册的任务处理器: %s"
)
+10 -10
View File
@@ -99,7 +99,7 @@ func AppendLog(ctx context.Context, format string, args ...interface{}) {
func DispatchTask(ctx context.Context, taskType string, payload []byte, triggeredBy string) (string, error) {
meta := GetTaskMeta(taskType)
if meta == nil {
return "", fmt.Errorf("未知的任务类型: %s", taskType)
return "", fmt.Errorf(errUnknownTaskType, taskType)
}
// 生成唯一的 TaskID
@@ -119,7 +119,7 @@ func DispatchTask(ctx context.Context, taskType string, payload []byte, triggere
}
if err := model.CreateTaskExecution(ctx, execution); err != nil {
return "", fmt.Errorf("创建任务执行记录失败: %w", err)
return "", fmt.Errorf(errCreateTaskExecutionFailed, err)
}
// 入队 Asynq
@@ -137,7 +137,7 @@ func DispatchTask(ctx context.Context, taskType string, payload []byte, triggere
execution.StartedAt = &now
execution.FinishedAt = &now
_ = model.UpdateTaskExecution(ctx, execution)
return "", fmt.Errorf("任务入队失败: %w", err)
return "", fmt.Errorf(errTaskEnqueueFailed, err)
}
return taskID, nil
@@ -147,19 +147,19 @@ func DispatchTask(ctx context.Context, taskType string, payload []byte, triggere
func RetryTask(ctx context.Context, id uint64) (string, error) {
execution, err := model.GetTaskExecutionByID(ctx, id)
if err != nil {
return "", fmt.Errorf("任务执行记录不存在: %w", err)
return "", fmt.Errorf(errTaskExecutionNotFound, err)
}
if execution.Status != model.TaskExecutionStatusFailed {
return "", fmt.Errorf("只有失败的任务才能重试,当前状态: %s", execution.Status)
return "", fmt.Errorf(errRetryOnlyFailedTask, execution.Status)
}
if !execution.Retryable {
return "", fmt.Errorf("该任务不支持重试")
return "", fmt.Errorf(errTaskNotRetryable)
}
if execution.RetryCount >= execution.MaxRetry {
return "", fmt.Errorf("已达到最大重试次数 %d", execution.MaxRetry)
return "", fmt.Errorf(errTaskMaxRetryExceeded, execution.MaxRetry)
}
// 生成新的 TaskID
@@ -179,7 +179,7 @@ func RetryTask(ctx context.Context, id uint64) (string, error) {
}
if err := model.CreateTaskExecution(ctx, newExecution); err != nil {
return "", fmt.Errorf("创建重试任务执行记录失败: %w", err)
return "", fmt.Errorf(errCreateRetryExecutionFailed, err)
}
// 入队 Asynq
@@ -196,7 +196,7 @@ func RetryTask(ctx context.Context, id uint64) (string, error) {
newExecution.StartedAt = &now
newExecution.FinishedAt = &now
_ = model.UpdateTaskExecution(ctx, newExecution)
return "", fmt.Errorf("重试任务入队失败: %w", err)
return "", fmt.Errorf(errRetryTaskEnqueueFailed, err)
}
return newTaskID, nil
@@ -224,7 +224,7 @@ func ProcessTask(ctx context.Context, t *asynq.Task) error {
// 查找处理器
handler, ok := getHandler(t.Type())
if !ok {
err := fmt.Errorf("未注册的任务处理器: %s", t.Type())
err := fmt.Errorf(errUnregisteredTaskHandler, t.Type())
logger.ErrorF(ctx, "[TaskExecutor] %v", err)
span.SetStatus(codes.Error, err.Error())
return err
+21
View File
@@ -0,0 +1,21 @@
/*
Copyright 2026 Arctel.net
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package scheduler
const (
errLoadLocationFailed = "failed to load location: %w"
)
+1 -1
View File
@@ -48,7 +48,7 @@ func StartScheduler() error {
schedulerOnce.Do(func() {
location, locErr := time.LoadLocation("Asia/Shanghai")
if locErr != nil {
err = fmt.Errorf("failed to load location: %w", locErr)
err = fmt.Errorf(errLoadLocationFailed, locErr)
return
}
scheduler = asynq.NewScheduler(
+9 -9
View File
@@ -105,10 +105,10 @@ func jwtSign(payload []byte, secret []byte) string {
func jwtVerify(token string, secret []byte) ([]byte, error) {
parts := strings.Split(token, ".")
if len(parts) != 3 {
return nil, errors.New("invalid token format")
return nil, errors.New(errInvalidTokenFormat)
}
if parts[0] != jwtHeaderB64 {
return nil, errors.New("invalid header")
return nil, errors.New(errInvalidHeader)
}
sigInput := parts[0] + "." + parts[1]
@@ -122,7 +122,7 @@ func jwtVerify(token string, secret []byte) ([]byte, error) {
}
if !hmac.Equal(expectedSig, actualSig) {
return nil, errors.New("signature mismatch")
return nil, errors.New(errSignatureMismatch)
}
payload, err := b64urlDecode(parts[1])
@@ -195,25 +195,25 @@ func GenerateChallenge(secret []byte, conf ChallengeConfig, scope string) (*Chal
func VerifyChallengeSolutions(token string, solutions []int, secret []byte, expectedScope string) (*ChallengePayload, error) {
payloadBytes, err := jwtVerify(token, secret)
if err != nil {
return nil, errors.New("invalid_token")
return nil, errors.New(errInvalidToken)
}
var payload ChallengePayload
if err := json.Unmarshal(payloadBytes, &payload); err != nil {
return nil, errors.New("invalid_token")
return nil, errors.New(errInvalidToken)
}
if expectedScope != "" && payload.Scope != expectedScope {
return nil, errors.New("scope_mismatch")
return nil, errors.New(errScopeMismatch)
}
now := time.Now().UnixNano() / int64(time.Millisecond)
if payload.Expires < now {
return nil, errors.New("expired")
return nil, errors.New(errExpired)
}
if len(solutions) != payload.Count {
return nil, errors.New("invalid_solutions")
return nil, errors.New(errInvalidSolutions)
}
tokenFnv := fnv1a(token)
@@ -229,7 +229,7 @@ func VerifyChallengeSolutions(token string, solutions []int, secret []byte, expe
hashHex := hex.EncodeToString(hashBytes[:])
if !strings.HasPrefix(hashHex, target) {
return nil, errors.New("invalid_solution")
return nil, errors.New(errInvalidSolution)
}
}
+28
View File
@@ -0,0 +1,28 @@
/*
Copyright 2026 Arctel.net
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package cap
const (
errInvalidTokenFormat = "invalid token format"
errInvalidHeader = "invalid header"
errSignatureMismatch = "signature mismatch"
errInvalidToken = "invalid_token"
errScopeMismatch = "scope_mismatch"
errExpired = "expired"
errInvalidSolutions = "invalid_solutions"
errInvalidSolution = "invalid_solution"
)
+12 -12
View File
@@ -54,28 +54,28 @@ func encryptBytes(signKey string, plaintext []byte) (string, error) {
// 将 hex 编码的密钥转换为字节
key, err := hex.DecodeString(signKey)
if err != nil {
return "", fmt.Errorf("invalid sign key: %w", err)
return "", fmt.Errorf(errInvalidSignKey, err)
}
if len(key) != 32 {
return "", errors.New("sign key must be 32 bytes (64 hex characters)")
return "", errors.New(errSignKeyLengthInvalid)
}
// 创建 AES cipher
block, err := aes.NewCipher(key)
if err != nil {
return "", fmt.Errorf("failed to create cipher: %w", err)
return "", fmt.Errorf(errCreateCipherFailed, err)
}
// 使用 GCM 模式(Galois/Counter Mode)
gcm, err := cipher.NewGCM(block)
if err != nil {
return "", fmt.Errorf("failed to create GCM: %w", err)
return "", fmt.Errorf(errCreateGCMFailed, err)
}
// 生成随机 nonce
nonce := make([]byte, gcm.NonceSize())
if _, err := io.ReadFull(rand.Reader, nonce); err != nil {
return "", fmt.Errorf("failed to generate nonce: %w", err)
return "", fmt.Errorf(errGenerateNonceFailed, err)
}
// 加密数据
@@ -90,34 +90,34 @@ func decryptBytes(signKey string, ciphertext string) ([]byte, error) {
// 将 hex 编码的密钥转换为字节
key, err := hex.DecodeString(signKey)
if err != nil {
return nil, fmt.Errorf("invalid sign key: %w", err)
return nil, fmt.Errorf(errInvalidSignKey, err)
}
if len(key) != 32 {
return nil, errors.New("sign key must be 32 bytes (64 hex characters)")
return nil, errors.New(errSignKeyLengthInvalid)
}
// 解码 base64 密文
data, err := Base64Decode(ciphertext)
if err != nil {
return nil, fmt.Errorf("failed to decode ciphertext: %w", err)
return nil, fmt.Errorf(errDecodeCiphertextFailed, err)
}
// 创建 AES cipher
block, err := aes.NewCipher(key)
if err != nil {
return nil, fmt.Errorf("failed to create cipher: %w", err)
return nil, fmt.Errorf(errCreateCipherFailed, err)
}
// 使用 GCM 模式
gcm, err := cipher.NewGCM(block)
if err != nil {
return nil, fmt.Errorf("failed to create GCM: %w", err)
return nil, fmt.Errorf(errCreateGCMFailed, err)
}
// 提取 nonce
nonceSize := gcm.NonceSize()
if len(data) < nonceSize {
return nil, errors.New("ciphertext too short")
return nil, errors.New(errCiphertextTooShort)
}
nonce, ciphertextBytes := data[:nonceSize], data[nonceSize:]
@@ -125,7 +125,7 @@ func decryptBytes(signKey string, ciphertext string) ([]byte, error) {
// 解密数据
plaintext, err := gcm.Open(nil, nonce, ciphertextBytes, nil)
if err != nil {
return nil, fmt.Errorf("failed to decrypt: %w", err)
return nil, fmt.Errorf(errDecryptFailed, err)
}
return plaintext, nil
+1 -1
View File
@@ -29,7 +29,7 @@ type StringArray []string
func (sa *StringArray) Scan(value interface{}) error {
bytesValue, ok := value.([]byte)
if !ok {
return fmt.Errorf("invalid value: %v", value)
return fmt.Errorf(errInvalidCustomValue, value)
}
return json.Unmarshal(bytesValue, sa)
}
+31
View File
@@ -0,0 +1,31 @@
/*
Copyright 2026 Arctel.net
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package util
const (
errCreateHTTPRequestFailed = "创建HTTP请求失败: %w"
errHTTPRequestFailed = "请求%s接口失败: %w"
errInvalidCustomValue = "invalid value: %v"
errInvalidSignKey = "invalid sign key: %w"
errSignKeyLengthInvalid = "sign key must be 32 bytes (64 hex characters)"
errCreateCipherFailed = "failed to create cipher: %w"
errCreateGCMFailed = "failed to create GCM: %w"
errGenerateNonceFailed = "failed to generate nonce: %w"
errDecodeCiphertextFailed = "failed to decode ciphertext: %w"
errCiphertextTooShort = "ciphertext too short"
errDecryptFailed = "failed to decrypt: %w"
)
+2 -2
View File
@@ -55,7 +55,7 @@ func SetHTTPClient(c *http.Client) {
func Request(ctx context.Context, method, url string, body io.Reader, headers, cookies map[string]string) (*http.Response, error) {
req, err := http.NewRequestWithContext(ctx, method, url, body)
if err != nil {
return nil, fmt.Errorf("创建HTTP请求失败: %w", err)
return nil, fmt.Errorf(errCreateHTTPRequestFailed, err)
}
if cookies != nil {
@@ -72,7 +72,7 @@ func Request(ctx context.Context, method, url string, body io.Reader, headers, c
resp, err := httpClient.Do(req)
if err != nil {
return nil, fmt.Errorf("请求%s接口失败: %w", url, err)
return nil, fmt.Errorf(errHTTPRequestFailed, url, err)
}
return resp, nil
+28
View File
@@ -0,0 +1,28 @@
/*
Copyright 2026 Arctel.net
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package mail
const (
errDialTLSFailed = "dial tls failed: %w"
errSMTPClientCreationFailed = "smtp client creation failed: %w"
errSMTPAuthFailed = "smtp auth failed: %w"
errSMTPMailCommandFailed = "smtp mail command failed: %w"
errSMTPRcptCommandFailed = "smtp rcpt command failed: %w"
errSMTPDataCommandFailed = "smtp data command failed: %w"
errSMTPWritingBodyFailed = "smtp writing body failed: %w"
errSendMailFailed = "send mail failed: %w"
)
+8 -8
View File
@@ -69,38 +69,38 @@ func SendMailHTML(cfg Config, to string, subject, body string) error {
dialer := &net.Dialer{Timeout: 5 * time.Second}
conn, err := tls.DialWithDialer(dialer, "tcp", addr, tlsConfig)
if err != nil {
return fmt.Errorf("dial tls failed: %w", err)
return fmt.Errorf(errDialTLSFailed, err)
}
defer conn.Close()
_ = conn.SetDeadline(time.Now().Add(10 * time.Second))
client, err := smtp.NewClient(conn, cfg.Host)
if err != nil {
return fmt.Errorf("smtp client creation failed: %w", err)
return fmt.Errorf(errSMTPClientCreationFailed, err)
}
defer client.Close()
if err = client.Auth(auth); err != nil {
return fmt.Errorf("smtp auth failed: %w", err)
return fmt.Errorf(errSMTPAuthFailed, err)
}
if err = client.Mail(cfg.Username); err != nil {
return fmt.Errorf("smtp mail command failed: %w", err)
return fmt.Errorf(errSMTPMailCommandFailed, err)
}
if err = client.Rcpt(to); err != nil {
return fmt.Errorf("smtp rcpt command failed: %w", err)
return fmt.Errorf(errSMTPRcptCommandFailed, err)
}
w, err := client.Data()
if err != nil {
return fmt.Errorf("smtp data command failed: %w", err)
return fmt.Errorf(errSMTPDataCommandFailed, err)
}
defer w.Close()
_, err = w.Write([]byte(message))
if err != nil {
return fmt.Errorf("smtp writing body failed: %w", err)
return fmt.Errorf(errSMTPWritingBodyFailed, err)
}
return nil
@@ -109,7 +109,7 @@ func SendMailHTML(cfg Config, to string, subject, body string) error {
// For standard port (587 / 25), use smtp.SendMail directly (handles STARTTLS automatically if server supports it)
err := smtp.SendMail(addr, auth, cfg.Username, []string{to}, []byte(message))
if err != nil {
return fmt.Errorf("send mail failed: %w", err)
return fmt.Errorf(errSendMailFailed, err)
}
return nil