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)