mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-08 00:26:37 +08:00
重构
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -18,5 +18,8 @@ limitations under the License.
|
||||
package admin
|
||||
|
||||
const (
|
||||
AdminRequired = "未经授权访问"
|
||||
AdminRequired = "未经授权访问"
|
||||
InvalidAuthSourceID = "认证源 ID 无效"
|
||||
InvalidCursorParam = "无效的 cursor 参数"
|
||||
InvalidTaskExecutionID = "无效的任务执行记录 ID"
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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 = "验证码校验失败或已过期,请重试"
|
||||
)
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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 (
|
||||
|
||||
@@ -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 无效"
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
// 没有更多数据,退出循环
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user