diff --git a/common/quota_math.go b/common/quota_math.go new file mode 100644 index 0000000..9d71d8d --- /dev/null +++ b/common/quota_math.go @@ -0,0 +1,21 @@ +package common + +import "math" + +// QuotaFromFloat converts a computed quota value to int with saturation. +// Quota products can include user-controlled multipliers such as image count, +// video seconds, or resolution ratios; oversized products must never wrap into +// a negative charge. The bound is int32 because quota columns are int fields +// used as 32-bit database integers in supported deployments. +func QuotaFromFloat(value float64) int { + if math.IsNaN(value) { + return 0 + } + if value >= math.MaxInt32 { + return math.MaxInt32 + } + if value <= math.MinInt32 { + return math.MinInt32 + } + return int(value) +} diff --git a/common/quota_math_test.go b/common/quota_math_test.go new file mode 100644 index 0000000..11ecdf5 --- /dev/null +++ b/common/quota_math_test.go @@ -0,0 +1,18 @@ +package common + +import ( + "math" + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestQuotaFromFloat(t *testing.T) { + assert.Equal(t, 42, QuotaFromFloat(42.4)) + assert.Equal(t, -42, QuotaFromFloat(-42.4)) + assert.Equal(t, math.MaxInt32, QuotaFromFloat(2000*1.8446744073686647e19)) + assert.Equal(t, math.MinInt32, QuotaFromFloat(-2000*1.8446744073686647e19)) + assert.Equal(t, math.MaxInt32, QuotaFromFloat(math.Inf(1))) + assert.Equal(t, math.MinInt32, QuotaFromFloat(math.Inf(-1))) + assert.Equal(t, 0, QuotaFromFloat(math.NaN())) +} diff --git a/common/ssrf_protection.go b/common/ssrf_protection.go index 8575528..d07fb3b 100644 --- a/common/ssrf_protection.go +++ b/common/ssrf_protection.go @@ -279,6 +279,10 @@ func (p *SSRFProtection) validateHostAndPort(host string, port int) error { return nil } +func (p *SSRFProtection) ValidateNetworkTarget(host string, port int) error { + return p.validateHostAndPort(host, port) +} + func (p *SSRFProtection) validateResolvedIP(host string, ip net.IP) error { if !p.IsIPAccessAllowed(ip) { if isPrivateIP(ip) && !p.AllowPrivateIp { @@ -292,6 +296,10 @@ func (p *SSRFProtection) validateResolvedIP(host string, ip net.IP) error { return nil } +func (p *SSRFProtection) ValidateResolvedIP(host string, ip net.IP) error { + return p.validateResolvedIP(host, ip) +} + func (p *SSRFProtection) resolveValidatedIPs(ctx context.Context, host string) ([]net.IP, error) { ips, err := net.DefaultResolver.LookupIPAddr(ctx, host) if err != nil { @@ -414,6 +422,11 @@ func NewSSRFProtectionWithFetchSetting(enableSSRFProtection, allowPrivateIp bool }, true, nil } +func NewSSRFProtectionFromFetchSetting(allowPrivateIp bool, domainFilterMode bool, ipFilterMode bool, domainList, ipList, allowedPorts []string, applyIPFilterForDomain bool) (*SSRFProtection, error) { + protection, _, err := NewSSRFProtectionWithFetchSetting(true, allowPrivateIp, domainFilterMode, ipFilterMode, domainList, ipList, allowedPorts, applyIPFilterForDomain) + return protection, err +} + // ValidateURLWithFetchSetting 使用FetchSetting配置验证URL func ValidateURLWithFetchSetting(urlStr string, enableSSRFProtection, allowPrivateIp bool, domainFilterMode bool, ipFilterMode bool, domainList, ipList, allowedPorts []string, applyIPFilterForDomain bool) error { protection, enabled, err := NewSSRFProtectionWithFetchSetting(enableSSRFProtection, allowPrivateIp, domainFilterMode, ipFilterMode, domainList, ipList, allowedPorts, applyIPFilterForDomain) diff --git a/controller/discord.go b/controller/discord.go index ffc53e0..e447b3e 100644 --- a/controller/discord.go +++ b/controller/discord.go @@ -211,7 +211,7 @@ func DiscordBind(c *gin.Context) { return } user.DiscordId = discordUser.UID - err = user.Update(false) + err = user.UpdateFields(false, model.UserUpdateFieldDiscordId) if err != nil { common.ApiError(c, err) return diff --git a/controller/github.go b/controller/github.go index a6e40d2..4ef8da4 100644 --- a/controller/github.go +++ b/controller/github.go @@ -207,7 +207,7 @@ func GitHubBind(c *gin.Context) { return } user.GitHubId = githubUser.Login - err = user.Update(false) + err = user.UpdateFields(false, model.UserUpdateFieldGitHubId) if err != nil { common.ApiError(c, err) return diff --git a/controller/linuxdo.go b/controller/linuxdo.go index 1b1a072..56fce9b 100644 --- a/controller/linuxdo.go +++ b/controller/linuxdo.go @@ -66,7 +66,7 @@ func LinuxDoBind(c *gin.Context) { } user.LinuxDOId = strconv.Itoa(linuxdoUser.Id) - err = user.Update(false) + err = user.UpdateFields(false, model.UserUpdateFieldLinuxDOId) if err != nil { common.ApiError(c, err) return diff --git a/controller/log.go b/controller/log.go index 9839641..ff550fd 100644 --- a/controller/log.go +++ b/controller/log.go @@ -25,7 +25,23 @@ func GetAllLogs(c *gin.Context) { group := c.Query("group") requestId := c.Query("request_id") upstreamRequestId := c.Query("upstream_request_id") - logs, total, err := model.GetAllLogs(logType, logFilter, startTimestamp, endTimestamp, modelName, username, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), channel, group, requestId, upstreamRequestId) + quotaFilter := c.Query("quota_filter") + logs, total, err := model.GetAllLogs(model.LogQueryParams{ + LogType: logType, + LogFilter: logFilter, + StartTimestamp: startTimestamp, + EndTimestamp: endTimestamp, + ModelName: modelName, + Username: username, + TokenName: tokenName, + StartIdx: pageInfo.GetStartIdx(), + Num: pageInfo.GetPageSize(), + Channel: channel, + Group: group, + RequestId: requestId, + UpstreamRequestId: upstreamRequestId, + QuotaFilter: quotaFilter, + }) if err != nil { common.ApiError(c, err) return @@ -77,7 +93,22 @@ func GetUserLogs(c *gin.Context) { group := c.Query("group") requestId := c.Query("request_id") upstreamRequestId := c.Query("upstream_request_id") - logs, total, err := model.GetUserLogs(userId, logType, logFilter, startTimestamp, endTimestamp, modelName, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), group, requestId, upstreamRequestId) + quotaFilter := c.Query("quota_filter") + logs, total, err := model.GetUserLogs(model.LogQueryParams{ + UserId: userId, + LogType: logType, + LogFilter: logFilter, + StartTimestamp: startTimestamp, + EndTimestamp: endTimestamp, + ModelName: modelName, + TokenName: tokenName, + StartIdx: pageInfo.GetStartIdx(), + Num: pageInfo.GetPageSize(), + Group: group, + RequestId: requestId, + UpstreamRequestId: upstreamRequestId, + QuotaFilter: quotaFilter, + }) if err != nil { common.ApiError(c, err) return @@ -165,7 +196,19 @@ func GetLogsStat(c *gin.Context) { modelName := c.Query("model_name") channel, _ := strconv.Atoi(c.Query("channel")) group := c.Query("group") - stat, err := model.SumUsedQuota(logType, logFilter, startTimestamp, endTimestamp, modelName, username, tokenName, channel, group) + quotaFilter := c.Query("quota_filter") + stat, err := model.SumUsedQuota(model.LogQueryParams{ + LogType: logType, + LogFilter: logFilter, + StartTimestamp: startTimestamp, + EndTimestamp: endTimestamp, + ModelName: modelName, + Username: username, + TokenName: tokenName, + Channel: channel, + Group: group, + QuotaFilter: quotaFilter, + }) if err != nil { common.ApiError(c, err) return @@ -193,7 +236,19 @@ func GetLogsSelfStat(c *gin.Context) { modelName := c.Query("model_name") channel, _ := strconv.Atoi(c.Query("channel")) group := c.Query("group") - quotaNum, err := model.SumUsedQuota(logType, logFilter, startTimestamp, endTimestamp, modelName, username, tokenName, channel, group) + quotaFilter := c.Query("quota_filter") + quotaNum, err := model.SumUsedQuota(model.LogQueryParams{ + LogType: logType, + LogFilter: logFilter, + StartTimestamp: startTimestamp, + EndTimestamp: endTimestamp, + ModelName: modelName, + Username: username, + TokenName: tokenName, + Channel: channel, + Group: group, + QuotaFilter: quotaFilter, + }) if err != nil { common.ApiError(c, err) return diff --git a/controller/midjourney.go b/controller/midjourney.go index 9229c5b..7bdabfd 100644 --- a/controller/midjourney.go +++ b/controller/midjourney.go @@ -263,6 +263,7 @@ func GetAllMidjourney(c *gin.Context) { MjID: c.Query("mj_id"), StartTimestamp: c.Query("start_timestamp"), EndTimestamp: c.Query("end_timestamp"), + QuotaFilter: c.Query("quota_filter"), } items := model.GetAllTasks(pageInfo.GetStartIdx(), pageInfo.GetPageSize(), queryParams) @@ -288,6 +289,7 @@ func GetUserMidjourney(c *gin.Context) { MjID: c.Query("mj_id"), StartTimestamp: c.Query("start_timestamp"), EndTimestamp: c.Query("end_timestamp"), + QuotaFilter: c.Query("quota_filter"), } items := model.GetAllUserTask(userId, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), queryParams) diff --git a/controller/misc.go b/controller/misc.go index 53d7eee..4adc489 100644 --- a/controller/misc.go +++ b/controller/misc.go @@ -1,13 +1,14 @@ package controller import ( - "encoding/json" + "errors" "fmt" "net/http" "strings" "github.com/MAX-API-Next/MAX-API/common" "github.com/MAX-API-Next/MAX-API/constant" + "github.com/MAX-API-Next/MAX-API/i18n" "github.com/MAX-API-Next/MAX-API/logger" "github.com/MAX-API-Next/MAX-API/middleware" "github.com/MAX-API-Next/MAX-API/model" @@ -238,12 +239,9 @@ func GetHomePageContent(c *gin.Context) { } func SendEmailVerification(c *gin.Context) { - email := c.Query("email") + email := model.NormalizeEmail(c.Query("email")) if err := common.Validate.Var(email, "required,email"); err != nil { - c.JSON(http.StatusOK, gin.H{ - "success": false, - "message": "无效的参数", - }) + common.ApiErrorI18n(c, i18n.MsgInvalidParams) return } parts := strings.Split(email, "@") @@ -284,10 +282,7 @@ func SendEmailVerification(c *gin.Context) { } if model.IsEmailAlreadyTaken(email) { - c.JSON(http.StatusOK, gin.H{ - "success": false, - "message": "邮箱地址已被占用", - }) + common.ApiErrorI18n(c, i18n.MsgUserEmailAlreadyTaken) return } code := common.GenerateVerificationCode(6) @@ -309,15 +304,12 @@ func SendEmailVerification(c *gin.Context) { } func SendPasswordResetEmail(c *gin.Context) { - email := c.Query("email") + email := model.NormalizeEmail(c.Query("email")) if err := common.Validate.Var(email, "required,email"); err != nil { - c.JSON(http.StatusOK, gin.H{ - "success": false, - "message": "无效的参数", - }) + common.ApiErrorI18n(c, i18n.MsgInvalidParams) return } - if model.IsEmailAlreadyTaken(email) { + if _, err := model.GetUniqueUserByEmail(email); err == nil { code := common.GenerateVerificationCode(0) common.RegisterVerificationCodeWithKey(email, code, common.PasswordResetPurpose) link := fmt.Sprintf("%s/user/reset?email=%s&token=%s", system_setting.ServerAddress, email, code) @@ -328,8 +320,10 @@ func SendPasswordResetEmail(c *gin.Context) { "
重置链接 %d 分钟内有效,如果不是本人操作,请忽略。
", common.SystemName, link, link, common.VerificationValidMinutes) err := common.SendEmail(subject, email, content) if err != nil { - logger.LogError(c.Request.Context(), fmt.Sprintf("failed to send password reset email to %s: %s", email, err.Error())) + logger.LogError(c.Request.Context(), fmt.Sprintf("failed to send password reset email to %s: %s", common.MaskEmail(email), err.Error())) } + } else if !errors.Is(err, model.ErrEmailNotFound) { + logger.LogWarn(c.Request.Context(), fmt.Sprintf("skip password reset email for %s: %s", common.MaskEmail(email), err.Error())) } c.JSON(http.StatusOK, gin.H{ "success": true, @@ -344,24 +338,27 @@ type PasswordResetRequest struct { func ResetPassword(c *gin.Context) { var req PasswordResetRequest - err := json.NewDecoder(c.Request.Body).Decode(&req) + err := common.DecodeJson(c.Request.Body, &req) + if err != nil { + common.ApiError(c, err) + return + } + req.Email = model.NormalizeEmail(req.Email) if req.Email == "" || req.Token == "" { - c.JSON(http.StatusOK, gin.H{ - "success": false, - "message": "无效的参数", - }) + common.ApiErrorI18n(c, i18n.MsgInvalidParams) return } if !common.VerifyCodeWithKey(req.Email, req.Token, common.PasswordResetPurpose) { - c.JSON(http.StatusOK, gin.H{ - "success": false, - "message": "重置链接非法或已过期", - }) + common.ApiErrorI18n(c, i18n.MsgUserPasswordResetLinkInvalid) return } password := common.GenerateVerificationCode(12) err = model.ResetUserPasswordByEmail(req.Email, password) if err != nil { + if errors.Is(err, model.ErrEmailNotFound) || errors.Is(err, model.ErrEmailAmbiguous) { + common.ApiErrorI18n(c, i18n.MsgUserPasswordResetLinkInvalid) + return + } common.ApiError(c, err) return } diff --git a/controller/model_list_test.go b/controller/model_list_test.go index d48759f..f4224d8 100644 --- a/controller/model_list_test.go +++ b/controller/model_list_test.go @@ -14,8 +14,11 @@ import ( "github.com/MAX-API-Next/MAX-API/model" "github.com/MAX-API-Next/MAX-API/setting/config" "github.com/MAX-API-Next/MAX-API/setting/operation_setting" + "github.com/gin-contrib/sessions" + "github.com/gin-contrib/sessions/cookie" "github.com/gin-gonic/gin" "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "gorm.io/gorm" ) @@ -240,3 +243,81 @@ func TestListModelsTokenLimitIncludesTieredBillingModel(t *testing.T) { require.NotContains(t, ids, "zz-token-tiered-missing-expr-model") require.NotContains(t, ids, "zz-token-unpriced-model") } + +func TestCheckUpdatePasswordRequiresCurrentPassword(t *testing.T) { + db := setupModelListControllerTestDB(t) + hashedPassword, err := common.Password2Hash("CurrentPassword123") + require.NoError(t, err) + user := &model.User{ + Username: "password-user", + Password: hashedPassword, + Status: common.UserStatusEnabled, + } + require.NoError(t, db.Create(user).Error) + + updatePassword, err := checkUpdatePassword("", "", user.Id) + require.NoError(t, err) + assert.False(t, updatePassword) + + updatePassword, err = checkUpdatePassword("", "NewPassword123", user.Id) + require.Error(t, err) + assert.False(t, updatePassword) + assert.ErrorIs(t, err, errOriginalPasswordFail) + + updatePassword, err = checkUpdatePassword("CurrentPassword123", "NewPassword123", user.Id) + require.NoError(t, err) + assert.True(t, updatePassword) +} + +func TestCheckUpdatePasswordRejectsHistoricalEmptyPassword(t *testing.T) { + db := setupModelListControllerTestDB(t) + user := &model.User{ + Username: "legacy-passwordless-user", + Password: "", + Status: common.UserStatusEnabled, + } + require.NoError(t, db.Create(user).Error) + + updatePassword, err := checkUpdatePassword("", "NewPassword123", user.Id) + require.Error(t, err) + assert.False(t, updatePassword) + assert.ErrorIs(t, err, errUserPasswordUnset) +} + +func TestSetupLoginDoesNotTouchPasswordWhenPasswordFieldOmitted(t *testing.T) { + db := setupModelListControllerTestDB(t) + require.NoError(t, db.AutoMigrate(&model.Log{})) + + hashedPassword, err := common.Password2Hash("CurrentPassword123") + require.NoError(t, err) + user := &model.User{ + Username: "twofa-user", + Password: hashedPassword, + Role: common.RoleCommonUser, + Status: common.UserStatusEnabled, + Group: "default", + } + require.NoError(t, db.Create(user).Error) + + router := gin.New() + store := cookie.NewStore([]byte("test-session-secret")) + router.Use(sessions.Sessions("session", store)) + router.GET("/", func(c *gin.Context) { + setupLogin(&model.User{ + Id: user.Id, + Username: user.Username, + Role: user.Role, + Status: user.Status, + Group: user.Group, + }, c) + }) + + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodGet, "/", nil) + router.ServeHTTP(recorder, request) + + require.Equal(t, http.StatusOK, recorder.Code) + var stored model.User + require.NoError(t, db.First(&stored, user.Id).Error) + assert.Equal(t, hashedPassword, stored.Password) +} diff --git a/controller/oauth.go b/controller/oauth.go index ce725b9..7db4337 100644 --- a/controller/oauth.go +++ b/controller/oauth.go @@ -1,6 +1,7 @@ package controller import ( + "errors" "fmt" "net/http" "strconv" @@ -106,11 +107,17 @@ func HandleOAuth(c *gin.Context) { // 7. Find or create user user, err := findOrCreateOAuthUser(c, provider, oauthUser, session) if err != nil { + if errors.Is(err, model.ErrEmailAlreadyTaken) { + common.ApiErrorI18n(c, i18n.MsgUserEmailAlreadyTaken) + return + } switch err.(type) { case *OAuthUserDeletedError: common.ApiErrorI18n(c, i18n.MsgOAuthUserDeleted) case *OAuthRegistrationDisabledError: common.ApiErrorI18n(c, i18n.MsgUserRegisterDisabled) + case *OAuthEmailAlreadyTakenError: + common.ApiErrorI18n(c, i18n.MsgUserEmailAlreadyTaken) default: common.ApiError(c, err) } @@ -257,7 +264,13 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o user.DisplayName = provider.GetName() + " User" } if oauthUser.Email != "" { - user.Email = oauthUser.Email + user.Email = model.NormalizeEmail(oauthUser.Email) + if err := model.EnsureEmailAvailable(user.Email, 0); err != nil { + if errors.Is(err, model.ErrEmailAlreadyTaken) { + return nil, &OAuthEmailAlreadyTakenError{} + } + return nil, err + } } user.Role = common.RoleCommonUser user.Status = common.UserStatusEnabled @@ -272,7 +285,7 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o // Use transaction to ensure user creation and OAuth binding are atomic if genericProvider, ok := provider.(*oauth.GenericOAuthProvider); ok { // Custom provider: create user and binding in a transaction - err := model.DB.Transaction(func(tx *gorm.DB) error { + err := model.WithNormalizedEmailWriteTx(user.Email, func(tx *gorm.DB) error { // Create user if err := user.InsertWithTx(tx, inviterId); err != nil { return err @@ -298,7 +311,7 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o user.FinalizeOAuthUserCreation(inviterId) } else { // Built-in provider: create user and update provider ID in a transaction - err := model.DB.Transaction(func(tx *gorm.DB) error { + err := model.WithNormalizedEmailWriteTx(user.Email, func(tx *gorm.DB) error { // Create user if err := user.InsertWithTx(tx, inviterId); err != nil { return err @@ -343,6 +356,12 @@ func (e *OAuthRegistrationDisabledError) Error() string { return "registration is disabled" } +type OAuthEmailAlreadyTakenError struct{} + +func (e *OAuthEmailAlreadyTakenError) Error() string { + return "email is already in use" +} + // handleOAuthError handles OAuth errors and returns translated message func handleOAuthError(c *gin.Context, err error) { switch e := err.(type) { diff --git a/controller/oidc.go b/controller/oidc.go index 1d78cd1..25f82b4 100644 --- a/controller/oidc.go +++ b/controller/oidc.go @@ -215,7 +215,7 @@ func OidcBind(c *gin.Context) { return } user.OidcId = oidcUser.OpenID - err = user.Update(false) + err = user.UpdateFields(false, model.UserUpdateFieldOidcId) if err != nil { common.ApiError(c, err) return diff --git a/controller/task.go b/controller/task.go index 76ec353..adb47fc 100644 --- a/controller/task.go +++ b/controller/task.go @@ -33,6 +33,7 @@ func GetAllTask(c *gin.Context) { StartTimestamp: startTimestamp, EndTimestamp: endTimestamp, ChannelID: c.Query("channel_id"), + QuotaFilter: c.Query("quota_filter"), } items := model.TaskGetAllTasks(pageInfo.GetStartIdx(), pageInfo.GetPageSize(), queryParams) @@ -57,6 +58,7 @@ func GetUserTask(c *gin.Context) { Action: c.Query("action"), StartTimestamp: startTimestamp, EndTimestamp: endTimestamp, + QuotaFilter: c.Query("quota_filter"), } items := model.TaskGetAllUserTask(userId, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), queryParams) diff --git a/controller/task_video.go b/controller/task_video.go index c5d979d..de84b49 100644 --- a/controller/task_video.go +++ b/controller/task_video.go @@ -189,7 +189,7 @@ func updateVideoSingleTask(ctx context.Context, adaptor channel.TaskAdaptor, cha } // 计算实际应扣费额度: totalTokens * modelRatio * groupRatio - actualQuota := int(float64(taskResult.TotalTokens) * modelRatio * finalGroupRatio) + actualQuota := common.QuotaFromFloat(float64(taskResult.TotalTokens) * modelRatio * finalGroupRatio) // 计算差额 preConsumedQuota := task.Quota diff --git a/controller/telegram.go b/controller/telegram.go index 990b2b7..07a0b98 100644 --- a/controller/telegram.go +++ b/controller/telegram.go @@ -58,7 +58,7 @@ func TelegramBind(c *gin.Context) { return } user.TelegramId = telegramId - if err := user.Update(false); err != nil { + if err := user.UpdateFields(false, model.UserUpdateFieldTelegramId); err != nil { c.JSON(200, gin.H{ "message": err.Error(), "success": false, diff --git a/controller/token_test.go b/controller/token_test.go index 13dedbd..a78c8f3 100644 --- a/controller/token_test.go +++ b/controller/token_test.go @@ -552,7 +552,7 @@ func TestAddTokenRejectsNonSelectableAutoRoute(t *testing.T) { if !strings.Contains(response.Message, "auto:internal") { t.Fatalf("expected error message to mention hidden route, got %q", response.Message) } - if !strings.Contains(response.Message, "无权访问") { + if !strings.Contains(response.Message, "无权分配") { t.Fatalf("expected localized denial message, got %q", response.Message) } @@ -633,7 +633,7 @@ func TestUpdateTokenRejectsNonSelectableAutoRoute(t *testing.T) { if !strings.Contains(response.Message, "auto:internal") { t.Fatalf("expected error message to mention hidden route, got %q", response.Message) } - if !strings.Contains(response.Message, "无权访问") { + if !strings.Contains(response.Message, "无权分配") { t.Fatalf("expected localized denial message, got %q", response.Message) } diff --git a/controller/user.go b/controller/user.go index a996d7f..2388f1a 100644 --- a/controller/user.go +++ b/controller/user.go @@ -29,6 +29,11 @@ type LoginRequest struct { Password string `json:"password"` } +var ( + errUserPasswordUnset = errors.New("user password is not set") + errOriginalPasswordFail = errors.New("original password is incorrect") +) + func Login(c *gin.Context) { if !common.PasswordLoginEnabled { common.ApiErrorI18n(c, i18n.MsgUserPasswordLoginDisabled) @@ -184,6 +189,12 @@ func Register(c *gin.Context) { common.ApiErrorI18n(c, i18n.MsgInvalidParams) return } + user.Username = strings.TrimSpace(user.Username) + user.Email = model.NormalizeEmail(user.Email) + if user.Username == "" { + common.ApiErrorI18n(c, i18n.MsgInvalidParams) + return + } if err := common.Validate.Struct(&user); err != nil { common.ApiErrorI18n(c, i18n.MsgUserInputInvalid, map[string]any{"Error": err.Error()}) return @@ -197,8 +208,20 @@ func Register(c *gin.Context) { common.ApiErrorI18n(c, i18n.MsgUserVerificationCodeError) return } + if err := model.EnsureEmailAvailable(user.Email, 0); err != nil { + if errors.Is(err, model.ErrEmailAlreadyTaken) { + common.ApiErrorI18n(c, i18n.MsgUserEmailAlreadyTaken) + return + } + common.ApiErrorI18n(c, i18n.MsgDatabaseError) + return + } + } + emailForExistCheck := "" + if common.EmailVerificationEnabled { + emailForExistCheck = user.Email } - exist, err := model.CheckUserExistOrDeleted(user.Username, user.Email) + exist, err := model.CheckUserExistOrDeleted(user.Username, emailForExistCheck) if err != nil { common.ApiErrorI18n(c, i18n.MsgDatabaseError) common.SysLog(fmt.Sprintf("CheckUserExistOrDeleted error: %v", err)) @@ -221,6 +244,10 @@ func Register(c *gin.Context) { cleanUser.Email = user.Email } if err := cleanUser.Insert(inviterId); err != nil { + if errors.Is(err, model.ErrEmailAlreadyTaken) { + common.ApiErrorI18n(c, i18n.MsgUserEmailAlreadyTaken) + return + } common.ApiError(c, err) return } @@ -361,7 +388,7 @@ func GenerateAccessToken(c *gin.Context) { return } - if err := user.Update(false); err != nil { + if err := user.UpdateFields(false, model.UserUpdateFieldAccessToken); err != nil { common.ApiError(c, err) return } @@ -411,7 +438,7 @@ func GetAffCode(c *gin.Context) { } if user.AffCode == "" { user.AffCode = common.GetRandomString(4) - if err := user.Update(false); err != nil { + if err := user.UpdateFields(false, model.UserUpdateFieldAffCode); err != nil { c.JSON(http.StatusOK, gin.H{ "success": false, "message": err.Error(), @@ -605,6 +632,11 @@ func UpdateUser(c *gin.Context) { common.ApiErrorI18n(c, i18n.MsgInvalidParams) return } + updatedUser.Username = strings.TrimSpace(updatedUser.Username) + if updatedUser.Username == "" { + common.ApiErrorI18n(c, i18n.MsgInvalidParams) + return + } if updatedUser.Password == "" { updatedUser.Password = "$I_LOVE_U" // make Validator happy :) } @@ -767,20 +799,35 @@ func UpdateSelf(c *gin.Context) { cleanUser := model.User{ Id: c.GetInt("id"), - Username: user.Username, Password: user.Password, DisplayName: user.DisplayName, } + updateFields := make([]model.UserUpdateField, 0, 2) + if _, ok := requestData["display_name"]; ok { + updateFields = append(updateFields, model.UserUpdateFieldDisplayName) + } + if _, ok := requestData["username"]; ok && user.Username != "" { + cleanUser.Username = user.Username + updateFields = append(updateFields, model.UserUpdateFieldUsername) + } if user.Password == "$I_LOVE_U" { user.Password = "" // rollback to what it should be cleanUser.Password = "" } updatePassword, err := checkUpdatePassword(user.OriginalPassword, user.Password, cleanUser.Id) if err != nil { + if errors.Is(err, errUserPasswordUnset) { + common.ApiErrorI18n(c, i18n.MsgUserPasswordUnset) + return + } + if errors.Is(err, errOriginalPasswordFail) { + common.ApiErrorI18n(c, i18n.MsgUserOriginalPasswordError) + return + } common.ApiError(c, err) return } - if err := cleanUser.Update(updatePassword); err != nil { + if err := cleanUser.UpdateFields(updatePassword, updateFields...); err != nil { common.ApiError(c, err) return } @@ -793,19 +840,21 @@ func UpdateSelf(c *gin.Context) { } func checkUpdatePassword(originalPassword string, newPassword string, userId int) (updatePassword bool, err error) { + if newPassword == "" { + return + } var currentUser *model.User currentUser, err = model.GetUserById(userId, true) if err != nil { return } - // 密码不为空,需要验证原密码 - // 支持第一次账号绑定时原密码为空的情况 - if !common.ValidatePasswordAndHash(originalPassword, currentUser.Password) && currentUser.Password != "" { - err = fmt.Errorf("原密码错误") + if currentUser.Password == "" { + err = errUserPasswordUnset return } - if newPassword == "" { + if !common.ValidatePasswordAndHash(originalPassword, currentUser.Password) { + err = errOriginalPasswordFail return } updatePassword = true @@ -1040,7 +1089,14 @@ func ManageUser(c *gin.Context) { return } - if err := user.Update(false); err != nil { + var updateFields []model.UserUpdateField + switch req.Action { + case "disable", "enable": + updateFields = []model.UserUpdateField{model.UserUpdateFieldStatus} + case "promote", "demote": + updateFields = []model.UserUpdateField{model.UserUpdateFieldRole} + } + if err := user.UpdateFields(false, updateFields...); err != nil { common.ApiError(c, err) return } @@ -1084,7 +1140,7 @@ func EmailBind(c *gin.Context) { common.ApiError(c, errors.New("invalid request body")) return } - email := req.Email + email := model.NormalizeEmail(req.Email) code := req.Code if !common.VerifyCodeWithKey(email, code, common.EmailVerificationPurpose) { common.ApiErrorI18n(c, i18n.MsgUserVerificationCodeError) @@ -1100,10 +1156,11 @@ func EmailBind(c *gin.Context) { common.ApiError(c, err) return } - user.Email = email - // no need to check if this email already taken, because we have used verification code to check it - err = user.Update(false) - if err != nil { + if err := model.BindEmailToUser(&user, email); err != nil { + if errors.Is(err, model.ErrEmailAlreadyTaken) { + common.ApiErrorI18n(c, i18n.MsgUserEmailAlreadyTaken) + return + } common.ApiError(c, err) return } diff --git a/controller/wechat.go b/controller/wechat.go index aca0071..2c73114 100644 --- a/controller/wechat.go +++ b/controller/wechat.go @@ -169,7 +169,7 @@ func WeChatBind(c *gin.Context) { return } user.WeChatId = wechatId - err = user.Update(false) + err = user.UpdateFields(false, model.UserUpdateFieldWeChatId) if err != nil { common.ApiError(c, err) return diff --git a/dto/openai_image.go b/dto/openai_image.go index 1991075..0ff5b2b 100644 --- a/dto/openai_image.go +++ b/dto/openai_image.go @@ -2,6 +2,7 @@ package dto import ( "encoding/json" + "fmt" "reflect" "strings" @@ -11,6 +12,19 @@ import ( "github.com/gin-gonic/gin" ) +// MaxImageN caps image generation count before it becomes a billing multiplier. +const MaxImageN = 128 + +func ValidateImageN(field string, n int) error { + if field == "" { + field = "n" + } + if n < 0 || n > MaxImageN { + return fmt.Errorf("%s must be an integer between 0 and %d", field, MaxImageN) + } + return nil +} + type ImageRequest struct { Model string `json:"model"` Prompt string `json:"prompt" binding:"required"` diff --git a/dto/openai_image_test.go b/dto/openai_image_test.go new file mode 100644 index 0000000..663d409 --- /dev/null +++ b/dto/openai_image_test.go @@ -0,0 +1,21 @@ +package dto + +import ( + "fmt" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestValidateImageN(t *testing.T) { + require.NoError(t, ValidateImageN("n", 0)) + require.NoError(t, ValidateImageN("n", MaxImageN)) + + err := ValidateImageN("", -1) + require.Error(t, err) + require.Equal(t, fmt.Sprintf("n must be an integer between 0 and %d", MaxImageN), err.Error()) + + err = ValidateImageN("parameters.n", MaxImageN+1) + require.Error(t, err) + require.Equal(t, fmt.Sprintf("parameters.n must be an integer between 0 and %d", MaxImageN), err.Error()) +} diff --git a/go.mod b/go.mod index e59b5cf..f33a45d 100644 --- a/go.mod +++ b/go.mod @@ -4,8 +4,8 @@ module github.com/MAX-API-Next/MAX-API go 1.25.1 require ( - github.com/Calcium-Ion/go-epay v0.0.4 github.com/Azure/go-ntlmssp v0.1.1 + github.com/Calcium-Ion/go-epay v0.0.4 github.com/abema/go-mp4 v1.4.1 github.com/andybalholm/brotli v1.1.1 github.com/anknown/ahocorasick v0.0.0-20190904063843-d75dbd5169c0 @@ -49,12 +49,12 @@ require ( github.com/tiktoken-go/tokenizer v0.6.2 github.com/waffo-com/waffo-go v1.3.1 github.com/yapingcat/gomedia v0.0.0-20240906162731-17feea57090c - golang.org/x/crypto v0.45.0 - golang.org/x/image v0.38.0 - golang.org/x/net v0.47.0 - golang.org/x/sync v0.20.0 - golang.org/x/sys v0.38.0 - golang.org/x/text v0.35.0 + golang.org/x/crypto v0.52.0 + golang.org/x/image v0.43.0 + golang.org/x/net v0.55.0 + golang.org/x/sync v0.21.0 + golang.org/x/sys v0.45.0 + golang.org/x/text v0.38.0 gopkg.in/yaml.v3 v3.0.1 gorm.io/driver/mysql v1.4.3 gorm.io/driver/postgres v1.5.2 diff --git a/go.sum b/go.sum index 592ed58..ff5e8e6 100644 --- a/go.sum +++ b/go.sum @@ -310,10 +310,6 @@ github.com/ugorji/go/codec v1.2.12 h1:9LC83zGrHhuUA9l16C9AHXAqEV/2wBQ4nkvumAE65E github.com/ugorji/go/codec v1.2.12/go.mod h1:UNopzCgEMSXjBc6AOMqYvWC1ktqTAfzJZUZgYf6w6lg= github.com/waffo-com/waffo-go v1.3.1 h1:NCYD3oQ59DTJj1bwS5T/659LI4h8PuAIW4Qj/w7fKPw= github.com/waffo-com/waffo-go v1.3.1/go.mod h1:IaXVYq6mmYtrLFFsLxPslNwuIZx0mIadWWjhe+eWb0g= -github.com/waffo-com/waffo-pancake-sdk-go v0.1.1 h1:YOI7+3zTBlTB7Ou6+ZXnJV2JvW/ag9d7CwE/TxH3Hls= -github.com/waffo-com/waffo-pancake-sdk-go v0.1.1/go.mod h1:5MBCGH/nqRRA5sHO/lQB/96r4BTAqy8QpWxn53m9htI= -github.com/waffo-com/waffo-pancake-sdk-go v0.2.0 h1:cCSgccM66p7feTtgRqUUGT50tYQOhahsoPXavd+ib1U= -github.com/waffo-com/waffo-pancake-sdk-go v0.2.0/go.mod h1:5MBCGH/nqRRA5sHO/lQB/96r4BTAqy8QpWxn53m9htI= github.com/waffo-com/waffo-pancake-sdk-go v0.3.1 h1:ngQSN/oVB35xTwFPLfg++bxPC+SptcF145Mb6c62YCc= github.com/waffo-com/waffo-pancake-sdk-go v0.3.1/go.mod h1:OB2MyFIQaefoPO0FV3J+yu9sDP8RVFQ+sbFsXqGuObc= github.com/x448/float16 v0.8.4 h1:qLwI1I70+NjRFUR3zs1JPUCgaCXSh3SW62uAKT1mSBM= @@ -331,20 +327,20 @@ go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= golang.org/x/arch v0.21.0 h1:iTC9o7+wP6cPWpDWkivCvQFGAHDQ59SrSxsLPcnkArw= golang.org/x/arch v0.21.0/go.mod h1:dNHoOeKiyja7GTvF9NJS1l3Z2yntpQNzgrjh1cU103A= golang.org/x/crypto v0.0.0-20210711020723-a769d52b0f97/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc= -golang.org/x/crypto v0.45.0 h1:jMBrvKuj23MTlT0bQEOBcAE0mjg8mK9RXFhRH6nyF3Q= -golang.org/x/crypto v0.45.0/go.mod h1:XTGrrkGJve7CYK7J8PEww4aY7gM3qMCElcJQ8n8JdX4= +golang.org/x/crypto v0.52.0 h1:RMs7fP2rXdep0CftQlK8Uf+kibLm7qkCcradZWYz988= +golang.org/x/crypto v0.52.0/go.mod h1:1QgfPxDqh0T2M/elOJtp9RvuR95kVjir0e6/BvEmGbc= golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b h1:M2rDM6z3Fhozi9O7NWsxAkg/yqS/lQJ6PmkyIV3YP+o= golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b/go.mod h1:3//PLf8L/X+8b4vuAfHzxeRUl04Adcb341+IGKfnqS8= -golang.org/x/image v0.38.0 h1:5l+q+Y9JDC7mBOMjo4/aPhMDcxEptsX+Tt3GgRQRPuE= -golang.org/x/image v0.38.0/go.mod h1:/3f6vaXC+6CEanU4KJxbcUZyEePbyKbaLoDOe4ehFYY= -golang.org/x/mod v0.33.0 h1:tHFzIWbBifEmbwtGz65eaWyGiGZatSrT9prnU8DbVL8= -golang.org/x/mod v0.33.0/go.mod h1:swjeQEj+6r7fODbD2cqrnje9PnziFuw4bmLbBZFrQ5w= +golang.org/x/image v0.43.0 h1:FLxcP4ec2350nTfOC8ysKtqYSIFbk/QGjw1ZHNP4tsY= +golang.org/x/image v0.43.0/go.mod h1:rrpelvGFt+kLPAjPM4HeWPgrl0FtafueU//e5N0qk/Q= +golang.org/x/mod v0.36.0 h1:JJjpVx6myfUsUdAzZuOSTTmRE0PfZeNWzzvKrP7amb4= +golang.org/x/mod v0.36.0/go.mod h1:moc6ELqsWcOw5Ef3xVprK5ul/MvtVvkIXLziUOICjUQ= golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= golang.org/x/net v0.0.0-20210520170846-37e1c6afe023/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y= -golang.org/x/net v0.47.0 h1:Mx+4dIFzqraBXUugkia1OOvlD6LemFo1ALMHjrXDOhY= -golang.org/x/net v0.47.0/go.mod h1:/jNxtkgq5yWUGYkaZGqo27cfGZ1c5Nen03aYrrKpVRU= -golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4= -golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/net v0.55.0 h1:bcvxaJn3e1U6InsFWt1JUq1aSjnRxLzT2rtD2KfkDF8= +golang.org/x/net v0.55.0/go.mod h1:L5U2KuzuOe1lY7Z+aWVIKK6qEeJXnXV9yzGA+WCHJww= +golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM= +golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sys v0.0.0-20190726091711-fc99dfbffb4e/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20200116001909-b77594299b42/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= @@ -356,18 +352,18 @@ golang.org/x/sys v0.0.0-20210806184541-e5e7981a1069/go.mod h1:oPkhp1MJrh7nUepCBc golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.8.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.11.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.38.0 h1:3yZWxaJjBmCWXqhN1qh02AkOnCQ1poK6oF+a7xWL6Gc= -golang.org/x/sys v0.38.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= +golang.org/x/sys v0.45.0 h1:dO4czNzziLiiXplLQgBCEpCvXQ3dnkn0SdaZSYdQ+FY= +golang.org/x/sys v0.45.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8= golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= -golang.org/x/text v0.35.0 h1:JOVx6vVDFokkpaq1AEptVzLTpDe9KGpj5tR4/X+ybL8= -golang.org/x/text v0.35.0/go.mod h1:khi/HExzZJ2pGnjenulevKNX1W67CUy0AsXcNubPGCA= +golang.org/x/text v0.38.0 h1:sXmwo9DwP3OK9EZ7PqAdaooSGozfl/3a6/xJcbzPRhE= +golang.org/x/text v0.38.0/go.mod h1:YXZt3QhHUKYT53r2lLKFIVi6Ao1jdzrTR/KQ09qyxF4= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= -golang.org/x/tools v0.42.0 h1:uNgphsn75Tdz5Ji2q36v/nsFSfR/9BRFvqhGBaJGd5k= -golang.org/x/tools v0.42.0/go.mod h1:Ma6lCIwGZvHK6XtgbswSoWroEkhugApmsXyrUmBhfr0= +golang.org/x/tools v0.45.0 h1:18qN3FAooORvApf5XjCXgsuayZOEtXf6JK18I3+ONa8= +golang.org/x/tools v0.45.0/go.mod h1:LuUGqqaXcXMEFEruIVJVm5mgDD8vww/z/SR1gQ4uE/0= golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= google.golang.org/protobuf v1.26.0-rc.1/go.mod h1:jlhhOSvTdKEhbULTjvd4ARK9grFBp09yW+WbY/TyQbw= google.golang.org/protobuf v1.28.0/go.mod h1:HV8QOd/L58Z+nl8r43ehVNZIU/HEI6OcFqwMG9pJV4I= diff --git a/i18n/keys.go b/i18n/keys.go index 5d4f4d1..5c1d8f2 100644 --- a/i18n/keys.go +++ b/i18n/keys.go @@ -88,6 +88,9 @@ const ( MsgUserRequire2FA = "user.require_2fa" MsgUserEmailVerificationRequired = "user.email_verification_required" MsgUserVerificationCodeError = "user.verification_code_error" + MsgUserEmailAlreadyTaken = "user.email_already_taken" + MsgUserPasswordUnset = "user.password_unset" + MsgUserPasswordResetLinkInvalid = "user.password_reset_link_invalid" MsgUserInputInvalid = "user.input_invalid" MsgUserNoPermissionSameLevel = "user.no_permission_same_level" MsgUserNoPermissionHigherLevel = "user.no_permission_higher_level" diff --git a/i18n/locales/en.yaml b/i18n/locales/en.yaml index ece1229..5380e59 100644 --- a/i18n/locales/en.yaml +++ b/i18n/locales/en.yaml @@ -76,6 +76,9 @@ user.session_save_failed: "Failed to save session, please try again" user.require_2fa: "Please enter two-factor authentication code" user.email_verification_required: "Email verification is enabled, please enter email address and verification code" user.verification_code_error: "Verification code is incorrect or has expired" +user.email_already_taken: "Email address is already in use" +user.password_unset: "This account has no password set. Please use password reset or contact an administrator to reset it." +user.password_reset_link_invalid: "Password reset link is invalid or has expired" user.input_invalid: "Invalid input {{.Error}}" user.no_permission_same_level: "No permission to access users of same or higher level" user.no_permission_higher_level: "No permission to update users of same or higher permission level" diff --git a/i18n/locales/zh-CN.yaml b/i18n/locales/zh-CN.yaml index 1b7a27e..7798a64 100644 --- a/i18n/locales/zh-CN.yaml +++ b/i18n/locales/zh-CN.yaml @@ -40,7 +40,7 @@ token.quota_negative: "额度值不能为负数" token.quota_exceed_max: "额度值超出有效范围,最大值为 {{.Max}}" token.generate_failed: "生成令牌失败" token.get_info_failed: "获取令牌信息失败,请稍后重试" -token.group_not_assignable: "无权访问 {{.Group}} 分组" +token.group_not_assignable: "无权分配 {{.Group}} 分组" token.expired_cannot_enable: "令牌已过期,无法启用,请先修改令牌过期时间,或者设置为永不过期" token.exhausted_cannot_enable: "令牌可用额度已用尽,无法启用,请先修改令牌剩余额度,或者设置为无限额度" token.invalid: "无效的令牌" @@ -77,6 +77,9 @@ user.session_save_failed: "无法保存会话信息,请重试" user.require_2fa: "请输入两步验证码" user.email_verification_required: "管理员开启了邮箱验证,请输入邮箱地址和验证码" user.verification_code_error: "验证码错误或已过期" +user.email_already_taken: "邮箱地址已被占用" +user.password_unset: "当前账号未设置密码,请使用密码重置或联系管理员重置密码" +user.password_reset_link_invalid: "重置链接非法或已过期" user.input_invalid: "输入不合法 {{.Error}}" user.no_permission_same_level: "无权获取同级或更高等级用户的信息" user.no_permission_higher_level: "无权更新同权限等级或更高权限等级的用户信息" diff --git a/i18n/locales/zh-TW.yaml b/i18n/locales/zh-TW.yaml index 466b731..3d3aed0 100644 --- a/i18n/locales/zh-TW.yaml +++ b/i18n/locales/zh-TW.yaml @@ -77,6 +77,9 @@ user.session_save_failed: "無法保存對話,請重試" user.require_2fa: "請輸入雙重驗證碼" user.email_verification_required: "管理員開啟了信箱驗證,請輸入信箱位址和驗證碼" user.verification_code_error: "驗證碼錯誤或已過期" +user.email_already_taken: "信箱位址已被占用" +user.password_unset: "目前帳號未設定密碼,請使用密碼重置或聯繫管理員重置密碼" +user.password_reset_link_invalid: "重置連結非法或已過期" user.input_invalid: "輸入不合法 {{.Error}}" user.no_permission_same_level: "無權獲取同級或更高等級使用者的資訊" user.no_permission_higher_level: "無權更新同權限等級或更高權限等級的使用者資訊" diff --git a/model/errors.go b/model/errors.go index a942a5b..7f53a03 100644 --- a/model/errors.go +++ b/model/errors.go @@ -11,6 +11,9 @@ var ( var ( ErrInvalidCredentials = errors.New("invalid credentials") ErrUserEmptyCredentials = errors.New("empty credentials") + ErrEmailAlreadyTaken = errors.New("email already taken") + ErrEmailNotFound = errors.New("email not found") + ErrEmailAmbiguous = errors.New("email matches multiple users") ) // Token auth errors diff --git a/model/log.go b/model/log.go index 8908317..87af760 100644 --- a/model/log.go +++ b/model/log.go @@ -33,13 +33,13 @@ func applyExplicitLogTextFilter(tx *gorm.DB, column string, value string) (*gorm type Log struct { Id int `json:"id" gorm:"index:idx_created_at_id,priority:2;index:idx_user_id_id,priority:2"` UserId int `json:"user_id" gorm:"index;index:idx_user_id_id,priority:1"` - CreatedAt int64 `json:"created_at" gorm:"bigint;index:idx_created_at_id,priority:1;index:idx_created_at_type"` - Type int `json:"type" gorm:"index:idx_created_at_type"` + CreatedAt int64 `json:"created_at" gorm:"bigint;index:idx_created_at_id,priority:1;index:idx_created_at_type;index:idx_logs_type_quota_created_at,priority:3"` + Type int `json:"type" gorm:"index:idx_created_at_type;index:idx_logs_type_quota_created_at,priority:1"` Content string `json:"content"` Username string `json:"username" gorm:"index;index:index_username_model_name,priority:2;default:''"` TokenName string `json:"token_name" gorm:"index;default:''"` ModelName string `json:"model_name" gorm:"index;index:index_username_model_name,priority:1;default:''"` - Quota int `json:"quota" gorm:"default:0"` + Quota int `json:"quota" gorm:"default:0;index:idx_logs_quota;index:idx_logs_type_quota_created_at,priority:2"` PromptTokens int `json:"prompt_tokens" gorm:"default:0"` CompletionTokens int `json:"completion_tokens" gorm:"default:0"` UseTime int `json:"use_time" gorm:"default:0"` @@ -91,6 +91,56 @@ const ( LogFilterEmptyRetry = "empty_retry" ) +const ( + LogQuotaFilterAbnormal = "abnormal" + LogQuotaFilterZero = "zero" + LogQuotaFilterNegative = "negative" +) + +type LogQueryParams struct { + UserId int + LogType int + LogFilter string + StartTimestamp int64 + EndTimestamp int64 + ModelName string + Username string + TokenName string + StartIdx int + Num int + Channel int + Group string + RequestId string + UpstreamRequestId string + QuotaFilter string +} + +func normalizeLogQuotaFilter(filter string) string { + switch strings.ToLower(strings.TrimSpace(filter)) { + case LogQuotaFilterAbnormal: + return LogQuotaFilterAbnormal + case LogQuotaFilterZero: + return LogQuotaFilterZero + case LogQuotaFilterNegative: + return LogQuotaFilterNegative + default: + return "" + } +} + +func applyQuotaFilter(tx *gorm.DB, column string, filter string) *gorm.DB { + switch normalizeLogQuotaFilter(filter) { + case LogQuotaFilterAbnormal: + return tx.Where(column+" <= ?", 0) + case LogQuotaFilterZero: + return tx.Where(column+" = ?", 0) + case LogQuotaFilterNegative: + return tx.Where(column+" < ?", 0) + default: + return tx + } +} + func applyLogTypeFilter(tx *gorm.DB, logType int) *gorm.DB { if logType == LogTypeUnknown { return tx @@ -604,44 +654,45 @@ func RecordTaskBillingLog(params RecordTaskBillingLogParams) { } } -func GetAllLogs(logType int, logFilter string, startTimestamp int64, endTimestamp int64, modelName string, username string, tokenName string, startIdx int, num int, channel int, group string, requestId string, upstreamRequestId string) (logs []*Log, total int64, err error) { - tx, err := applyLogFilter(applyLogTypeFilter(LOG_DB, logType), logFilter) +func GetAllLogs(params LogQueryParams) (logs []*Log, total int64, err error) { + tx, err := applyLogFilter(applyLogTypeFilter(LOG_DB, params.LogType), params.LogFilter) if err != nil { return nil, 0, err } + tx = applyQuotaFilter(tx, "logs.quota", params.QuotaFilter) - if tx, err = applyExplicitLogTextFilter(tx, "logs.model_name", modelName); err != nil { + if tx, err = applyExplicitLogTextFilter(tx, "logs.model_name", params.ModelName); err != nil { return nil, 0, err } - if tx, err = applyExplicitLogTextFilter(tx, "logs.username", username); err != nil { + if tx, err = applyExplicitLogTextFilter(tx, "logs.username", params.Username); err != nil { return nil, 0, err } - if tokenName != "" { - tx = tx.Where("logs.token_name = ?", tokenName) + if params.TokenName != "" { + tx = tx.Where("logs.token_name = ?", params.TokenName) } - if requestId != "" { - tx = tx.Where("logs.request_id = ?", requestId) + if params.RequestId != "" { + tx = tx.Where("logs.request_id = ?", params.RequestId) } - if upstreamRequestId != "" { - tx = tx.Where("logs.upstream_request_id = ?", upstreamRequestId) + if params.UpstreamRequestId != "" { + tx = tx.Where("logs.upstream_request_id = ?", params.UpstreamRequestId) } - if startTimestamp != 0 { - tx = tx.Where("logs.created_at >= ?", startTimestamp) + if params.StartTimestamp != 0 { + tx = tx.Where("logs.created_at >= ?", params.StartTimestamp) } - if endTimestamp != 0 { - tx = tx.Where("logs.created_at <= ?", endTimestamp) + if params.EndTimestamp != 0 { + tx = tx.Where("logs.created_at <= ?", params.EndTimestamp) } - if channel != 0 { - tx = tx.Where("logs.channel_id = ?", channel) + if params.Channel != 0 { + tx = tx.Where("logs.channel_id = ?", params.Channel) } - if group != "" { - tx = tx.Where("logs."+logGroupCol+" = ?", group) + if params.Group != "" { + tx = tx.Where("logs."+logGroupCol+" = ?", params.Group) } err = tx.Model(&Log{}).Count(&total).Error if err != nil { return nil, 0, err } - err = tx.Order("logs.created_at desc, logs.id desc").Limit(num).Offset(startIdx).Find(&logs).Error + err = tx.Order("logs.created_at desc, logs.id desc").Limit(params.Num).Offset(params.StartIdx).Find(&logs).Error if err != nil { return nil, 0, err } @@ -692,45 +743,46 @@ func GetAllLogs(logType int, logFilter string, startTimestamp int64, endTimestam const logSearchCountLimit = 10000 -func GetUserLogs(userId int, logType int, logFilter string, startTimestamp int64, endTimestamp int64, modelName string, tokenName string, startIdx int, num int, group string, requestId string, upstreamRequestId string) (logs []*Log, total int64, err error) { - tx, err := applyLogFilter(applyLogTypeFilter(LOG_DB.Where("logs.user_id = ?", userId), logType), logFilter) +func GetUserLogs(params LogQueryParams) (logs []*Log, total int64, err error) { + tx, err := applyLogFilter(applyLogTypeFilter(LOG_DB.Where("logs.user_id = ?", params.UserId), params.LogType), params.LogFilter) if err != nil { return nil, 0, err } + tx = applyQuotaFilter(tx, "logs.quota", params.QuotaFilter) - if tx, err = applyExplicitLogTextFilter(tx, "logs.model_name", modelName); err != nil { + if tx, err = applyExplicitLogTextFilter(tx, "logs.model_name", params.ModelName); err != nil { return nil, 0, err } - if tokenName != "" { - tx = tx.Where("logs.token_name = ?", tokenName) + if params.TokenName != "" { + tx = tx.Where("logs.token_name = ?", params.TokenName) } - if requestId != "" { - tx = tx.Where("logs.request_id = ?", requestId) + if params.RequestId != "" { + tx = tx.Where("logs.request_id = ?", params.RequestId) } - if upstreamRequestId != "" { - tx = tx.Where("logs.upstream_request_id = ?", upstreamRequestId) + if params.UpstreamRequestId != "" { + tx = tx.Where("logs.upstream_request_id = ?", params.UpstreamRequestId) } - if startTimestamp != 0 { - tx = tx.Where("logs.created_at >= ?", startTimestamp) + if params.StartTimestamp != 0 { + tx = tx.Where("logs.created_at >= ?", params.StartTimestamp) } - if endTimestamp != 0 { - tx = tx.Where("logs.created_at <= ?", endTimestamp) + if params.EndTimestamp != 0 { + tx = tx.Where("logs.created_at <= ?", params.EndTimestamp) } - if group != "" { - tx = tx.Where("logs."+logGroupCol+" = ?", group) + if params.Group != "" { + tx = tx.Where("logs."+logGroupCol+" = ?", params.Group) } err = tx.Model(&Log{}).Limit(logSearchCountLimit).Count(&total).Error if err != nil { common.SysError("failed to count user logs: " + err.Error()) return nil, 0, errors.New("查询日志失败") } - err = tx.Order("logs.id desc").Limit(num).Offset(startIdx).Find(&logs).Error + err = tx.Order("logs.id desc").Limit(params.Num).Offset(params.StartIdx).Find(&logs).Error if err != nil { common.SysError("failed to search user logs: " + err.Error()) return nil, 0, errors.New("查询日志失败") } - formatUserLogs(logs, startIdx) + formatUserLogs(logs, params.StartIdx) return logs, total, err } @@ -771,56 +823,58 @@ type Stat struct { Tpm int `json:"tpm"` } -func SumUsedQuota(logType int, logFilter string, startTimestamp int64, endTimestamp int64, modelName string, username string, tokenName string, channel int, group string) (stat Stat, err error) { +func SumUsedQuota(params LogQueryParams) (stat Stat, err error) { tx := LOG_DB.Table("logs").Select("sum(quota) quota") // 为rpm和tpm创建单独的查询 rpmTpmQuery := LOG_DB.Table("logs").Select("count(*) rpm, sum(prompt_tokens) + sum(completion_tokens) tpm") - tx, err = applyLogFilter(tx, logFilter) + tx, err = applyLogFilter(tx, params.LogFilter) if err != nil { return stat, err } - rpmTpmQuery, err = applyLogFilter(rpmTpmQuery, logFilter) + rpmTpmQuery, err = applyLogFilter(rpmTpmQuery, params.LogFilter) if err != nil { return stat, err } + tx = applyQuotaFilter(tx, "logs.quota", params.QuotaFilter) + rpmTpmQuery = applyQuotaFilter(rpmTpmQuery, "logs.quota", params.QuotaFilter) - if tx, err = applyExplicitLogTextFilter(tx, "username", username); err != nil { + if tx, err = applyExplicitLogTextFilter(tx, "username", params.Username); err != nil { return stat, err } - if rpmTpmQuery, err = applyExplicitLogTextFilter(rpmTpmQuery, "username", username); err != nil { + if rpmTpmQuery, err = applyExplicitLogTextFilter(rpmTpmQuery, "username", params.Username); err != nil { return stat, err } - if tokenName != "" { - tx = tx.Where("token_name = ?", tokenName) - rpmTpmQuery = rpmTpmQuery.Where("token_name = ?", tokenName) + if params.TokenName != "" { + tx = tx.Where("token_name = ?", params.TokenName) + rpmTpmQuery = rpmTpmQuery.Where("token_name = ?", params.TokenName) } - if startTimestamp != 0 { - tx = tx.Where("created_at >= ?", startTimestamp) + if params.StartTimestamp != 0 { + tx = tx.Where("created_at >= ?", params.StartTimestamp) } - if endTimestamp != 0 { - tx = tx.Where("created_at <= ?", endTimestamp) + if params.EndTimestamp != 0 { + tx = tx.Where("created_at <= ?", params.EndTimestamp) } - if tx, err = applyExplicitLogTextFilter(tx, "model_name", modelName); err != nil { + if tx, err = applyExplicitLogTextFilter(tx, "model_name", params.ModelName); err != nil { return stat, err } - if rpmTpmQuery, err = applyExplicitLogTextFilter(rpmTpmQuery, "model_name", modelName); err != nil { + if rpmTpmQuery, err = applyExplicitLogTextFilter(rpmTpmQuery, "model_name", params.ModelName); err != nil { return stat, err } - if channel != 0 { - tx = tx.Where("channel_id = ?", channel) - rpmTpmQuery = rpmTpmQuery.Where("channel_id = ?", channel) + if params.Channel != 0 { + tx = tx.Where("channel_id = ?", params.Channel) + rpmTpmQuery = rpmTpmQuery.Where("channel_id = ?", params.Channel) } - if group != "" { - tx = tx.Where(logGroupCol+" = ?", group) - rpmTpmQuery = rpmTpmQuery.Where(logGroupCol+" = ?", group) + if params.Group != "" { + tx = tx.Where(logGroupCol+" = ?", params.Group) + rpmTpmQuery = rpmTpmQuery.Where(logGroupCol+" = ?", params.Group) } - if logType != LogTypeUnknown { - tx = tx.Where("logs.type = ?", logType) - rpmTpmQuery = rpmTpmQuery.Where("logs.type = ?", logType) - } else if !isRetryLogFilter(logFilter) { + if params.LogType != LogTypeUnknown { + tx = tx.Where("logs.type = ?", params.LogType) + rpmTpmQuery = rpmTpmQuery.Where("logs.type = ?", params.LogType) + } else { tx = tx.Where("logs.type = ?", LogTypeConsume) rpmTpmQuery = rpmTpmQuery.Where("logs.type = ?", LogTypeConsume) } diff --git a/model/log_test.go b/model/log_test.go index c105733..36af575 100644 --- a/model/log_test.go +++ b/model/log_test.go @@ -24,6 +24,13 @@ func withLogAuditSettings(t *testing.T, requestEnabled bool, responseEnabled boo }) } +func TestLogQuotaFilterIndexes(t *testing.T) { + db := newRetryBackfillTestDB(t, &Log{}) + + require.True(t, db.Migrator().HasIndex(&Log{}, "idx_logs_quota")) + require.True(t, db.Migrator().HasIndex(&Log{}, "idx_logs_type_quota_created_at")) +} + func TestGetAllLogsRetryFilter(t *testing.T) { require.NoError(t, LOG_DB.Where("1 = 1").Delete(&Log{}).Error) t.Cleanup(func() { @@ -32,7 +39,7 @@ func TestGetAllLogsRetryFilter(t *testing.T) { logs := createRetryFilterLogs(t) - got, total, err := GetAllLogs(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "") + got, total, err := GetAllLogs(LogQueryParams{LogType: LogTypeUnknown, LogFilter: LogFilterRetry, Num: 10}) require.NoError(t, err) require.EqualValues(t, 5, total) require.Len(t, got, 5) @@ -47,7 +54,7 @@ func TestGetUserLogsRetryFilter(t *testing.T) { logs := createRetryFilterLogs(t) - got, total, err := GetUserLogs(1, LogTypeUnknown, LogFilterRetry, 0, 0, "", "", 0, 10, "", "", "") + got, total, err := GetUserLogs(LogQueryParams{UserId: 1, LogType: LogTypeUnknown, LogFilter: LogFilterRetry, Num: 10}) require.NoError(t, err) require.EqualValues(t, 4, total) require.Len(t, got, 4) @@ -67,10 +74,36 @@ func TestSumUsedQuotaRetryFilter(t *testing.T) { createRetryFilterLogs(t) - stat, err := SumUsedQuota(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, "") + stat, err := SumUsedQuota(LogQueryParams{LogType: LogTypeUnknown, LogFilter: LogFilterRetry}) + require.NoError(t, err) + require.Equal(t, 1550, stat.Quota) + require.Equal(t, 4, stat.Rpm) + require.Equal(t, 96, stat.Tpm) +} + +func TestSumUsedQuotaRetryFilterIgnoresNonConsumeQuota(t *testing.T) { + require.NoError(t, LOG_DB.Where("1 = 1").Delete(&Log{}).Error) + t.Cleanup(func() { + require.NoError(t, LOG_DB.Where("1 = 1").Delete(&Log{}).Error) + }) + + createRetryFilterLogs(t) + require.NoError(t, LOG_DB.Create(&Log{ + UserId: 1, + CreatedAt: time.Now().Unix(), + Type: LogTypeTopup, + Quota: 999, + PromptTokens: 1000, + CompletionTokens: 1000, + Other: common.MapToJsonStr(map[string]interface{}{ + "retry_log": true, + }), + }).Error) + + stat, err := SumUsedQuota(LogQueryParams{LogType: LogTypeUnknown, LogFilter: LogFilterRetry}) require.NoError(t, err) require.Equal(t, 1550, stat.Quota) - require.Equal(t, 5, stat.Rpm) + require.Equal(t, 4, stat.Rpm) require.Equal(t, 96, stat.Tpm) } @@ -82,13 +115,13 @@ func TestGetAllLogsRetrySubtypeFilters(t *testing.T) { logs := createRetryFilterLogs(t) - errorLogs, total, err := GetAllLogs(LogTypeUnknown, LogFilterErrorRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "") + errorLogs, total, err := GetAllLogs(LogQueryParams{LogType: LogTypeUnknown, LogFilter: LogFilterErrorRetry, Num: 10}) require.NoError(t, err) require.EqualValues(t, 4, total) require.Len(t, errorLogs, 4) require.ElementsMatch(t, []int{logs[0].Id, logs[4].Id, logs[5].Id, logs[6].Id}, []int{errorLogs[0].Id, errorLogs[1].Id, errorLogs[2].Id, errorLogs[3].Id}) - emptyLogs, total, err := GetAllLogs(LogTypeUnknown, LogFilterEmptyRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "") + emptyLogs, total, err := GetAllLogs(LogQueryParams{LogType: LogTypeUnknown, LogFilter: LogFilterEmptyRetry, Num: 10}) require.NoError(t, err) require.EqualValues(t, 2, total) require.Len(t, emptyLogs, 2) @@ -103,13 +136,13 @@ func TestSumUsedQuotaRetrySubtypeFilters(t *testing.T) { createRetryFilterLogs(t) - errorStat, err := SumUsedQuota(LogTypeUnknown, LogFilterErrorRetry, 0, 0, "", "", "", 0, "") + errorStat, err := SumUsedQuota(LogQueryParams{LogType: LogTypeUnknown, LogFilter: LogFilterErrorRetry}) require.NoError(t, err) require.Equal(t, 1350, errorStat.Quota) - require.Equal(t, 4, errorStat.Rpm) + require.Equal(t, 3, errorStat.Rpm) require.Equal(t, 81, errorStat.Tpm) - emptyStat, err := SumUsedQuota(LogTypeUnknown, LogFilterEmptyRetry, 0, 0, "", "", "", 0, "") + emptyStat, err := SumUsedQuota(LogQueryParams{LogType: LogTypeUnknown, LogFilter: LogFilterEmptyRetry}) require.NoError(t, err) require.Equal(t, 550, emptyStat.Quota) require.Equal(t, 2, emptyStat.Rpm) @@ -124,7 +157,7 @@ func TestSumUsedQuotaAppliesExplicitLogType(t *testing.T) { createRetryFilterLogs(t) - stat, err := SumUsedQuota(LogTypeError, "", 0, 0, "", "", "", 0, "") + stat, err := SumUsedQuota(LogQueryParams{LogType: LogTypeError}) require.NoError(t, err) require.Equal(t, 0, stat.Quota) require.Equal(t, 1, stat.Rpm) @@ -158,7 +191,11 @@ func TestSumUsedQuotaKeepsRpmTpmLiveForHistoricalWindow(t *testing.T) { } require.NoError(t, LOG_DB.Create(&logs).Error) - stat, err := SumUsedQuota(LogTypeUnknown, "", now-90000, now-80000, "", "", "", 0, "") + stat, err := SumUsedQuota(LogQueryParams{ + LogType: LogTypeUnknown, + StartTimestamp: now - 90000, + EndTimestamp: now - 80000, + }) require.NoError(t, err) require.Equal(t, 200, stat.Quota) @@ -166,6 +203,68 @@ func TestSumUsedQuotaKeepsRpmTpmLiveForHistoricalWindow(t *testing.T) { require.Equal(t, 10, stat.Tpm) } +func TestLogQuotaFilters(t *testing.T) { + require.NoError(t, LOG_DB.Where("1 = 1").Delete(&Log{}).Error) + t.Cleanup(func() { + require.NoError(t, LOG_DB.Where("1 = 1").Delete(&Log{}).Error) + }) + + now := time.Now().Unix() + logs := []Log{ + { + UserId: 1, + CreatedAt: now - 10, + Type: LogTypeConsume, + Quota: 0, + PromptTokens: 1, + CompletionTokens: 2, + }, + { + UserId: 1, + CreatedAt: now - 20, + Type: LogTypeConsume, + Quota: -50, + PromptTokens: 3, + CompletionTokens: 4, + }, + { + UserId: 1, + CreatedAt: now - 30, + Type: LogTypeConsume, + Quota: 100, + PromptTokens: 5, + CompletionTokens: 6, + }, + { + UserId: 2, + CreatedAt: now - 40, + Type: LogTypeConsume, + Quota: -75, + PromptTokens: 7, + CompletionTokens: 8, + }, + } + require.NoError(t, LOG_DB.Create(&logs).Error) + + zeroLogs, total, err := GetAllLogs(LogQueryParams{LogType: LogTypeConsume, Num: 10, QuotaFilter: LogQuotaFilterZero}) + require.NoError(t, err) + require.EqualValues(t, 1, total) + require.Len(t, zeroLogs, 1) + require.Equal(t, logs[0].Id, zeroLogs[0].Id) + + negativeLogs, total, err := GetUserLogs(LogQueryParams{UserId: 1, LogType: LogTypeConsume, Num: 10, QuotaFilter: LogQuotaFilterNegative}) + require.NoError(t, err) + require.EqualValues(t, 1, total) + require.Len(t, negativeLogs, 1) + require.Equal(t, logs[1].Id, negativeLogs[0].LogId) + + abnormalStat, err := SumUsedQuota(LogQueryParams{LogType: LogTypeUnknown, QuotaFilter: LogQuotaFilterAbnormal}) + require.NoError(t, err) + require.Equal(t, -125, abnormalStat.Quota) + require.Equal(t, 3, abnormalStat.Rpm) + require.Equal(t, 25, abnormalStat.Tpm) +} + func TestRetryFilterIgnoresNestedRetryMarker(t *testing.T) { require.NoError(t, LOG_DB.Where("1 = 1").Delete(&Log{}).Error) t.Cleanup(func() { @@ -201,13 +300,13 @@ func TestRetryFilterIgnoresNestedRetryMarker(t *testing.T) { } require.NoError(t, LOG_DB.Create(&logs).Error) - got, total, err := GetAllLogs(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "") + got, total, err := GetAllLogs(LogQueryParams{LogType: LogTypeUnknown, LogFilter: LogFilterRetry, Num: 10}) require.NoError(t, err) require.EqualValues(t, 1, total) require.Len(t, got, 1) require.Equal(t, logs[1].Id, got[0].Id) - stat, err := SumUsedQuota(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, "") + stat, err := SumUsedQuota(LogQueryParams{LogType: LogTypeUnknown, LogFilter: LogFilterRetry}) require.NoError(t, err) require.Equal(t, 200, stat.Quota) require.Equal(t, 1, stat.Rpm) @@ -250,7 +349,7 @@ func TestRetryFilterBackfillsLegacyMarkersBeforeCompletion(t *testing.T) { require.NoError(t, LOG_DB.Create(&logs).Error) require.NoError(t, LOG_DB.Model(&Log{}).Where("1 = 1").UpdateColumn("is_retry", false).Error) - got, total, err := GetAllLogs(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "") + got, total, err := GetAllLogs(LogQueryParams{LogType: LogTypeUnknown, LogFilter: LogFilterRetry, Num: 10}) require.NoError(t, err) require.EqualValues(t, 2, total) require.Len(t, got, 2) @@ -289,7 +388,7 @@ func TestRetryFilterUsesIsRetryAfterBackfillCompletion(t *testing.T) { require.NoError(t, LOG_DB.Model(&Log{}).Where("id = ?", log.Id).UpdateColumn("is_retry", false).Error) require.NoError(t, markLogRetryMarkerBackfillCompleted()) - got, total, err := GetAllLogs(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "") + got, total, err := GetAllLogs(LogQueryParams{LogType: LogTypeUnknown, LogFilter: LogFilterRetry, Num: 10}) require.NoError(t, err) require.EqualValues(t, 0, total) require.Empty(t, got) @@ -309,17 +408,17 @@ func TestRetryFilterReadPathsReturnReadinessError(t *testing.T) { ensureLogRetryMarkerBackfillCompletedForRead = originalEnsure }) - got, total, err := GetAllLogs(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "") + got, total, err := GetAllLogs(LogQueryParams{LogType: LogTypeUnknown, LogFilter: LogFilterRetry, Num: 10}) require.ErrorIs(t, err, expectedErr) require.Nil(t, got) require.Zero(t, total) - got, total, err = GetUserLogs(1, LogTypeUnknown, LogFilterRetry, 0, 0, "", "", 0, 10, "", "", "") + got, total, err = GetUserLogs(LogQueryParams{UserId: 1, LogType: LogTypeUnknown, LogFilter: LogFilterRetry, Num: 10}) require.ErrorIs(t, err, expectedErr) require.Nil(t, got) require.Zero(t, total) - stat, err := SumUsedQuota(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, "") + stat, err := SumUsedQuota(LogQueryParams{LogType: LogTypeUnknown, LogFilter: LogFilterRetry}) require.ErrorIs(t, err, expectedErr) require.Zero(t, stat) } diff --git a/model/main.go b/model/main.go index b986b9b..b0524b4 100644 --- a/model/main.go +++ b/model/main.go @@ -311,6 +311,9 @@ func migrateDB() error { if err != nil { return err } + if err := backfillUserNormalizedEmails(); err != nil { + return err + } if common.UsingSQLite { if err := ensureSubscriptionPlanTableSQLite(); err != nil { return err @@ -386,6 +389,9 @@ func migrateDBFast() error { return err } } + if err := backfillUserNormalizedEmails(); err != nil { + return err + } if common.UsingSQLite { if err := ensureSubscriptionPlanTableSQLite(); err != nil { return err @@ -399,6 +405,19 @@ func migrateDBFast() error { return nil } +func backfillUserNormalizedEmails() error { + if DB == nil || !DB.Migrator().HasTable(&User{}) || !DB.Migrator().HasColumn(&User{}, "normalized_email") { + return nil + } + result := DB.Model(&User{}). + Where("email <> ? AND (normalized_email = ? OR normalized_email IS NULL)", "", ""). + Update("normalized_email", gorm.Expr("LOWER(TRIM(email))")) + if result.Error != nil { + return fmt.Errorf("failed to backfill user normalized emails: %w", result.Error) + } + return nil +} + func migrateLOGDB() error { var err error if err = LOG_DB.AutoMigrate(&Log{}); err != nil { diff --git a/model/midjourney.go b/model/midjourney.go index e1a8d77..3abf627 100644 --- a/model/midjourney.go +++ b/model/midjourney.go @@ -31,6 +31,7 @@ type TaskQueryParams struct { MjID string StartTimestamp string EndTimestamp string + QuotaFilter string } func GetAllUserTask(userId int, startIdx int, num int, queryParams TaskQueryParams) []*Midjourney { @@ -39,6 +40,7 @@ func GetAllUserTask(userId int, startIdx int, num int, queryParams TaskQueryPara // 初始化查询构建器 query := DB.Where("user_id = ?", userId) + query = applyQuotaFilter(query, "quota", queryParams.QuotaFilter) if queryParams.MjID != "" { query = query.Where("mj_id = ?", queryParams.MjID) @@ -66,6 +68,7 @@ func GetAllTasks(startIdx int, num int, queryParams TaskQueryParams) []*Midjourn // 初始化查询构建器 query := DB + query = applyQuotaFilter(query, "quota", queryParams.QuotaFilter) // 添加过滤条件 if queryParams.ChannelID != "" { @@ -186,6 +189,7 @@ func MjBulkUpdateByTaskIds(taskIDs []int, params map[string]any) error { func CountAllTasks(queryParams TaskQueryParams) int64 { var total int64 query := DB.Model(&Midjourney{}) + query = applyQuotaFilter(query, "quota", queryParams.QuotaFilter) if queryParams.ChannelID != "" { query = query.Where("channel_id = ?", queryParams.ChannelID) } @@ -206,6 +210,7 @@ func CountAllTasks(queryParams TaskQueryParams) int64 { func CountAllUserTask(userId int, queryParams TaskQueryParams) int64 { var total int64 query := DB.Model(&Midjourney{}).Where("user_id = ?", userId) + query = applyQuotaFilter(query, "quota", queryParams.QuotaFilter) if queryParams.MjID != "" { query = query.Where("mj_id = ?", queryParams.MjID) } diff --git a/model/option.go b/model/option.go index 987c3f8..394a6f1 100644 --- a/model/option.go +++ b/model/option.go @@ -211,6 +211,11 @@ func SyncOptions(frequency int) { } func UpdateOption(key string, value string) error { + var err error + value, err = normalizeOptionUpdateValue(key, value) + if err != nil { + return err + } if err := validateOptionUpdate(key, value); err != nil { return err } @@ -239,13 +244,21 @@ func UpdateOptionsBulk(values map[string]string) error { if len(values) == 0 { return nil } + normalizedValues := make(map[string]string, len(values)) for k, v := range values { + normalized, err := normalizeOptionUpdateValue(k, v) + if err != nil { + return err + } + normalizedValues[k] = normalized + } + for k, v := range normalizedValues { if err := validateOptionUpdate(k, v); err != nil { return err } } err := DB.Transaction(func(tx *gorm.DB) error { - for k, v := range values { + for k, v := range normalizedValues { option := Option{Key: k} if err := tx.FirstOrCreate(&option, Option{Key: k}).Error; err != nil { return err @@ -260,7 +273,7 @@ func UpdateOptionsBulk(values map[string]string) error { if err != nil { return err } - for k, v := range values { + for k, v := range normalizedValues { if err := updateOptionMap(k, v); err != nil { return err } @@ -268,6 +281,15 @@ func UpdateOptionsBulk(values map[string]string) error { return nil } +func normalizeOptionUpdateValue(key string, value string) (string, error) { + switch key { + case "GroupRatio", "group_ratio_setting.group_ratio": + return ratio_setting.NormalizeGroupRatioJSONString(value) + default: + return value, nil + } +} + func validateOptionUpdate(key string, value string) error { switch key { case "Chats": @@ -307,6 +329,10 @@ func validateJSONOption[T any](value string) error { } func updateOptionMap(key string, value string) (err error) { + value, err = normalizeOptionUpdateValue(key, value) + if err != nil { + return err + } if err := validateOptionUpdate(key, value); err != nil { return err } diff --git a/model/option_test.go b/model/option_test.go index 281b774..c4efc8c 100644 --- a/model/option_test.go +++ b/model/option_test.go @@ -44,34 +44,55 @@ func optionExistsForTest(t *testing.T, key string) bool { return count > 0 } -func TestUpdateOptionRejectsAutoRouteGroupRatioNamesBeforePersistence(t *testing.T) { +func optionValueForTest(t *testing.T, key string) string { + t.Helper() + var option Option + require.NoError(t, DB.First(&option, commonKeyCol+" = ?", key).Error) + return option.Value +} + +func optionMapValueForTest(key string) string { + common.OptionMapRWMutex.RLock() + defer common.OptionMapRWMutex.RUnlock() + return common.OptionMap[key] +} + +func TestUpdateOptionFiltersAutoRouteGroupRatioNamesBeforePersistence(t *testing.T) { setupOptionMapTestState(t) + deleteOptionsForTest(t, "GroupRatio", "group_ratio_setting.group_ratio") + t.Cleanup(func() { + deleteOptionsForTest(t, "GroupRatio", "group_ratio_setting.group_ratio") + }) - err := UpdateOption("GroupRatio", `{"auto":1}`) - require.Error(t, err) - require.Contains(t, err.Error(), "auto route namespace") - require.False(t, optionMapContainsForTest("GroupRatio")) + err := UpdateOption("GroupRatio", `{"auto":1,"default":1.25}`) + require.NoError(t, err) + require.NotContains(t, optionValueForTest(t, "GroupRatio"), "auto") + require.NotContains(t, optionMapValueForTest("GroupRatio"), "auto") + require.NotContains(t, ratio_setting.GetGroupRatioCopy(), "auto") + require.Equal(t, 1.25, ratio_setting.GetGroupRatio("default")) err = UpdateOptionsBulk(map[string]string{ - "group_ratio_setting.group_ratio": `{"auto:fast":1}`, + "group_ratio_setting.group_ratio": `{"auto:fast":1,"vip":0.5}`, }) - require.Error(t, err) - require.Contains(t, err.Error(), "auto route namespace") - require.False(t, optionMapContainsForTest("group_ratio_setting.group_ratio")) + require.NoError(t, err) + require.NotContains(t, optionValueForTest(t, "group_ratio_setting.group_ratio"), "auto:fast") + require.NotContains(t, optionMapValueForTest("group_ratio_setting.group_ratio"), "auto:fast") + require.NotContains(t, ratio_setting.GetGroupRatioCopy(), "auto:fast") + require.Equal(t, 0.5, ratio_setting.GetGroupRatio("vip")) } -func TestUpdateOptionMapRejectsAutoRouteGroupRatioNames(t *testing.T) { +func TestUpdateOptionMapFiltersAutoRouteGroupRatioNames(t *testing.T) { setupOptionMapTestState(t) - err := updateOptionMap("GroupRatio", `{"auto":1}`) - require.Error(t, err) - require.Contains(t, err.Error(), "auto route namespace") + err := updateOptionMap("GroupRatio", `{"auto":1,"default":1.25}`) + require.NoError(t, err) require.NotContains(t, ratio_setting.GetGroupRatioCopy(), "auto") + require.Equal(t, 1.25, ratio_setting.GetGroupRatio("default")) - err = updateOptionMap("group_ratio_setting.group_ratio", `{"auto:fast":1}`) - require.Error(t, err) - require.Contains(t, err.Error(), "auto route namespace") + err = updateOptionMap("group_ratio_setting.group_ratio", `{"auto:fast":1,"vip":0.5}`) + require.NoError(t, err) require.NotContains(t, ratio_setting.GetGroupRatioCopy(), "auto:fast") + require.Equal(t, 0.5, ratio_setting.GetGroupRatio("vip")) require.NoError(t, updateOptionMap("GroupRatio", `{"default":1,"vip":0.5}`)) require.Equal(t, 0.5, ratio_setting.GetGroupRatio("vip")) diff --git a/model/payment_method_guard_test.go b/model/payment_method_guard_test.go index 9502ead..c523331 100644 --- a/model/payment_method_guard_test.go +++ b/model/payment_method_guard_test.go @@ -312,6 +312,39 @@ func TestRechargeCreemRejectsZeroQuotaBeforeCompletingOrder(t *testing.T) { assert.Equal(t, 0, getUserQuotaForPaymentGuardTest(t, 610)) } +func TestRechargeCreemSkipsDuplicateCustomerEmailBinding(t *testing.T) { + truncateTables(t) + + require.NoError(t, DB.Create(&User{ + Id: 611, + Username: "creem-email-owner", + Email: "taken@example.com", + AffCode: "creem611", + Status: common.UserStatusEnabled, + }).Error) + require.NoError(t, DB.Create(&User{ + Id: 612, + Username: "creem-empty-email", + AffCode: "creem612", + Status: common.UserStatusEnabled, + }).Error) + insertTopUpForPaymentGuardTest(t, "creem-duplicate-email", 612, PaymentProviderCreem) + + err := RechargeCreem("creem-duplicate-email", " Taken@Example.COM ", "", "127.0.0.1") + require.NoError(t, err) + + var got User + require.NoError(t, DB.First(&got, 612).Error) + assert.Empty(t, got.Email) + assert.Empty(t, got.NormalizedEmail) + assert.Equal(t, 2, got.Quota) + assert.Equal(t, common.TopUpStatusSuccess, getTopUpStatusForPaymentGuardTest(t, "creem-duplicate-email")) + + count, err := CountUsersByEmail("taken@example.com") + require.NoError(t, err) + assert.EqualValues(t, 1, count) +} + func TestRefundSubscriptionPreConsume_IdempotentDoesNotDoubleRefund(t *testing.T) { truncateTables(t) diff --git a/model/task.go b/model/task.go index 9c1929f..978b8b5 100644 --- a/model/task.go +++ b/model/task.go @@ -170,6 +170,7 @@ type SyncTaskQueryParams struct { StartTimestamp int64 EndTimestamp int64 UserIDs []int + QuotaFilter string } func InitTask(platform constant.TaskPlatform, relayInfo *commonRelay.RelayInfo) *Task { @@ -217,6 +218,7 @@ func TaskGetAllUserTask(userId int, startIdx int, num int, queryParams SyncTaskQ // 初始化查询构建器 query := DB.Where("user_id = ?", userId) + query = applyQuotaFilter(query, "quota", queryParams.QuotaFilter) if queryParams.TaskID != "" { query = query.Where("task_id = ?", queryParams.TaskID) @@ -253,6 +255,7 @@ func TaskGetAllTasks(startIdx int, num int, queryParams SyncTaskQueryParams) []* // 初始化查询构建器 query := DB + query = applyQuotaFilter(query, "quota", queryParams.QuotaFilter) // 添加过滤条件 if queryParams.ChannelID != "" { @@ -402,6 +405,10 @@ func (Task *Task) Update() error { return err } +func (t *Task) UpdateQuota() error { + return DB.Model(t).Update("quota", t.Quota).Error +} + // UpdateWithStatus performs a conditional UPDATE guarded by fromStatus (CAS). // Returns (true, nil) if this caller won the update, (false, nil) if // another process already moved the task out of fromStatus. @@ -451,6 +458,7 @@ type TaskQuotaUsage struct { func TaskCountAllTasks(queryParams SyncTaskQueryParams) int64 { var total int64 query := DB.Model(&Task{}) + query = applyQuotaFilter(query, "quota", queryParams.QuotaFilter) if queryParams.ChannelID != "" { query = query.Where("channel_id = ?", queryParams.ChannelID) } @@ -486,6 +494,7 @@ func TaskCountAllTasks(queryParams SyncTaskQueryParams) int64 { func TaskCountAllUserTask(userId int, queryParams SyncTaskQueryParams) int64 { var total int64 query := DB.Model(&Task{}).Where("user_id = ?", userId) + query = applyQuotaFilter(query, "quota", queryParams.QuotaFilter) if queryParams.TaskID != "" { query = query.Where("task_id = ?", queryParams.TaskID) } diff --git a/model/task_cas_test.go b/model/task_cas_test.go index b7cbe4a..4a977d0 100644 --- a/model/task_cas_test.go +++ b/model/task_cas_test.go @@ -244,6 +244,85 @@ func TestUpdateWithStatus_Lose(t *testing.T) { assert.EqualValues(t, TaskStatusFailure, reloaded.Status) // unchanged } +func TestUpdateQuotaScopesByTaskPrimaryKey(t *testing.T) { + truncateTables(t) + + target := &Task{ + TaskID: "task_quota_target", + Status: TaskStatusInProgress, + Quota: 100, + Data: json.RawMessage(`{}`), + } + other := &Task{ + TaskID: "task_quota_other", + Status: TaskStatusInProgress, + Quota: 200, + Data: json.RawMessage(`{}`), + } + insertTask(t, target) + insertTask(t, other) + + target.Quota = 350 + require.NoError(t, target.UpdateQuota()) + + var reloadedTarget Task + require.NoError(t, DB.First(&reloadedTarget, target.ID).Error) + assert.Equal(t, 350, reloadedTarget.Quota) + + var reloadedOther Task + require.NoError(t, DB.First(&reloadedOther, other.ID).Error) + assert.Equal(t, 200, reloadedOther.Quota) +} + +func TestTaskQuotaFilters(t *testing.T) { + truncateTables(t) + + tasks := []*Task{ + { + TaskID: "task_quota_zero", + UserId: 1, + Status: TaskStatusSuccess, + Quota: 0, + Data: json.RawMessage(`{}`), + }, + { + TaskID: "task_quota_negative", + UserId: 1, + Status: TaskStatusSuccess, + Quota: -50, + Data: json.RawMessage(`{}`), + }, + { + TaskID: "task_quota_positive", + UserId: 1, + Status: TaskStatusSuccess, + Quota: 100, + Data: json.RawMessage(`{}`), + }, + { + TaskID: "task_quota_other_user_negative", + UserId: 2, + Status: TaskStatusSuccess, + Quota: -75, + Data: json.RawMessage(`{}`), + }, + } + for _, task := range tasks { + insertTask(t, task) + } + + zeroTasks := TaskGetAllTasks(0, 10, SyncTaskQueryParams{QuotaFilter: LogQuotaFilterZero}) + require.Len(t, zeroTasks, 1) + assert.Equal(t, "task_quota_zero", zeroTasks[0].TaskID) + + negativeUserTasks := TaskGetAllUserTask(1, 0, 10, SyncTaskQueryParams{QuotaFilter: LogQuotaFilterNegative}) + require.Len(t, negativeUserTasks, 1) + assert.Equal(t, "task_quota_negative", negativeUserTasks[0].TaskID) + + assert.EqualValues(t, 3, TaskCountAllTasks(SyncTaskQueryParams{QuotaFilter: LogQuotaFilterAbnormal})) + assert.EqualValues(t, 2, TaskCountAllUserTask(1, SyncTaskQueryParams{QuotaFilter: LogQuotaFilterAbnormal})) +} + func TestUpdateWithStatus_ConcurrentWinner(t *testing.T) { truncateTables(t) diff --git a/model/topup.go b/model/topup.go index c5155e9..763d9ca 100644 --- a/model/topup.go +++ b/model/topup.go @@ -438,7 +438,7 @@ func RechargeCreem(referenceId string, customerEmail string, customerName string refCol = `"trade_no"` } - err = DB.Transaction(func(tx *gorm.DB) error { + err = WithNormalizedEmailWriteTx(customerEmail, func(tx *gorm.DB) error { err := withRowLock(tx).Where(refCol+" = ?", referenceId).First(topUp).Error if err != nil { return errors.New("充值订单不存在") @@ -467,18 +467,8 @@ func RechargeCreem(referenceId string, customerEmail string, customerName string } // 如果有客户邮箱,尝试更新用户邮箱(仅当用户邮箱为空时) - if customerEmail != "" { - // 先检查用户当前邮箱是否为空 - var user User - err = tx.Where("id = ?", topUp.UserId).First(&user).Error - if err != nil { - return err - } - - // 如果用户邮箱为空,则更新为支付时使用的邮箱 - if user.Email == "" { - updateFields["email"] = customerEmail - } + if err := addCreemCustomerEmailUpdateIfAvailable(tx, topUp.UserId, customerEmail, updateFields); err != nil { + return err } result := tx.Model(&User{}).Where("id = ?", topUp.UserId).Updates(updateFields) @@ -499,6 +489,31 @@ func RechargeCreem(referenceId string, customerEmail string, customerName string return nil } +func addCreemCustomerEmailUpdateIfAvailable(tx *gorm.DB, userId int, customerEmail string, updateFields map[string]interface{}) error { + customerEmail = NormalizeEmail(customerEmail) + if customerEmail == "" { + return nil + } + + var user User + if err := tx.Where("id = ?", userId).First(&user).Error; err != nil { + return err + } + if user.Email != "" { + return nil + } + + if err := ensureEmailAvailableWithTx(tx, customerEmail, user.Id); err != nil { + if errors.Is(err, ErrEmailAlreadyTaken) { + return nil + } + return err + } + updateFields["email"] = customerEmail + updateFields["normalized_email"] = customerEmail + return nil +} + func RechargeWaffo(tradeNo string, callerIp string) (err error) { if tradeNo == "" { return errors.New("未提供支付单号") diff --git a/model/user.go b/model/user.go index 572d3aa..5a6031b 100644 --- a/model/user.go +++ b/model/user.go @@ -18,6 +18,33 @@ import ( const UserNameMaxLength = 20 +type UserUpdateField string + +const ( + UserUpdateFieldUsername UserUpdateField = "username" + UserUpdateFieldDisplayName UserUpdateField = "display_name" + UserUpdateFieldRole UserUpdateField = "role" + UserUpdateFieldStatus UserUpdateField = "status" + UserUpdateFieldEmail UserUpdateField = "email" + UserUpdateFieldGitHubId UserUpdateField = "github_id" + UserUpdateFieldDiscordId UserUpdateField = "discord_id" + UserUpdateFieldOidcId UserUpdateField = "oidc_id" + UserUpdateFieldWeChatId UserUpdateField = "wechat_id" + UserUpdateFieldTelegramId UserUpdateField = "telegram_id" + UserUpdateFieldAccessToken UserUpdateField = "access_token" + UserUpdateFieldGroup UserUpdateField = "group" + UserUpdateFieldAffCode UserUpdateField = "aff_code" + UserUpdateFieldAffCount UserUpdateField = "aff_count" + UserUpdateFieldAffQuota UserUpdateField = "aff_quota" + UserUpdateFieldAffHistoryQuota UserUpdateField = "aff_history" + UserUpdateFieldInviterId UserUpdateField = "inviter_id" + UserUpdateFieldLinuxDOId UserUpdateField = "linux_do_id" + UserUpdateFieldSetting UserUpdateField = "setting" + UserUpdateFieldRemark UserUpdateField = "remark" + UserUpdateFieldStripeCustomer UserUpdateField = "stripe_customer" + UserUpdateFieldLastLoginAt UserUpdateField = "last_login_at" +) + // User if you add sensitive fields, don't forget to clean them in setupLogin function. // Otherwise, the sensitive information will be saved on local storage in plain text! type User struct { @@ -29,6 +56,7 @@ type User struct { Role int `json:"role" gorm:"type:int;default:1"` // admin, common Status int `json:"status" gorm:"type:int;default:1"` // enabled, disabled Email string `json:"email" gorm:"index" validate:"max=50"` + NormalizedEmail string `json:"-" gorm:"column:normalized_email;size:50;index"` GitHubId string `json:"github_id" gorm:"column:github_id;index"` DiscordId string `json:"discord_id" gorm:"column:discord_id;index"` OidcId string `json:"oidc_id" gorm:"column:oidc_id;index"` @@ -54,6 +82,16 @@ type User struct { LastLoginAt int64 `json:"last_login_at" gorm:"default:0;column:last_login_at"` } +func (user *User) BeforeSave(_ *gorm.DB) error { + user.NormalizedEmail = NormalizeEmail(user.Email) + return nil +} + +func (user *User) normalizeEmailFields() { + user.Email = NormalizeEmail(user.Email) + user.NormalizedEmail = user.Email +} + func (user *User) ToBaseUser() *UserBase { cache := &UserBase{ Id: user.Id, @@ -187,10 +225,11 @@ func CheckUserExistOrDeleted(username string, email string) (bool, error) { // err := DB.Unscoped().First(&user, "username = ? or email = ?", username, email).Error // check email if empty var err error + email = NormalizeEmail(email) if email == "" { err = DB.Unscoped().First(&user, "username = ?", username).Error } else { - err = DB.Unscoped().First(&user, "username = ? or email = ?", username, email).Error + err = DB.Unscoped().First(&user, "username = ? or normalized_email = ?", username, email).Error } if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { @@ -204,6 +243,146 @@ func CheckUserExistOrDeleted(username string, email string) (bool, error) { return true, nil } +func NormalizeEmail(email string) string { + return strings.ToLower(strings.TrimSpace(email)) +} + +func normalizedEmailLockName(email string) string { + return "maxapi:user-email:" + common.Sha1([]byte(NormalizeEmail(email))) +} + +func emailQuery(tx *gorm.DB, email string) *gorm.DB { + if tx == nil { + tx = DB + } + return tx.Unscoped().Model(&User{}).Where("normalized_email = ?", NormalizeEmail(email)) +} + +func CountUsersByEmail(email string) (int64, error) { + email = NormalizeEmail(email) + if email == "" { + return 0, nil + } + var count int64 + err := emailQuery(DB, email).Count(&count).Error + return count, err +} + +func IsEmailAvailable(email string, excludeUserID int) (bool, error) { + email = NormalizeEmail(email) + if email == "" { + return true, nil + } + query := emailQuery(DB, email) + if excludeUserID > 0 { + query = query.Where("id <> ?", excludeUserID) + } + var count int64 + if err := query.Count(&count).Error; err != nil { + return false, err + } + return count == 0, nil +} + +func EnsureEmailAvailable(email string, excludeUserID int) error { + available, err := IsEmailAvailable(email, excludeUserID) + if err != nil { + return err + } + if !available { + return ErrEmailAlreadyTaken + } + return nil +} + +// withNormalizedEmailLock serializes concurrent writers targeting the same +// normalized email inside tx. SQLite's single-writer model already serializes +// the write path, so it needs no explicit lock here. +func withNormalizedEmailLock(tx *gorm.DB, email string, fn func(tx *gorm.DB) error) error { + email = NormalizeEmail(email) + if email == "" { + return fn(tx) + } + switch { + case common.UsingPostgreSQL: + if err := tx.Exec("SELECT pg_advisory_xact_lock(hashtext(?))", email).Error; err != nil { + return err + } + return fn(tx) + case common.UsingMySQL: + lockName := normalizedEmailLockName(email) + var acquired sql.NullInt64 + if err := tx.Raw("SELECT GET_LOCK(?, ?)", lockName, 10).Scan(&acquired).Error; err != nil { + return err + } + if !acquired.Valid || acquired.Int64 != 1 { + return errors.New("failed to acquire user email lock") + } + released := false + defer func() { + if !released { + _ = releaseMySQLNamedLock(tx, lockName) + } + }() + err := fn(tx) + releaseErr := releaseMySQLNamedLock(tx, lockName) + released = true + if releaseErr != nil && err == nil { + return releaseErr + } + return err + default: + return fn(tx) + } +} + +// WithNormalizedEmailWriteTx runs fn in a transaction while serializing writers +// for the same normalized email. MySQL named locks are connection-scoped, so the +// lock must be acquired on a pinned connection and released only after the +// transaction has committed or rolled back. +func WithNormalizedEmailWriteTx(email string, fn func(tx *gorm.DB) error) error { + email = NormalizeEmail(email) + if !common.UsingMySQL { + return DB.Transaction(func(tx *gorm.DB) error { + return withNormalizedEmailLock(tx, email, fn) + }) + } + if email == "" { + return DB.Transaction(fn) + } + + lockName := normalizedEmailLockName(email) + return DB.Connection(func(conn *gorm.DB) error { + var acquired sql.NullInt64 + if err := conn.Raw("SELECT GET_LOCK(?, ?)", lockName, 10).Scan(&acquired).Error; err != nil { + return err + } + if !acquired.Valid || acquired.Int64 != 1 { + return errors.New("failed to acquire user email lock") + } + + released := false + defer func() { + if !released { + _ = releaseMySQLNamedLock(conn, lockName) + } + }() + + err := conn.Transaction(fn) + releaseErr := releaseMySQLNamedLock(conn, lockName) + released = true + if err != nil { + return err + } + return releaseErr + }) +} + +func releaseMySQLNamedLock(tx *gorm.DB, lockName string) error { + var released sql.NullInt64 + return tx.Raw("SELECT RELEASE_LOCK(?)", lockName).Scan(&released).Error +} + func GetMaxUserId() int { var user User DB.Unscoped().Last(&user) @@ -430,28 +609,79 @@ func (user *User) TransferAffQuotaToQuota(quota int) error { return tx.Commit().Error } -func (user *User) Insert(inviterId int) error { +func (user *User) prepareForInsert(tx *gorm.DB) error { + user.normalizeEmailFields() + if err := ensureEmailAvailableWithTx(tx, user.Email, 0); err != nil { + return err + } + if user.Password == "" { + return nil + } var err error - if user.Password != "" { - user.Password, err = common.Password2Hash(user.Password) - if err != nil { + user.Password, err = common.Password2Hash(user.Password) + return err +} + +func ensureEmailAvailableWithTx(tx *gorm.DB, email string, excludeUserID int) error { + email = NormalizeEmail(email) + if email == "" { + return nil + } + query := emailQuery(tx, email) + if excludeUserID > 0 { + query = query.Where("id <> ?", excludeUserID) + } + var count int64 + if err := query.Count(&count).Error; err != nil { + return err + } + if count > 0 { + return ErrEmailAlreadyTaken + } + return nil +} + +// BindEmailToUser atomically checks email availability and assigns it to the +// user, preventing concurrent binds from sharing the same normalized address. +func BindEmailToUser(user *User, email string) error { + email = NormalizeEmail(email) + if err := WithNormalizedEmailWriteTx(email, func(tx *gorm.DB) error { + if err := ensureEmailAvailableWithTx(tx, email, user.Id); err != nil { + return err + } + if err := tx.Model(&User{}).Where("id = ?", user.Id).Updates(map[string]interface{}{ + "email": email, + "normalized_email": email, + }).Error; err != nil { return err } + user.Email = email + user.NormalizedEmail = email + return tx.First(user, user.Id).Error + }); err != nil { + return err } - user.Quota = common.QuotaForNewUser - //user.SetAccessToken(common.GetUUID()) - user.AffCode = common.GetRandomString(4) + return updateUserCache(*user) +} - // 初始化用户设置,包括默认的边栏配置 - if user.Setting == "" { - defaultSetting := dto.UserSetting{} - // 这里暂时不设置SidebarModules,因为需要在用户创建后根据角色设置 - user.SetSetting(defaultSetting) - } +func (user *User) Insert(inviterId int) error { + if err := WithNormalizedEmailWriteTx(user.Email, func(tx *gorm.DB) error { + if err := user.prepareForInsert(tx); err != nil { + return err + } + user.Quota = common.QuotaForNewUser + user.AffCode = common.GetRandomString(4) + + // 初始化用户设置,包括默认的边栏配置 + if user.Setting == "" { + defaultSetting := dto.UserSetting{} + // 这里暂时不设置SidebarModules,因为需要在用户创建后根据角色设置 + user.SetSetting(defaultSetting) + } - result := DB.Create(user) - if result.Error != nil { - return result.Error + return tx.Create(user).Error + }); err != nil { + return err } // 用户创建成功后,根据角色初始化边栏配置 @@ -464,7 +694,7 @@ func (user *User) Insert(inviterId int) error { currentSetting := createdUser.GetSetting() currentSetting.SidebarModules = defaultSidebarConfig createdUser.SetSetting(currentSetting) - createdUser.Update(false) + createdUser.UpdateFields(false, UserUpdateFieldSetting) common.SysLog(fmt.Sprintf("为新用户 %s (角色: %d) 初始化边栏配置", createdUser.Username, createdUser.Role)) } } @@ -487,31 +717,26 @@ func (user *User) Insert(inviterId int) error { } // InsertWithTx inserts a new user within an existing transaction. +// Callers that own the outer transaction should use WithNormalizedEmailWriteTx +// so MySQL's connection-scoped email lock covers the final commit. // This is used for OAuth registration where user creation and binding need to be atomic. // Post-creation tasks (sidebar config, logs, inviter rewards) are handled after the transaction commits. func (user *User) InsertWithTx(tx *gorm.DB, inviterId int) error { - var err error - if user.Password != "" { - user.Password, err = common.Password2Hash(user.Password) - if err != nil { + return withNormalizedEmailLock(tx, user.Email, func(tx *gorm.DB) error { + if err := user.prepareForInsert(tx); err != nil { return err } - } - user.Quota = common.QuotaForNewUser - user.AffCode = common.GetRandomString(4) - - // 初始化用户设置 - if user.Setting == "" { - defaultSetting := dto.UserSetting{} - user.SetSetting(defaultSetting) - } + user.Quota = common.QuotaForNewUser + user.AffCode = common.GetRandomString(4) - result := tx.Create(user) - if result.Error != nil { - return result.Error - } + // 初始化用户设置 + if user.Setting == "" { + defaultSetting := dto.UserSetting{} + user.SetSetting(defaultSetting) + } - return nil + return tx.Create(user).Error + }) } // FinalizeOAuthUserCreation performs post-transaction tasks for OAuth user creation. @@ -525,7 +750,7 @@ func (user *User) FinalizeOAuthUserCreation(inviterId int) { currentSetting := createdUser.GetSetting() currentSetting.SidebarModules = defaultSidebarConfig createdUser.SetSetting(currentSetting) - createdUser.Update(false) + createdUser.UpdateFields(false, UserUpdateFieldSetting) common.SysLog(fmt.Sprintf("为新用户 %s (角色: %d) 初始化边栏配置", createdUser.Username, createdUser.Role)) } } @@ -546,6 +771,34 @@ func (user *User) FinalizeOAuthUserCreation(inviterId int) { } func (user *User) Update(updatePassword bool) error { + if err := user.updateWithTx(DB, updatePassword, nil); err != nil { + return err + } + if err := updateUserCache(*user); err != nil { + common.SysLog(fmt.Sprintf("failed to update user cache: user_id=%d, error=%v", user.Id, err)) + } + return nil +} + +func (user *User) UpdateFields(updatePassword bool, fields ...UserUpdateField) error { + if err := user.updateWithTx(DB, updatePassword, fields); err != nil { + return err + } + if err := updateUserCache(*user); err != nil { + common.SysLog(fmt.Sprintf("failed to update user cache: user_id=%d, error=%v", user.Id, err)) + } + return nil +} + +func (user *User) UpdateWithTx(tx *gorm.DB, updatePassword bool) error { + return user.updateWithTx(tx, updatePassword, nil) +} + +func (user *User) UpdateFieldsWithTx(tx *gorm.DB, updatePassword bool, fields ...UserUpdateField) error { + return user.updateWithTx(tx, updatePassword, fields) +} + +func (user *User) updateWithTx(tx *gorm.DB, updatePassword bool, fields []UserUpdateField) error { var err error if updatePassword { user.Password, err = common.Password2Hash(user.Password) @@ -555,21 +808,185 @@ func (user *User) Update(updatePassword bool) error { } newUser := *user current := User{} - if err = DB.First(¤t, user.Id).Error; err != nil { + if err = tx.First(¤t, user.Id).Error; err != nil { return err } - result := DB.Model(¤t).Omit("quota", "used_quota", "request_count").Updates(newUser) - if err = ensureUserUpdateMatchedTx(DB, result, user.Id, errors.New("用户不存在")); err != nil { + result := tx.Model(¤t).Updates(buildUserUpdateValues(current, newUser, updatePassword, fields...)) + if err = ensureUserUpdateMatchedTx(tx, result, user.Id, errors.New("用户不存在")); err != nil { return err } - if err = DB.First(user, user.Id).Error; err != nil { - return err + return tx.First(user, user.Id).Error +} + +func buildUserUpdateValues(current User, newUser User, updatePassword bool, fields ...UserUpdateField) map[string]interface{} { + updates := map[string]interface{}{} + + if len(fields) > 0 { + for _, field := range fields { + applyUserUpdateField(updates, newUser, field) + } + } else { + applyNonZeroUserUpdateValues(updates, newUser) + } + if updatePassword { + updates["password"] = newUser.Password } - if err = updateUserCache(*user); err != nil { - common.SysLog(fmt.Sprintf("failed to update user cache: user_id=%d, error=%v", user.Id, err)) + copyUnspecifiedUserUpdateValues(updates, current) + return updates +} + +func applyNonZeroUserUpdateValues(updates map[string]interface{}, newUser User) { + if newUser.Username != "" { + updates["username"] = newUser.Username + } + if newUser.DisplayName != "" { + updates["display_name"] = newUser.DisplayName + } + if newUser.Role != 0 { + updates["role"] = newUser.Role + } + if newUser.Status != 0 { + updates["status"] = newUser.Status + } + if newUser.Email != "" { + email := NormalizeEmail(newUser.Email) + updates["email"] = email + updates["normalized_email"] = email + } + if newUser.GitHubId != "" { + updates["github_id"] = newUser.GitHubId + } + if newUser.DiscordId != "" { + updates["discord_id"] = newUser.DiscordId + } + if newUser.OidcId != "" { + updates["oidc_id"] = newUser.OidcId + } + if newUser.WeChatId != "" { + updates["wechat_id"] = newUser.WeChatId + } + if newUser.TelegramId != "" { + updates["telegram_id"] = newUser.TelegramId + } + if newUser.AccessToken != nil { + updates["access_token"] = newUser.AccessToken + } + if newUser.Group != "" { + updates["group"] = newUser.Group + } + if newUser.AffCode != "" { + updates["aff_code"] = newUser.AffCode + } + if newUser.AffCount != 0 { + updates["aff_count"] = newUser.AffCount + } + if newUser.AffQuota != 0 { + updates["aff_quota"] = newUser.AffQuota + } + if newUser.AffHistoryQuota != 0 { + updates["aff_history"] = newUser.AffHistoryQuota + } + if newUser.InviterId != 0 { + updates["inviter_id"] = newUser.InviterId + } + if newUser.LinuxDOId != "" { + updates["linux_do_id"] = newUser.LinuxDOId + } + if newUser.Setting != "" { + updates["setting"] = newUser.Setting + } + if newUser.Remark != "" { + updates["remark"] = newUser.Remark + } + if newUser.StripeCustomer != "" { + updates["stripe_customer"] = newUser.StripeCustomer + } + if newUser.LastLoginAt != 0 { + updates["last_login_at"] = newUser.LastLoginAt + } +} + +func applyUserUpdateField(updates map[string]interface{}, newUser User, field UserUpdateField) { + switch field { + case UserUpdateFieldUsername: + updates["username"] = newUser.Username + case UserUpdateFieldDisplayName: + updates["display_name"] = newUser.DisplayName + case UserUpdateFieldRole: + updates["role"] = newUser.Role + case UserUpdateFieldStatus: + updates["status"] = newUser.Status + case UserUpdateFieldEmail: + email := NormalizeEmail(newUser.Email) + updates["email"] = email + updates["normalized_email"] = email + case UserUpdateFieldGitHubId: + updates["github_id"] = newUser.GitHubId + case UserUpdateFieldDiscordId: + updates["discord_id"] = newUser.DiscordId + case UserUpdateFieldOidcId: + updates["oidc_id"] = newUser.OidcId + case UserUpdateFieldWeChatId: + updates["wechat_id"] = newUser.WeChatId + case UserUpdateFieldTelegramId: + updates["telegram_id"] = newUser.TelegramId + case UserUpdateFieldAccessToken: + updates["access_token"] = newUser.AccessToken + case UserUpdateFieldGroup: + updates["group"] = newUser.Group + case UserUpdateFieldAffCode: + updates["aff_code"] = newUser.AffCode + case UserUpdateFieldAffCount: + updates["aff_count"] = newUser.AffCount + case UserUpdateFieldAffQuota: + updates["aff_quota"] = newUser.AffQuota + case UserUpdateFieldAffHistoryQuota: + updates["aff_history"] = newUser.AffHistoryQuota + case UserUpdateFieldInviterId: + updates["inviter_id"] = newUser.InviterId + case UserUpdateFieldLinuxDOId: + updates["linux_do_id"] = newUser.LinuxDOId + case UserUpdateFieldSetting: + updates["setting"] = newUser.Setting + case UserUpdateFieldRemark: + updates["remark"] = newUser.Remark + case UserUpdateFieldStripeCustomer: + updates["stripe_customer"] = newUser.StripeCustomer + case UserUpdateFieldLastLoginAt: + updates["last_login_at"] = newUser.LastLoginAt + } +} + +func copyUnspecifiedUserUpdateValues(updates map[string]interface{}, current User) { + defaults := map[string]interface{}{ + "role": current.Role, + "status": current.Status, + "email": current.Email, + "normalized_email": NormalizeEmail(current.Email), + "github_id": current.GitHubId, + "discord_id": current.DiscordId, + "oidc_id": current.OidcId, + "wechat_id": current.WeChatId, + "telegram_id": current.TelegramId, + "access_token": current.AccessToken, + "group": current.Group, + "aff_code": current.AffCode, + "aff_count": current.AffCount, + "aff_quota": current.AffQuota, + "aff_history": current.AffHistoryQuota, + "inviter_id": current.InviterId, + "linux_do_id": current.LinuxDOId, + "setting": current.Setting, + "remark": current.Remark, + "stripe_customer": current.StripeCustomer, + "last_login_at": current.LastLoginAt, + } + for key, value := range defaults { + if _, ok := updates[key]; !ok { + updates[key] = value + } } - return nil } func (user *User) Edit(updatePassword bool) error { @@ -620,7 +1037,11 @@ func (user *User) ClearBinding(bindingType string) error { return errors.New("invalid binding type") } - if err := DB.Model(&User{}).Where("id = ?", user.Id).Update(column, "").Error; err != nil { + updates := map[string]interface{}{column: ""} + if bindingType == "email" { + updates["normalized_email"] = "" + } + if err := DB.Model(&User{}).Where("id = ?", user.Id).Updates(updates).Error; err != nil { return err } @@ -671,6 +1092,9 @@ func (user *User) ValidateAndFill() (err error) { } return fmt.Errorf("%w: %v", ErrDatabase, err) } + if user.Password == "" { + return ErrInvalidCredentials + } okay := common.ValidatePasswordAndHash(password, user.Password) if !okay || user.Status != common.UserStatusEnabled { return ErrInvalidCredentials @@ -746,7 +1170,27 @@ func (user *User) FillUserByTelegramId() error { } func IsEmailAlreadyTaken(email string) bool { - return DB.Unscoped().Where("email = ?", email).Find(&User{}).RowsAffected == 1 + count, err := CountUsersByEmail(email) + return err == nil && count > 0 +} + +func GetUniqueUserByEmail(email string) (*User, error) { + email = NormalizeEmail(email) + if email == "" { + return nil, ErrEmailNotFound + } + var users []User + if err := DB.Where("normalized_email = ?", email).Limit(2).Find(&users).Error; err != nil { + return nil, err + } + switch len(users) { + case 0: + return nil, ErrEmailNotFound + case 1: + return &users[0], nil + default: + return nil, ErrEmailAmbiguous + } } func IsWeChatIdAlreadyTaken(wechatId string) bool { @@ -773,11 +1217,15 @@ func ResetUserPasswordByEmail(email string, password string) error { if email == "" || password == "" { return errors.New("邮箱地址或密码为空!") } + user, err := GetUniqueUserByEmail(email) + if err != nil { + return err + } hashedPassword, err := common.Password2Hash(password) if err != nil { return err } - err = DB.Model(&User{}).Where("email = ?", email).Update("password", hashedPassword).Error + err = DB.Model(&User{}).Where("id = ?", user.Id).Update("password", hashedPassword).Error return err } diff --git a/model/user_update_test.go b/model/user_update_test.go index e75bbe2..4bc68c5 100644 --- a/model/user_update_test.go +++ b/model/user_update_test.go @@ -3,7 +3,11 @@ package model import ( "context" "errors" + "fmt" "net" + "net/url" + "os" + "strings" "testing" "github.com/MAX-API-Next/MAX-API/common" @@ -11,6 +15,7 @@ import ( "github.com/go-redis/redis/v8" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + gormmysql "gorm.io/driver/mysql" "gorm.io/gorm" ) @@ -49,6 +54,85 @@ func useFailingUserUpdateRedis(t *testing.T) { }) } +func openUserUpdateMySQLTestDB(t *testing.T, dsn string) *gorm.DB { + t.Helper() + + db, err := gorm.Open(gormmysql.Open(mysqlDSNWithClientFoundRowsFalse(dsn)), &gorm.Config{}) + require.NoError(t, err) + if db.Migrator().HasTable(&User{}) { + t.Skip("refusing to run mysql user update test against external database because users table already exists") + } + require.NoError(t, db.AutoMigrate(&User{})) + t.Cleanup(func() { + _ = db.Migrator().DropTable(&User{}) + if sqlDB, err := db.DB(); err == nil { + _ = sqlDB.Close() + } + }) + return db +} + +func TestNormalizedEmailLockNameIsStableAndShort(t *testing.T) { + lockName := normalizedEmailLockName(" User@Example.COM ") + + assert.Equal(t, normalizedEmailLockName("user@example.com"), lockName) + assert.LessOrEqual(t, len(lockName), 64) + assert.Contains(t, lockName, "maxapi:user-email:") +} + +func mysqlDSNWithClientFoundRowsFalse(dsn string) string { + base, rawQuery, ok := strings.Cut(dsn, "?") + values, err := url.ParseQuery(rawQuery) + if err != nil { + if ok { + return dsn + "&clientFoundRows=false" + } + return dsn + "?clientFoundRows=false" + } + for key := range values { + if strings.EqualFold(key, "clientFoundRows") { + delete(values, key) + } + } + values.Set("clientFoundRows", "false") + if !ok { + return base + "?" + values.Encode() + } + return base + "?" + values.Encode() +} + +func useUserUpdateTestDB(t *testing.T, db *gorm.DB) { + t.Helper() + + oldDB := DB + oldLOGDB := LOG_DB + oldUsingSQLite := common.UsingSQLite + oldUsingMySQL := common.UsingMySQL + oldUsingPostgreSQL := common.UsingPostgreSQL + oldRedisEnabled := common.RedisEnabled + oldBatchUpdateEnabled := common.BatchUpdateEnabled + + DB = db + LOG_DB = db + common.UsingSQLite = false + common.UsingMySQL = true + common.UsingPostgreSQL = false + common.RedisEnabled = false + common.BatchUpdateEnabled = false + initCol() + + t.Cleanup(func() { + DB = oldDB + LOG_DB = oldLOGDB + common.UsingSQLite = oldUsingSQLite + common.UsingMySQL = oldUsingMySQL + common.UsingPostgreSQL = oldUsingPostgreSQL + common.RedisEnabled = oldRedisEnabled + common.BatchUpdateEnabled = oldBatchUpdateEnabled + initCol() + }) +} + func TestUserUpdateDoesNotOverwriteAccountingFields(t *testing.T) { setupUserUpdateTestState(t) @@ -84,6 +168,69 @@ func TestUserUpdateDoesNotOverwriteAccountingFields(t *testing.T) { assert.Equal(t, 4, got.RequestCount) } +func TestUserUpdatePersistsZeroValueProfileFields(t *testing.T) { + setupUserUpdateTestState(t) + + user := User{ + Id: 5, + Username: "zero-value-user", + Password: "password", + DisplayName: "display", + Status: common.UserStatusEnabled, + AffCount: 7, + Quota: 1000, + UsedQuota: 20, + RequestCount: 3, + } + require.NoError(t, DB.Create(&user).Error) + + loaded, err := GetUserById(user.Id, true) + require.NoError(t, err) + loaded.DisplayName = "" + loaded.AffCount = 0 + loaded.Quota = 1 + loaded.UsedQuota = 1 + loaded.RequestCount = 1 + + require.NoError(t, loaded.UpdateFields(false, UserUpdateFieldDisplayName, UserUpdateFieldAffCount)) + + var got User + require.NoError(t, DB.First(&got, user.Id).Error) + assert.Empty(t, got.DisplayName) + assert.Zero(t, got.AffCount) + assert.Equal(t, 1000, got.Quota) + assert.Equal(t, 20, got.UsedQuota) + assert.Equal(t, 3, got.RequestCount) +} + +func TestUserUpdateDoesNotClearEmailFromLoadedPartialUpdate(t *testing.T) { + setupUserUpdateTestState(t) + + user := User{ + Id: 6, + Username: "email-preserve-user", + Password: "password", + DisplayName: "before", + Email: "Keep@Example.COM", + Status: common.UserStatusEnabled, + } + require.NoError(t, DB.Create(&user).Error) + + loaded, err := GetUserById(user.Id, true) + require.NoError(t, err) + loaded.Email = "" + loaded.NormalizedEmail = "" + loaded.DisplayName = "after" + + require.NoError(t, loaded.Update(false)) + + var got User + require.NoError(t, DB.First(&got, user.Id).Error) + assert.Equal(t, "after", got.DisplayName) + assert.Equal(t, "Keep@Example.COM", got.Email) + assert.Equal(t, "keep@example.com", got.NormalizedEmail) +} + func TestUserUpdateIgnoresCacheWriteFailure(t *testing.T) { setupUserUpdateTestState(t) @@ -137,6 +284,94 @@ func TestUpdateUserSettingOnlyUpdatesSetting(t *testing.T) { require.NoError(t, UpdateUserSetting(user.Id, dto.UserSetting{Language: "zh"})) } +func TestUpdateUserSettingOnlyUpdatesSettingMySQL(t *testing.T) { + dsn := os.Getenv("TEST_MYSQL_DSN") + if dsn == "" { + t.Skip("set TEST_MYSQL_DSN to run mysql UpdateUserSetting coverage") + } + db := openUserUpdateMySQLTestDB(t, dsn) + useUserUpdateTestDB(t, db) + + user := User{ + Id: 2, + Username: "mysql-setting-user", + Password: "password", + Status: common.UserStatusEnabled, + Quota: 1000, + UsedQuota: 20, + RequestCount: 3, + } + require.NoError(t, DB.Create(&user).Error) + + require.NoError(t, DB.Model(&User{}).Where("id = ?", user.Id).Updates(map[string]interface{}{ + "quota": gorm.Expr("quota - ?", 250), + "used_quota": gorm.Expr("used_quota + ?", 250), + "request_count": gorm.Expr("request_count + ?", 1), + }).Error) + + require.NoError(t, UpdateUserSetting(user.Id, dto.UserSetting{Language: "zh"})) + + var got User + require.NoError(t, DB.First(&got, user.Id).Error) + assert.Equal(t, 750, got.Quota) + assert.Equal(t, 270, got.UsedQuota) + assert.Equal(t, 4, got.RequestCount) + assert.Equal(t, "zh", got.GetSetting().Language) + + require.NoError(t, UpdateUserSetting(user.Id, dto.UserSetting{Language: "zh"})) +} + +func TestInsertRejectsConcurrentDuplicateEmailMySQL(t *testing.T) { + dsn := os.Getenv("TEST_MYSQL_DSN") + if dsn == "" { + t.Skip("set TEST_MYSQL_DSN to run mysql concurrent email insert coverage") + } + db := openUserUpdateMySQLTestDB(t, dsn) + useUserUpdateTestDB(t, db) + + oldQuotaForNewUser := common.QuotaForNewUser + common.QuotaForNewUser = 0 + t.Cleanup(func() { + common.QuotaForNewUser = oldQuotaForNewUser + }) + + start := make(chan struct{}) + errs := make(chan error, 2) + for i := range 2 { + go func(i int) { + <-start + user := &User{ + Username: fmt.Sprintf("mysql-race-user-%d", i), + Email: "Race@Example.COM", + Role: common.RoleCommonUser, + Status: common.UserStatusEnabled, + } + errs <- user.Insert(0) + }(i) + } + close(start) + + var successCount int + var duplicateCount int + for range 2 { + err := <-errs + switch { + case err == nil: + successCount++ + case errors.Is(err, ErrEmailAlreadyTaken): + duplicateCount++ + default: + require.NoError(t, err) + } + } + + require.Equal(t, 1, successCount) + require.Equal(t, 1, duplicateCount) + count, err := CountUsersByEmail("race@example.com") + require.NoError(t, err) + require.EqualValues(t, 1, count) +} + func TestUpdateUserSettingIgnoresCacheWriteFailure(t *testing.T) { setupUserUpdateTestState(t) @@ -164,3 +399,153 @@ func TestUpdateUserSettingMissingUserReturnsError(t *testing.T) { require.Error(t, err) assert.Contains(t, err.Error(), "用户不存在") } + +func TestEnsureEmailAvailableRejectsExistingEmailCaseInsensitive(t *testing.T) { + setupUserUpdateTestState(t) + + require.NoError(t, DB.Create(&User{ + Username: "existing", + Password: "old-password", + Email: "Taken@Example.com", + Status: common.UserStatusEnabled, + }).Error) + + err := EnsureEmailAvailable(" taken@example.COM ", 0) + require.ErrorIs(t, err, ErrEmailAlreadyTaken) + + user, err := GetUniqueUserByEmail("TAKEN@example.com") + require.NoError(t, err) + assert.Equal(t, "existing", user.Username) + + require.NoError(t, EnsureEmailAvailable("taken@example.com", user.Id)) +} + +func TestBackfillUserNormalizedEmails(t *testing.T) { + setupUserUpdateTestState(t) + + require.NoError(t, DB.Exec( + "INSERT INTO users (id, username, password, email, status) VALUES (?, ?, ?, ?, ?)", + 21, + "legacy-email-user", + "password", + "Legacy@Example.COM", + common.UserStatusEnabled, + ).Error) + + require.NoError(t, backfillUserNormalizedEmails()) + + var got User + require.NoError(t, DB.First(&got, 21).Error) + assert.Equal(t, "Legacy@Example.COM", got.Email) + assert.Equal(t, "legacy@example.com", got.NormalizedEmail) + + require.ErrorIs(t, EnsureEmailAvailable(" legacy@example.com ", 0), ErrEmailAlreadyTaken) + require.NoError(t, EnsureEmailAvailable("legacy@example.com", got.Id)) +} + +func TestInsertRejectsDuplicateEmailWithoutUniqueIndex(t *testing.T) { + setupUserUpdateTestState(t) + + require.NoError(t, DB.Create(&User{ + Username: "existing", + Password: "old-password", + Email: "taken@example.com", + Status: common.UserStatusEnabled, + }).Error) + + user := &User{ + Username: "oauth-user", + Email: "TAKEN@example.com", + Role: common.RoleCommonUser, + Status: common.UserStatusEnabled, + } + + err := user.Insert(0) + require.ErrorIs(t, err, ErrEmailAlreadyTaken) + + var count int64 + require.NoError(t, DB.Model(&User{}).Where("username = ?", "oauth-user").Count(&count).Error) + assert.Zero(t, count) +} + +func TestInsertKeepsBlankPasswordForPasswordlessUser(t *testing.T) { + setupUserUpdateTestState(t) + + user := &User{ + Username: "passwordless-user", + Role: common.RoleCommonUser, + Status: common.UserStatusEnabled, + } + + require.NoError(t, user.Insert(0)) + + var stored User + require.NoError(t, DB.Where("username = ?", user.Username).First(&stored).Error) + assert.Empty(t, stored.Password) +} + +func TestValidateAndFillRejectsPasswordlessUser(t *testing.T) { + setupUserUpdateTestState(t) + + require.NoError(t, DB.Create(&User{ + Username: "passwordless-user", + Password: "", + Status: common.UserStatusEnabled, + }).Error) + + loginUser := User{ + Username: "passwordless-user", + Password: "NewPassword123", + } + err := loginUser.ValidateAndFill() + require.ErrorIs(t, err, ErrInvalidCredentials) + + var stored User + require.NoError(t, DB.Where("username = ?", "passwordless-user").First(&stored).Error) + assert.Empty(t, stored.Password) +} + +func TestResetUserPasswordByEmailRequiresSingleActiveMatch(t *testing.T) { + setupUserUpdateTestState(t) + + require.NoError(t, DB.Create(&User{ + Username: "duplicate-1", + Password: "old-1", + Email: "legacy@example.com", + AffCode: "dupe1", + Status: common.UserStatusEnabled, + }).Error) + require.NoError(t, DB.Create(&User{ + Username: "duplicate-2", + Password: "old-2", + Email: "LEGACY@example.com", + AffCode: "dupe2", + Status: common.UserStatusEnabled, + }).Error) + + err := ResetUserPasswordByEmail("legacy@example.com", "NewPassword123") + require.ErrorIs(t, err, ErrEmailAmbiguous) + + var duplicates []User + require.NoError(t, DB.Where("LOWER(email) = ?", "legacy@example.com").Order("username asc").Find(&duplicates).Error) + require.Len(t, duplicates, 2) + assert.Equal(t, "old-1", duplicates[0].Password) + assert.Equal(t, "old-2", duplicates[1].Password) + + require.NoError(t, DB.Create(&User{ + Username: "unique", + Password: "old", + Email: "unique@example.com", + AffCode: "unique", + Status: common.UserStatusEnabled, + }).Error) + + require.NoError(t, ResetUserPasswordByEmail("UNIQUE@example.com", "NewPassword123")) + + var unique User + require.NoError(t, DB.Where("username = ?", "unique").First(&unique).Error) + assert.True(t, common.ValidatePasswordAndHash("NewPassword123", unique.Password)) + + err = ResetUserPasswordByEmail("missing@example.com", "NewPassword123") + require.True(t, errors.Is(err, ErrEmailNotFound)) +} diff --git a/pkg/billingexpr/billingexpr_test.go b/pkg/billingexpr/billingexpr_test.go index 699bc18..5ae29f6 100644 --- a/pkg/billingexpr/billingexpr_test.go +++ b/pkg/billingexpr/billingexpr_test.go @@ -292,6 +292,10 @@ func TestQuotaRound(t *testing.T) { {999.4999, 999}, {999.5, 1000}, {1e9 + 0.5, 1e9 + 1}, + {3.6893488147419103e19, math.MaxInt32}, + {-3.6893488147419103e19, math.MinInt32}, + {math.Inf(1), math.MaxInt32}, + {math.NaN(), 0}, } for _, tt := range tests { got := billingexpr.QuotaRound(tt.in) diff --git a/pkg/billingexpr/round.go b/pkg/billingexpr/round.go index 35a5534..fbdc64d 100644 --- a/pkg/billingexpr/round.go +++ b/pkg/billingexpr/round.go @@ -5,6 +5,19 @@ import "math" // QuotaRound converts a float64 quota value to int using half-away-from-zero // rounding. Every tiered billing path (pre-consume, settlement, breakdown // validation, log fields) MUST use this function to avoid +-1 discrepancies. +// +// The result saturates at int32 bounds: quota columns are 32-bit integers in +// the database, and oversized expression results must never wrap around. func QuotaRound(f float64) int { - return int(math.Round(f)) + r := math.Round(f) + if math.IsNaN(r) { + return 0 + } + if r >= math.MaxInt32 { + return math.MaxInt32 + } + if r <= math.MinInt32 { + return math.MinInt32 + } + return int(r) } diff --git a/relay/channel/ali/image.go b/relay/channel/ali/image.go index d065657..f681c82 100644 --- a/relay/channel/ali/image.go +++ b/relay/channel/ali/image.go @@ -54,6 +54,9 @@ func oaiImage2AliImageRequest(info *relaycommon.RelayInfo, request dto.ImageRequ } } + if err := dto.ValidateImageN("parameters.n", imageRequest.Parameters.N); err != nil { + return nil, err + } if imageRequest.Parameters.N != 0 { info.PriceData.AddOtherRatio("n", float64(imageRequest.Parameters.N)) } diff --git a/relay/channel/api_request.go b/relay/channel/api_request.go index 0ee2a73..53aac67 100644 --- a/relay/channel/api_request.go +++ b/relay/channel/api_request.go @@ -64,6 +64,8 @@ const ( headerPassthroughRegexPrefixV2 = "regex:" ) +var sendPingDataTimeout = 10 * time.Second + var passthroughSkipHeaderNamesLower = map[string]struct{}{ // RFC 7230 hop-by-hop headers. "connection": {}, @@ -400,10 +402,12 @@ func DoWssRequest(a Adaptor, c *gin.Context, info *common.RelayInfo, requestBody return targetConn, nil } -func startPingKeepAlive(c *gin.Context, pingInterval time.Duration) context.CancelFunc { +func startPingKeepAlive(c *gin.Context, pingInterval time.Duration) (context.CancelFunc, <-chan struct{}) { pingerCtx, stopPinger := context.WithCancel(context.Background()) + done := make(chan struct{}) gopool.Go(func() { + defer close(done) defer func() { // 增加panic恢复处理 if r := recover(); r != nil { @@ -453,16 +457,16 @@ func startPingKeepAlive(c *gin.Context, pingInterval time.Duration) context.Canc } }) - return stopPinger + return stopPinger, done } func sendPingData(c *gin.Context, mutex *sync.Mutex) error { - // 增加超时控制,防止锁死等待 done := make(chan error, 1) go func() { mutex.Lock() defer mutex.Unlock() + helper.ExtendWriteDeadline(c) err := helper.PingData(c) if err != nil { logger.LogError(c, "SSE ping error: "+err.Error()) @@ -474,14 +478,24 @@ func sendPingData(c *gin.Context, mutex *sync.Mutex) error { done <- nil }() - // 设置发送ping数据的超时时间 + timer := time.NewTimer(sendPingDataTimeout) + defer timer.Stop() + + var requestDone <-chan struct{} + if c != nil && c.Request != nil { + requestDone = c.Request.Context().Done() + } + select { case err := <-done: return err - case <-time.After(10 * time.Second): - return errors.New("SSE ping data send timeout") - case <-c.Request.Context().Done(): - return errors.New("request context cancelled during ping") + case <-requestDone: + if c != nil && c.Request != nil && c.Request.Context().Err() != nil { + return fmt.Errorf("SSE ping request context done: %w", c.Request.Context().Err()) + } + return errors.New("SSE ping request context done") + case <-timer.C: + return fmt.Errorf("SSE ping write timed out after %s", sendPingDataTimeout) } } @@ -501,17 +515,19 @@ func doRequest(c *gin.Context, req *http.Request, info *common.RelayInfo) (*http } var stopPinger context.CancelFunc + var pingerDone <-chan struct{} if info.IsStream { helper.SetEventStreamHeaders(c) // 处理流式请求的 ping 保活 generalSettings := operation_setting.GetGeneralSetting() if generalSettings.PingIntervalEnabled && !info.DisablePing { pingInterval := time.Duration(generalSettings.PingIntervalSeconds) * time.Second - stopPinger = startPingKeepAlive(c, pingInterval) + stopPinger, pingerDone = startPingKeepAlive(c, pingInterval) // 使用defer确保在任何情况下都能停止ping goroutine defer func() { if stopPinger != nil { stopPinger() + <-pingerDone logger.LogDebug(c, "SSE ping goroutine stopped by defer") } }() @@ -580,6 +596,9 @@ func attachSeekableGetBody(req *http.Request, reader io.Reader) { } func DoTaskApiRequest(a TaskAdaptor, c *gin.Context, info *common.RelayInfo, requestBody io.Reader) (*http.Response, error) { + if info != nil && info.ChannelMeta == nil { + info.InitChannelMeta(c) + } fullRequestURL, err := a.BuildRequestURL(info) if err != nil { return nil, err diff --git a/relay/channel/api_request_test.go b/relay/channel/api_request_test.go index b226b0b..e5646a7 100644 --- a/relay/channel/api_request_test.go +++ b/relay/channel/api_request_test.go @@ -5,8 +5,12 @@ import ( "net/http" "net/http/httptest" "strings" + "sync" "testing" + "time" + common2 "github.com/MAX-API-Next/MAX-API/common" + "github.com/MAX-API-Next/MAX-API/constant" "github.com/MAX-API-Next/MAX-API/dto" "github.com/MAX-API-Next/MAX-API/model" "github.com/MAX-API-Next/MAX-API/relay/channel/task/taskcommon" @@ -89,6 +93,23 @@ func (r *customReadSeeker) Seek(offset int64, whence int) (int64, error) { return r.reader.Seek(offset, whence) } +type blockingPingWriter struct { + gin.ResponseWriter + started chan struct{} + release chan struct{} + finished chan struct{} + once sync.Once +} + +func (w *blockingPingWriter) Write(p []byte) (int, error) { + w.once.Do(func() { + close(w.started) + }) + <-w.release + close(w.finished) + return len(p), nil +} + func TestNewTaskHTTPRequestDoesNotPreReadRequestBody(t *testing.T) { t.Parallel() @@ -229,6 +250,38 @@ func TestProcessHeaderOverride_RuntimeOverrideIsFinalHeaderMap(t *testing.T) { require.False(t, exists) } +func TestDoTaskApiRequestInitializesChannelMetaForApiKeyPlaceholder(t *testing.T) { + t.Parallel() + + gin.SetMode(gin.TestMode) + service.InitHttpClient() + var gotAuthorization string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotAuthorization = r.Header.Get("Authorization") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"id":"task_123"}`)) + })) + t.Cleanup(server.Close) + + recorder := httptest.NewRecorder() + ctx, _ := gin.CreateTestContext(recorder) + ctx.Request = httptest.NewRequest(http.MethodPost, "/v1/videos", nil) + common2.SetContextKey(ctx, constant.ContextKeyChannelKey, "sk-task") + + info := &relaycommon.RelayInfo{ + UseRuntimeHeadersOverride: true, + RuntimeHeadersOverride: map[string]any{ + "Authorization": "Bearer {api_key}", + }, + } + + resp, err := DoTaskApiRequest(taskHeaderAdaptor{url: server.URL}, ctx, info, strings.NewReader(`{"prompt":"test"}`)) + require.NoError(t, err) + require.NotNil(t, resp) + require.NoError(t, resp.Body.Close()) + require.Equal(t, "Bearer sk-task", gotAuthorization) +} + func TestProcessHeaderOverride_PassthroughSkipsAcceptEncoding(t *testing.T) { t.Parallel() @@ -346,3 +399,49 @@ func TestDoTaskApiRequestAppliesRuntimeHeaderOverride(t *testing.T) { require.Equal(t, "overridden", gotDefault) require.Equal(t, "enabled", gotRuntime) } + +func TestSendPingDataReturnsWhenWriterBlocks(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + ctx, _ := gin.CreateTestContext(recorder) + ctx.Request = httptest.NewRequest(http.MethodGet, "/stream", nil) + writer := &blockingPingWriter{ + ResponseWriter: ctx.Writer, + started: make(chan struct{}), + release: make(chan struct{}), + finished: make(chan struct{}), + } + ctx.Writer = writer + + oldTimeout := sendPingDataTimeout + sendPingDataTimeout = 20 * time.Millisecond + t.Cleanup(func() { + sendPingDataTimeout = oldTimeout + }) + + errCh := make(chan error, 1) + var mutex sync.Mutex + go func() { + errCh <- sendPingData(ctx, &mutex) + }() + + select { + case <-writer.started: + case <-time.After(time.Second): + t.Fatal("ping write did not start") + } + + select { + case err := <-errCh: + require.ErrorContains(t, err, "timed out") + case <-time.After(time.Second): + t.Fatal("sendPingData did not return after timeout") + } + + close(writer.release) + select { + case <-writer.finished: + case <-time.After(time.Second): + t.Fatal("blocked ping writer did not finish after release") + } +} diff --git a/relay/channel/gemini/relay_responses.go b/relay/channel/gemini/relay_responses.go index b0535a7..6780ed4 100644 --- a/relay/channel/gemini/relay_responses.go +++ b/relay/channel/gemini/relay_responses.go @@ -89,7 +89,10 @@ func GeminiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, r streamErr = types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError) return false } - helper.ResponseChunkData(c, dto.ResponsesStreamResponse{Type: event.Type}, string(data)) + if err := helper.ResponseChunkData(c, dto.ResponsesStreamResponse{Type: event.Type}, string(data)); err != nil { + streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError) + return false + } return true } sendChunk := func(chunk *dto.ChatCompletionsStreamResponse) bool { diff --git a/relay/channel/openai/audio.go b/relay/channel/openai/audio.go index 6425e79..76bb136 100644 --- a/relay/channel/openai/audio.go +++ b/relay/channel/openai/audio.go @@ -103,8 +103,8 @@ func OpenaiTTSHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel usage.CompletionTokens = estimatedTokens usage.CompletionTokenDetails.AudioTokens = estimatedTokens } else if duration > 0 { - // 计算 token: ceil(duration) / 60.0 * 1000,即每分钟 1000 tokens - completionTokens := int(math.Round(math.Ceil(duration) / 60.0 * 1000)) + // 计算 token: ceil(duration) / 60.0 * 1000,即每分钟 1000 tokens。 + completionTokens := common.QuotaFromFloat(math.Round(math.Ceil(duration) / 60.0 * 1000)) usage.CompletionTokens = completionTokens usage.CompletionTokenDetails.AudioTokens = completionTokens } diff --git a/relay/channel/openai/helper.go b/relay/channel/openai/helper.go index c0cec80..2b89771 100644 --- a/relay/channel/openai/helper.go +++ b/relay/channel/openai/helper.go @@ -206,5 +206,5 @@ func sendResponsesStreamData(c *gin.Context, streamResponse dto.ResponsesStreamR if data == "" { return } - helper.ResponseChunkData(c, streamResponse, data) + _ = helper.ResponseChunkData(c, streamResponse, data) } diff --git a/relay/channel/openai/relay_image.go b/relay/channel/openai/relay_image.go index 36c4510..91a08da 100644 --- a/relay/channel/openai/relay_image.go +++ b/relay/channel/openai/relay_image.go @@ -131,10 +131,10 @@ func writeOpenaiImageStreamChunk(c *gin.Context, data []byte) { } _ = common.Unmarshal(data, &payload) if eventName := strings.TrimSpace(payload.Type); eventName != "" { - c.Render(-1, common.CustomEvent{Data: fmt.Sprintf("event: %s\n", eventName)}) + _ = helper.ResponseChunkData(c, dto.ResponsesStreamResponse{Type: eventName}, string(data)) + return } - c.Render(-1, common.CustomEvent{Data: "data: " + string(data)}) - _ = helper.FlushWriter(c) + _ = helper.StringData(c, string(data)) } func isOpenAIImageStreamErrorEvent(data []byte) bool { @@ -291,19 +291,11 @@ func writeOpenaiImageStreamPayload(c *gin.Context, eventName string, payload any return err } if eventName != "" { - if _, err := fmt.Fprintf(c.Writer, "event: %s\n", eventName); err != nil { - return err - } - } - if _, err := fmt.Fprintf(c.Writer, "data: %s\n\n", data); err != nil { - return err + return helper.ResponseChunkData(c, dto.ResponsesStreamResponse{Type: eventName}, string(data)) } - return helper.FlushWriter(c) + return helper.StringData(c, string(data)) } func writeOpenaiImageStreamDone(c *gin.Context) error { - if _, err := fmt.Fprint(c.Writer, "data: [DONE]\n\n"); err != nil { - return err - } - return helper.FlushWriter(c) + return helper.StringData(c, "[DONE]") } diff --git a/relay/channel/openai/responses_via_chat.go b/relay/channel/openai/responses_via_chat.go index 214b4fe..4b090f7 100644 --- a/relay/channel/openai/responses_via_chat.go +++ b/relay/channel/openai/responses_via_chat.go @@ -74,7 +74,10 @@ func OaiChatToResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo streamErr = types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError) return false } - helper.ResponseChunkData(c, dto.ResponsesStreamResponse{Type: event.Type}, string(data)) + if err := helper.ResponseChunkData(c, dto.ResponsesStreamResponse{Type: event.Type}, string(data)); err != nil { + streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError) + return false + } return true } diff --git a/relay/channel/task/ali/adaptor.go b/relay/channel/task/ali/adaptor.go index 90ea3d5..15438ca 100644 --- a/relay/channel/task/ali/adaptor.go +++ b/relay/channel/task/ali/adaptor.go @@ -510,19 +510,11 @@ func (a *TaskAdaptor) convertToAliRequest(info *relaycommon.RelayInfo, req relay } } - // 处理时长 - if duration := req.DurationValue(); duration > 0 { - aliReq.Parameters.Duration = duration - } else if req.Seconds != "" { - seconds, err := strconv.Atoi(req.Seconds) - if err != nil { - return nil, errors.Wrap(err, "convert seconds to int failed") - } else { - aliReq.Parameters.Duration = seconds - } - } else { - aliReq.Parameters.Duration = 5 // 默认5秒 + duration, err := req.ResolvedSecondsOrDefault(5) + if err != nil { + return nil, err } + aliReq.Parameters.Duration = duration // 从 metadata 中提取额外参数 if err := applyAliMetadata(req.Metadata, aliReq); err != nil { @@ -556,15 +548,11 @@ func (a *TaskAdaptor) convertToAliKlingRequest(upstreamModel string, req relayco }, } - if duration := req.DurationValue(); duration > 0 { - aliReq.Parameters.Duration = duration - } else if req.Seconds != "" { - seconds, err := strconv.Atoi(req.Seconds) - if err != nil { - return nil, errors.Wrap(err, "convert seconds to int failed") - } - aliReq.Parameters.Duration = seconds + duration, err := req.ResolvedSecondsOrDefault(5) + if err != nil { + return nil, err } + aliReq.Parameters.Duration = duration if imageURL := firstNonEmpty(req.InputReference, req.Image); imageURL != "" { aliReq.Input.Media = []map[string]interface{}{ { @@ -773,7 +761,7 @@ func (a *TaskAdaptor) EstimateBilling(c *gin.Context, info *relaycommon.RelayInf } otherRatios := map[string]float64{ - "seconds": float64(aliReq.Parameters.Duration), + "seconds": float64(min(aliReq.Parameters.Duration, relaycommon.MaxTaskDurationSeconds)), } ratios, err := ProcessAliOtherRatios(aliReq) if err != nil { diff --git a/relay/channel/task/ali/adaptor_kling_test.go b/relay/channel/task/ali/adaptor_kling_test.go index 1e0af60..64ceb33 100644 --- a/relay/channel/task/ali/adaptor_kling_test.go +++ b/relay/channel/task/ali/adaptor_kling_test.go @@ -87,6 +87,58 @@ func TestConvertToAliKlingRequestMapsKlingOfficialVideoList(t *testing.T) { } } +func TestConvertToAliKlingRequestResolvesDurationFromTaskSubmitReq(t *testing.T) { + adaptor := &TaskAdaptor{} + duration := 6 + + tests := []struct { + name string + req relaycommon.TaskSubmitReq + want int + }{ + { + name: "duration takes precedence over seconds", + req: relaycommon.TaskSubmitReq{ + Model: "kling/kling-v3-omni-video-generation", + Prompt: "make a video", + Duration: &duration, + Seconds: "9", + }, + want: 6, + }, + { + name: "seconds is used when duration is absent", + req: relaycommon.TaskSubmitReq{ + Model: "kling/kling-v3-omni-video-generation", + Prompt: "make a video", + Seconds: "9", + }, + want: 9, + }, + { + name: "zero duration defaults to five seconds", + req: relaycommon.TaskSubmitReq{ + Model: "kling/kling-v3-omni-video-generation", + Prompt: "make a video", + }, + want: 5, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := adaptor.convertToAliKlingRequest(tt.req.Model, tt.req) + + if err != nil { + t.Fatalf("convertToAliKlingRequest returned error: %v", err) + } + if got.Parameters == nil || got.Parameters.Duration != tt.want { + t.Fatalf("duration = %+v, want %d", got.Parameters, tt.want) + } + }) + } +} + func TestBuildAliKlingOfficialVideoResponseUsesPublicTaskID(t *testing.T) { resp := buildAliKlingOfficialVideoResponse( "task_public", diff --git a/relay/channel/task/ali/adaptor_wan27_test.go b/relay/channel/task/ali/adaptor_wan27_test.go index 87edbb9..a0c441c 100644 --- a/relay/channel/task/ali/adaptor_wan27_test.go +++ b/relay/channel/task/ali/adaptor_wan27_test.go @@ -44,6 +44,54 @@ func TestConvertToAliRequestWan27I2VBuildsMediaFromImage(t *testing.T) { assert.NotContains(t, string(body), `"img_url"`) } +func TestConvertToAliRequestResolvesDurationFromTaskSubmitReq(t *testing.T) { + adaptor := &TaskAdaptor{} + duration := 8 + + tests := []struct { + name string + req relaycommon.TaskSubmitReq + want int + }{ + { + name: "duration takes precedence over seconds", + req: relaycommon.TaskSubmitReq{ + Model: "wan2.5-t2v-preview", + Prompt: "make a video", + Duration: &duration, + Seconds: "12", + }, + want: 8, + }, + { + name: "seconds is used when duration is absent", + req: relaycommon.TaskSubmitReq{ + Model: "wan2.5-t2v-preview", + Prompt: "make a video", + Seconds: "12", + }, + want: 12, + }, + { + name: "zero duration defaults to five seconds", + req: relaycommon.TaskSubmitReq{ + Model: "wan2.5-t2v-preview", + Prompt: "make a video", + }, + want: 5, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + aliReq, err := adaptor.convertToAliRequest(testRelayInfo(), tt.req) + + require.NoError(t, err) + assert.Equal(t, tt.want, aliReq.Parameters.Duration) + }) + } +} + func TestConvertToAliRequestWan27I2VPrefersInputReferenceOverImage(t *testing.T) { adaptor := &TaskAdaptor{} req := relaycommon.TaskSubmitReq{ diff --git a/relay/channel/task/gemini/billing.go b/relay/channel/task/gemini/billing.go index f1d2a62..4be2f63 100644 --- a/relay/channel/task/gemini/billing.go +++ b/relay/channel/task/gemini/billing.go @@ -50,19 +50,20 @@ func ParseVeoResolution(metadata map[string]any) string { // ResolveVeoDuration returns the effective duration in seconds. // Priority: metadata["durationSeconds"] > stdDuration > stdSeconds > default (8). +// The result is capped because it is used as a billing multiplier. func ResolveVeoDuration(metadata map[string]any, stdDuration int, stdSeconds string) int { if metadata != nil { if _, exists := metadata["durationSeconds"]; exists { if d := ParseVeoDurationSeconds(metadata); d > 0 { - return d + return min(d, relaycommon.MaxTaskDurationSeconds) } } } if stdDuration > 0 { - return stdDuration + return min(stdDuration, relaycommon.MaxTaskDurationSeconds) } if s, err := strconv.Atoi(stdSeconds); err == nil && s > 0 { - return s + return min(s, relaycommon.MaxTaskDurationSeconds) } return 8 } diff --git a/relay/channel/task/kling/adaptor.go b/relay/channel/task/kling/adaptor.go index 50cc117..1a9a8dd 100644 --- a/relay/channel/task/kling/adaptor.go +++ b/relay/channel/task/kling/adaptor.go @@ -502,7 +502,7 @@ func (a *TaskAdaptor) ParseTaskResult(respBody []byte) (*relaycommon.TaskInfo, e taskInfo.Url = video.Url } if tokens, err := strconv.ParseFloat(resPayload.Data.FinalUnitDeduction, 64); err == nil { - rounded := int(math.Ceil(tokens)) + rounded := common.QuotaFromFloat(math.Ceil(tokens)) if rounded > 0 { taskInfo.CompletionTokens = rounded taskInfo.TotalTokens = rounded diff --git a/relay/common/relay_info.go b/relay/common/relay_info.go index fbe6044..ff4aa7c 100644 --- a/relay/common/relay_info.go +++ b/relay/common/relay_info.go @@ -733,6 +733,32 @@ func (t *TaskSubmitReq) DurationValue() int { return *t.Duration } +func (t *TaskSubmitReq) ResolvedSeconds() (int, error) { + if t == nil { + return 0, nil + } + seconds := t.DurationValue() + if seconds == 0 && t.Seconds != "" { + parsed, err := strconv.Atoi(t.Seconds) + if err != nil { + return 0, fmt.Errorf("invalid seconds value: %s", t.Seconds) + } + seconds = parsed + } + return seconds, nil +} + +func (t *TaskSubmitReq) ResolvedSecondsOrDefault(defaultSeconds int) (int, error) { + seconds, err := t.ResolvedSeconds() + if err != nil { + return 0, err + } + if seconds <= 0 { + return defaultSeconds, nil + } + return seconds, nil +} + func (t *TaskSubmitReq) HasImage() bool { return len(t.Images) > 0 } diff --git a/relay/common/relay_info_test.go b/relay/common/relay_info_test.go index 3788b3b..b80c152 100644 --- a/relay/common/relay_info_test.go +++ b/relay/common/relay_info_test.go @@ -51,3 +51,34 @@ func TestTaskSubmitReqPreservesExplicitZeroDuration(t *testing.T) { require.NoError(t, err) require.Contains(t, string(data), `"duration":0`) } + +func TestTaskSubmitReqResolvedSeconds(t *testing.T) { + duration := 8 + req := TaskSubmitReq{ + Duration: &duration, + Seconds: "12", + } + + seconds, err := req.ResolvedSeconds() + require.NoError(t, err) + require.Equal(t, 8, seconds) + + req = TaskSubmitReq{Seconds: "12"} + seconds, err = req.ResolvedSeconds() + require.NoError(t, err) + require.Equal(t, 12, seconds) + + req = TaskSubmitReq{Seconds: "abc"} + _, err = req.ResolvedSeconds() + require.Error(t, err) + require.Contains(t, err.Error(), "invalid seconds value: abc") +} + +func TestTaskSubmitReqResolvedSecondsOrDefault(t *testing.T) { + zero := 0 + req := TaskSubmitReq{Duration: &zero} + + seconds, err := req.ResolvedSecondsOrDefault(5) + require.NoError(t, err) + require.Equal(t, 5, seconds) +} diff --git a/relay/common/relay_utils.go b/relay/common/relay_utils.go index ec19987..1cd1cd0 100644 --- a/relay/common/relay_utils.go +++ b/relay/common/relay_utils.go @@ -112,6 +112,21 @@ func validatePrompt(prompt string) *dto.TaskError { return nil } +// MaxTaskDurationSeconds caps user-supplied video duration before it becomes +// an OtherRatio billing multiplier. +const MaxTaskDurationSeconds = 3600 + +func validateTaskDurationBounds(req TaskSubmitReq) *dto.TaskError { + seconds, err := req.ResolvedSeconds() + if err != nil { + return createTaskError(err, "invalid_seconds", http.StatusBadRequest, true) + } + if seconds < 0 || seconds > MaxTaskDurationSeconds { + return createTaskError(fmt.Errorf("seconds must be between 0 and %d", MaxTaskDurationSeconds), "invalid_seconds", http.StatusBadRequest, true) + } + return nil +} + func validateMultipartTaskRequest(c *gin.Context, info *RelayInfo, action string) (TaskSubmitReq, error) { var req TaskSubmitReq if _, err := c.MultipartForm(); err != nil { @@ -168,10 +183,11 @@ func ValidateMultipartDirect(c *gin.Context, info *RelayInfo) *dto.TaskError { prompt = req.Prompt model = req.Model size = req.Size - seconds, _ = strconv.Atoi(req.Seconds) - if seconds == 0 { - seconds = req.DurationValue() + resolvedSeconds, err := req.ResolvedSeconds() + if err != nil { + return createTaskError(err, "invalid_seconds", http.StatusBadRequest, true) } + seconds = resolvedSeconds if inputReference := strings.TrimSpace(req.InputReference); inputReference != "" { req.Images = []string{inputReference} } else if len(req.Images) == 0 && strings.TrimSpace(req.Image) != "" { @@ -190,6 +206,10 @@ func ValidateMultipartDirect(c *gin.Context, info *RelayInfo) *dto.TaskError { return taskErr } + if taskErr := validateTaskDurationBounds(req); taskErr != nil { + return taskErr + } + action := constant.TaskActionTextGenerate if hasInputReference { action = constant.TaskActionGenerate @@ -251,6 +271,10 @@ func ValidateBasicTaskRequest(c *gin.Context, info *RelayInfo, action string) *d return taskErr } + if taskErr := validateTaskDurationBounds(req); taskErr != nil { + return taskErr + } + if len(req.Images) == 0 && strings.TrimSpace(req.Image) != "" { // 兼容单图上传 req.Images = []string{strings.TrimSpace(req.Image)} diff --git a/relay/common/relay_utils_test.go b/relay/common/relay_utils_test.go index 863a33e..7db18de 100644 --- a/relay/common/relay_utils_test.go +++ b/relay/common/relay_utils_test.go @@ -53,6 +53,80 @@ func TestValidateMultipartDirectNormalizesInputReferenceField(t *testing.T) { require.Equal(t, constant.TaskActionGenerate, info.Action) } +func TestTaskDurationBounds(t *testing.T) { + gin.SetMode(gin.TestMode) + + newContext := func(body string) (*gin.Context, *RelayInfo) { + request := httptest.NewRequest(http.MethodPost, "/v1/video/generations", strings.NewReader(body)) + request.Header.Set("Content-Type", "application/json") + context, _ := gin.CreateTestContext(httptest.NewRecorder()) + context.Request = request + return context, &RelayInfo{TaskRelayInfo: &TaskRelayInfo{}} + } + + tests := []struct { + name string + body string + wantErr bool + wantMessage string + }{ + { + name: "huge duration is rejected", + body: `{"model":"sora-2","prompt":"a cat","duration":9999999999}`, + wantErr: true, + wantMessage: "seconds must be between 0 and", + }, + { + name: "huge seconds string is rejected", + body: `{"model":"sora-2","prompt":"a cat","seconds":"9999999999"}`, + wantErr: true, + wantMessage: "seconds must be between 0 and", + }, + { + name: "negative duration is rejected", + body: `{"model":"sora-2","prompt":"a cat","duration":-8}`, + wantErr: true, + wantMessage: "seconds must be between 0 and", + }, + { + name: "non numeric seconds string is rejected", + body: `{"model":"sora-2","prompt":"a cat","seconds":"abc"}`, + wantErr: true, + wantMessage: "invalid seconds value: abc", + }, + { + name: "normal duration is accepted", + body: `{"model":"sora-2","prompt":"a cat","seconds":"8"}`, + }, + } + + for _, tt := range tests { + t.Run(tt.name+" multipart direct", func(t *testing.T) { + context, info := newContext(tt.body) + taskErr := ValidateMultipartDirect(context, info) + if tt.wantErr { + require.NotNil(t, taskErr) + require.Equal(t, "invalid_seconds", taskErr.Code) + require.Contains(t, taskErr.Message, tt.wantMessage) + return + } + require.Nil(t, taskErr) + }) + + t.Run(tt.name+" basic task request", func(t *testing.T) { + context, info := newContext(tt.body) + taskErr := ValidateBasicTaskRequest(context, info, constant.TaskActionGenerate) + if tt.wantErr { + require.NotNil(t, taskErr) + require.Equal(t, "invalid_seconds", taskErr.Code) + require.Contains(t, taskErr.Message, tt.wantMessage) + return + } + require.Nil(t, taskErr) + }) + } +} + func TestValidateMultipartDirectIgnoresBlankInputReference(t *testing.T) { gin.SetMode(gin.TestMode) body := strings.NewReader(`{"model":"wan2.7-i2v","prompt":"animate","input_reference":" ","image":" https://example.com/first.png "}`) diff --git a/relay/helper/common.go b/relay/helper/common.go index 623df22..72fb876 100644 --- a/relay/helper/common.go +++ b/relay/helper/common.go @@ -25,7 +25,7 @@ func FlushWriter(c *gin.Context) (err error) { return nil } - if c.Request != nil && c.Request.Context().Err() != nil { + if requestContextDone(c) { return fmt.Errorf("request context done: %w", c.Request.Context().Err()) } @@ -38,6 +38,10 @@ func FlushWriter(c *gin.Context) (err error) { return nil } +func requestContextDone(c *gin.Context) bool { + return c != nil && c.Request != nil && c.Request.Context().Err() != nil +} + func SetEventStreamHeaders(c *gin.Context) { // 检查是否已经设置过头部 if _, exists := c.Get("event_stream_headers_set"); exists { @@ -55,6 +59,14 @@ func SetEventStreamHeaders(c *gin.Context) { } func ClaudeData(c *gin.Context, resp dto.ClaudeResponse) error { + if c == nil || c.Writer == nil { + return errors.New("context or writer is nil") + } + + if requestContextDone(c) { + return nil + } + jsonData, err := common.Marshal(resp) if err != nil { common.SysError("error marshalling stream response: " + err.Error()) @@ -67,15 +79,31 @@ func ClaudeData(c *gin.Context, resp dto.ClaudeResponse) error { } func ClaudeChunkData(c *gin.Context, resp dto.ClaudeResponse, data string) { + if c == nil || c.Writer == nil { + return + } + + if requestContextDone(c) { + return + } + c.Render(-1, common.CustomEvent{Data: fmt.Sprintf("event: %s\n", resp.Type)}) c.Render(-1, common.CustomEvent{Data: fmt.Sprintf("data: %s\n", data)}) _ = FlushWriter(c) } -func ResponseChunkData(c *gin.Context, resp dto.ResponsesStreamResponse, data string) { +func ResponseChunkData(c *gin.Context, resp dto.ResponsesStreamResponse, data string) error { + if c == nil || c.Writer == nil { + return errors.New("context or writer is nil") + } + + if requestContextDone(c) { + return fmt.Errorf("request context done: %w", c.Request.Context().Err()) + } + c.Render(-1, common.CustomEvent{Data: fmt.Sprintf("event: %s\n", resp.Type)}) c.Render(-1, common.CustomEvent{Data: fmt.Sprintf("data: %s", data)}) - _ = FlushWriter(c) + return FlushWriter(c) } func StringData(c *gin.Context, str string) error { @@ -83,7 +111,7 @@ func StringData(c *gin.Context, str string) error { return errors.New("context or writer is nil") } - if c.Request != nil && c.Request.Context().Err() != nil { + if requestContextDone(c) { return fmt.Errorf("request context done: %w", c.Request.Context().Err()) } @@ -96,7 +124,7 @@ func PingData(c *gin.Context) error { return errors.New("context or writer is nil") } - if c.Request != nil && c.Request.Context().Err() != nil { + if requestContextDone(c) { return fmt.Errorf("request context done: %w", c.Request.Context().Err()) } diff --git a/relay/helper/common_test.go b/relay/helper/common_test.go new file mode 100644 index 0000000..c099a38 --- /dev/null +++ b/relay/helper/common_test.go @@ -0,0 +1,22 @@ +package helper + +import ( + "testing" + + "github.com/MAX-API-Next/MAX-API/dto" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestStreamDataHelpersHandleNilContextOrWriter(t *testing.T) { + claudeResp := dto.ClaudeResponse{Type: "message_delta"} + responsesResp := dto.ResponsesStreamResponse{Type: "response.output_text.delta"} + + for _, c := range []*gin.Context{nil, &gin.Context{}} { + require.ErrorContains(t, ClaudeData(c, claudeResp), "context or writer is nil") + require.NotPanics(t, func() { + ClaudeChunkData(c, claudeResp, "{}") + }) + require.ErrorContains(t, ResponseChunkData(c, responsesResp, "{}"), "context or writer is nil") + } +} diff --git a/relay/helper/max_tokens_bounds_test.go b/relay/helper/max_tokens_bounds_test.go new file mode 100644 index 0000000..dcbab5f --- /dev/null +++ b/relay/helper/max_tokens_bounds_test.go @@ -0,0 +1,67 @@ +package helper + +import ( + "bytes" + "net/http" + "net/http/httptest" + "testing" + + relayconstant "github.com/MAX-API-Next/MAX-API/relay/constant" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestMaxTokensBounds(t *testing.T) { + gin.SetMode(gin.TestMode) + + newJSONContext := func(body string) *gin.Context { + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest(http.MethodPost, "/relay", bytes.NewBufferString(body)) + c.Request.Header.Set("Content-Type", "application/json") + return c + } + + const tooLarge = "1073741824" + + t.Run("openai max_tokens rejected", func(t *testing.T) { + c := newJSONContext(`{"model":"gpt-4o","messages":[{"role":"user","content":"hi"}],"max_tokens":` + tooLarge + `}`) + _, err := GetAndValidateTextRequest(c, relayconstant.RelayModeChatCompletions) + require.Error(t, err) + require.Contains(t, err.Error(), "max_tokens is invalid") + }) + + t.Run("openai max_completion_tokens rejected", func(t *testing.T) { + c := newJSONContext(`{"model":"gpt-4o","messages":[{"role":"user","content":"hi"}],"max_completion_tokens":` + tooLarge + `}`) + _, err := GetAndValidateTextRequest(c, relayconstant.RelayModeChatCompletions) + require.Error(t, err) + require.Contains(t, err.Error(), "max_tokens is invalid") + }) + + t.Run("claude max_tokens rejected", func(t *testing.T) { + c := newJSONContext(`{"model":"claude-sonnet-4","messages":[{"role":"user","content":"hi"}],"max_tokens":` + tooLarge + `}`) + _, err := GetAndValidateClaudeRequest(c) + require.Error(t, err) + require.Contains(t, err.Error(), "max_tokens is invalid") + }) + + t.Run("claude normal max_tokens accepted", func(t *testing.T) { + c := newJSONContext(`{"model":"claude-sonnet-4","messages":[{"role":"user","content":"hi"}],"max_tokens":8192}`) + req, err := GetAndValidateClaudeRequest(c) + require.NoError(t, err) + require.EqualValues(t, 8192, *req.MaxTokens) + }) + + t.Run("gemini maxOutputTokens rejected", func(t *testing.T) { + c := newJSONContext(`{"contents":[{"parts":[{"text":"hi"}]}],"generationConfig":{"maxOutputTokens":` + tooLarge + `}}`) + _, err := GetAndValidateGeminiRequest(c) + require.Error(t, err) + require.Contains(t, err.Error(), "maxOutputTokens is invalid") + }) + + t.Run("responses max_output_tokens rejected", func(t *testing.T) { + c := newJSONContext(`{"model":"gpt-4o","input":"hi","max_output_tokens":` + tooLarge + `}`) + _, err := GetAndValidateResponsesRequest(c) + require.Error(t, err) + require.Contains(t, err.Error(), "max_output_tokens is invalid") + }) +} diff --git a/relay/helper/openai_image_request_test.go b/relay/helper/openai_image_request_test.go new file mode 100644 index 0000000..f88a48e --- /dev/null +++ b/relay/helper/openai_image_request_test.go @@ -0,0 +1,82 @@ +package helper + +import ( + "bytes" + "fmt" + "mime/multipart" + "net/http" + "net/http/httptest" + "testing" + + "github.com/MAX-API-Next/MAX-API/dto" + relayconstant "github.com/MAX-API-Next/MAX-API/relay/constant" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestGetAndValidOpenAIImageRequestNBounds(t *testing.T) { + gin.SetMode(gin.TestMode) + + newJSONContext := func(body string) *gin.Context { + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", bytes.NewBufferString(body)) + c.Request.Header.Set("Content-Type", "application/json") + return c + } + + boundErr := fmt.Sprintf("n must be an integer between 0 and %d", dto.MaxImageN) + + tests := []struct { + name string + body string + wantErr string + wantN uint + }{ + { + name: "n above max is rejected", + body: fmt.Sprintf(`{"model":"gpt-image-1","prompt":"a cat","n":%d}`, dto.MaxImageN+1), + wantErr: boundErr, + }, + { + name: "n at max is accepted", + body: fmt.Sprintf(`{"model":"gpt-image-1","prompt":"a cat","n":%d}`, dto.MaxImageN), + wantN: dto.MaxImageN, + }, + { + name: "absent n defaults to 1", + body: `{"model":"gpt-image-1","prompt":"a cat"}`, + wantN: 1, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + req, err := GetAndValidOpenAIImageRequest(newJSONContext(tt.body), relayconstant.RelayModeImagesGenerations) + if tt.wantErr != "" { + require.Error(t, err) + require.Contains(t, err.Error(), tt.wantErr) + return + } + require.NoError(t, err) + require.NotNil(t, req.N) + require.Equal(t, tt.wantN, *req.N) + }) + } + + t.Run("negative multipart n is rejected", func(t *testing.T) { + var body bytes.Buffer + writer := multipart.NewWriter(&body) + require.NoError(t, writer.WriteField("model", "gpt-image-1")) + require.NoError(t, writer.WriteField("prompt", "edit this image")) + require.NoError(t, writer.WriteField("n", "-22904832")) + require.NoError(t, writer.Close()) + + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/edits", &body) + c.Request.Header.Set("Content-Type", writer.FormDataContentType()) + + _, err := GetAndValidOpenAIImageRequest(c, relayconstant.RelayModeImagesEdits) + require.Error(t, err) + require.Contains(t, err.Error(), boundErr) + }) +} diff --git a/relay/helper/price.go b/relay/helper/price.go index 691257d..10585bf 100644 --- a/relay/helper/price.go +++ b/relay/helper/price.go @@ -113,12 +113,12 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens audioRatio = ratio_setting.GetAudioRatio(info.OriginModelName) audioCompletionRatio = ratio_setting.GetAudioCompletionRatio(info.OriginModelName) ratio := modelRatio * groupRatioInfo.GroupRatio - preConsumedQuota = int(float64(preConsumedTokens) * ratio) + preConsumedQuota = common.QuotaFromFloat(float64(preConsumedTokens) * ratio) } else { if meta.ImagePriceRatio != 0 { modelPrice = modelPrice * meta.ImagePriceRatio } - preConsumedQuota = int(modelPrice * common.QuotaPerUnit * groupRatioInfo.GroupRatio) + preConsumedQuota = common.QuotaFromFloat(modelPrice * common.QuotaPerUnit * groupRatioInfo.GroupRatio) } // check if free model pre-consume is disabled @@ -199,7 +199,7 @@ func ModelPriceHelperPerCall(c *gin.Context, info *relaycommon.RelayInfo) (types freeModel := false if usePrice { - quota = int(modelPrice * common.QuotaPerUnit * groupRatioInfo.GroupRatio) + quota = common.QuotaFromFloat(modelPrice * common.QuotaPerUnit * groupRatioInfo.GroupRatio) if !operation_setting.GetQuotaSetting().EnableFreeModelPreConsume { if groupRatioInfo.GroupRatio == 0 || (modelPrice == 0 && !rateCardPriced) { quota = 0 @@ -208,7 +208,7 @@ func ModelPriceHelperPerCall(c *gin.Context, info *relaycommon.RelayInfo) (types } } else { // 按量计费:以模型倍率的一半作为预扣额度 - quota = int(modelRatio / 2 * common.QuotaPerUnit * groupRatioInfo.GroupRatio) + quota = common.QuotaFromFloat(modelRatio / 2 * common.QuotaPerUnit * groupRatioInfo.GroupRatio) modelPrice = -1 if !operation_setting.GetQuotaSetting().EnableFreeModelPreConsume { if groupRatioInfo.GroupRatio == 0 || modelRatio == 0 { diff --git a/relay/helper/stream_scanner.go b/relay/helper/stream_scanner.go index 3c79eac..0a8fc05 100644 --- a/relay/helper/stream_scanner.go +++ b/relay/helper/stream_scanner.go @@ -25,6 +25,7 @@ const ( InitialScannerBufferSize = 64 << 10 // 64KB (64*1024) DefaultMaxScannerBufferSize = 64 << 20 // 64MB (64*1024*1024) default SSE buffer size DefaultPingInterval = 10 * time.Second + streamWriteTimeout = 30 * time.Second ) func getScannerBufferSize() int { @@ -40,6 +41,13 @@ func NewStreamScanner(reader io.Reader) *bufio.Scanner { return scanner } +func ExtendWriteDeadline(c *gin.Context) { + if c == nil || c.Writer == nil { + return + } + _ = http.NewResponseController(c.Writer).SetWriteDeadline(time.Now().Add(streamWriteTimeout)) +} + func newStreamingTimeoutTicker(timeoutSeconds int) (*time.Ticker, <-chan time.Time) { if timeoutSeconds <= 0 { return nil, nil @@ -59,25 +67,28 @@ func StreamScannerHandler(c *gin.Context, resp *http.Response, info *relaycommon info.StreamStatus = relaycommon.NewStreamStatus() } - // 确保响应体总是被关闭 - defer func() { - if resp.Body != nil { - resp.Body.Close() - } - }() + ctx, cancel := context.WithCancel(context.Background()) streamingTimeoutSeconds := constant.StreamingTimeout streamingTimeout := time.Duration(streamingTimeoutSeconds) * time.Second ticker, timeoutChan := newStreamingTimeoutTicker(streamingTimeoutSeconds) var ( - stopChan = make(chan bool, 3) // 增加缓冲区避免阻塞 - scanner = NewStreamScanner(resp.Body) - pingTicker *time.Ticker - writeMutex sync.Mutex // Mutex to protect concurrent writes - wg sync.WaitGroup // 用于等待所有 goroutine 退出 + stopChan = make(chan bool, 3) // 增加缓冲区避免阻塞 + scanner = NewStreamScanner(resp.Body) + pingTicker *time.Ticker + writeMutex sync.Mutex // Mutex to protect concurrent writes + wg sync.WaitGroup // 用于等待所有 goroutine 退出 + cleanupOnce sync.Once + stopOnce sync.Once ) + stop := func() { + stopOnce.Do(func() { + close(stopChan) + }) + } + generalSettings := operation_setting.GetGeneralSetting() pingEnabled := generalSettings.PingIntervalEnabled && !info.DisablePing pingInterval := time.Duration(generalSettings.PingIntervalSeconds) * time.Second @@ -95,40 +106,27 @@ func StreamScannerHandler(c *gin.Context, resp *http.Response, info *relaycommon logger.LogDebug(c, "streaming timeout seconds: %d", int64(streamingTimeout.Seconds())) logger.LogDebug(c, "ping interval seconds: %d", int64(pingInterval.Seconds())) - // 改进资源清理,确保所有 goroutine 正确退出 - defer func() { - // 通知所有 goroutine 停止 - common.SafeSendBool(stopChan, true) - - if ticker != nil { - ticker.Stop() - } - if pingTicker != nil { - pingTicker.Stop() - } - - // 等待所有 goroutine 退出,最多等待5秒 - done := make(chan struct{}) - gopool.Go(func() { + cleanup := func() { + cleanupOnce.Do(func() { + cancel() + stop() + if resp.Body != nil { + _ = resp.Body.Close() + } + if ticker != nil { + ticker.Stop() + } + if pingTicker != nil { + pingTicker.Stop() + } wg.Wait() - close(done) }) - - select { - case <-done: - case <-time.After(5 * time.Second): - logger.LogError(c, "timeout waiting for goroutines to exit") - } - - close(stopChan) - }() + } + defer cleanup() scanner.Split(bufio.ScanLines) SetEventStreamHeaders(c) - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - ctx = context.WithValue(ctx, "stop_chan", stopChan) // Handle ping data sending with improved error handling @@ -136,13 +134,13 @@ func StreamScannerHandler(c *gin.Context, resp *http.Response, info *relaycommon wg.Add(1) gopool.Go(func() { defer func() { - wg.Done() if r := recover(); r != nil { logger.LogError(c, fmt.Sprintf("ping goroutine panic: %v", r)) info.StreamStatus.SetEndReason(relaycommon.StreamEndReasonPanic, fmt.Errorf("ping panic: %v", r)) - common.SafeSendBool(stopChan, true) + stop() } logger.LogDebug(c, "ping goroutine exited") + wg.Done() }() // 添加超时保护,防止 goroutine 无限运行 @@ -153,31 +151,19 @@ func StreamScannerHandler(c *gin.Context, resp *http.Response, info *relaycommon for { select { case <-pingTicker.C: - // 使用超时机制防止写操作阻塞 - done := make(chan error, 1) - gopool.Go(func() { + var err error + func() { writeMutex.Lock() defer writeMutex.Unlock() - done <- PingData(c) - }) - - select { - case err := <-done: - if err != nil { - logger.LogError(c, "ping data error: "+err.Error()) - info.StreamStatus.SetEndReason(relaycommon.StreamEndReasonPingFail, err) - return - } - logger.LogDebug(c, "ping data sent") - case <-time.After(10 * time.Second): - logger.LogError(c, "ping data send timeout") - info.StreamStatus.SetEndReason(relaycommon.StreamEndReasonPingFail, fmt.Errorf("ping send timeout")) - return - case <-ctx.Done(): - return - case <-stopChan: + ExtendWriteDeadline(c) + err = PingData(c) + }() + if err != nil { + logger.LogError(c, "ping data error: "+err.Error()) + info.StreamStatus.SetEndReason(relaycommon.StreamEndReasonPingFail, err) return } + logger.LogDebug(c, "ping data sent") case <-ctx.Done(): return case <-stopChan: @@ -198,19 +184,22 @@ func StreamScannerHandler(c *gin.Context, resp *http.Response, info *relaycommon wg.Add(1) gopool.Go(func() { defer func() { - wg.Done() if r := recover(); r != nil { logger.LogError(c, fmt.Sprintf("data handler goroutine panic: %v", r)) info.StreamStatus.SetEndReason(relaycommon.StreamEndReasonPanic, fmt.Errorf("handler panic: %v", r)) } - common.SafeSendBool(stopChan, true) + stop() + wg.Done() }() sr := newStreamResult(info.StreamStatus) for data := range dataChan { sr.reset() - writeMutex.Lock() - dataHandler(data, sr) - writeMutex.Unlock() + func() { + writeMutex.Lock() + defer writeMutex.Unlock() + ExtendWriteDeadline(c) + dataHandler(data, sr) + }() if sr.IsStopped() { return } @@ -222,13 +211,13 @@ func StreamScannerHandler(c *gin.Context, resp *http.Response, info *relaycommon common.RelayCtxGo(ctx, func() { defer func() { close(dataChan) - wg.Done() if r := recover(); r != nil { logger.LogError(c, fmt.Sprintf("scanner goroutine panic: %v", r)) info.StreamStatus.SetEndReason(relaycommon.StreamEndReasonPanic, fmt.Errorf("scanner panic: %v", r)) } - common.SafeSendBool(stopChan, true) + stop() logger.LogDebug(c, "scanner goroutine exited") + wg.Done() }() for scanner.Scan() { @@ -298,6 +287,7 @@ func StreamScannerHandler(c *gin.Context, resp *http.Response, info *relaycommon info.StreamStatus.SetEndReason(relaycommon.StreamEndReasonClientGone, c.Request.Context().Err()) } + cleanup() if info.StreamStatus.IsNormalEnd() && !info.StreamStatus.HasErrors() { logger.LogInfo(c, fmt.Sprintf("stream ended: %s", info.StreamStatus.Summary())) } else { diff --git a/relay/helper/stream_scanner_test.go b/relay/helper/stream_scanner_test.go index 1ecd645..bff7ae0 100644 --- a/relay/helper/stream_scanner_test.go +++ b/relay/helper/stream_scanner_test.go @@ -2,6 +2,7 @@ package helper import ( "bufio" + "context" "fmt" "io" "net/http" @@ -20,6 +21,19 @@ import ( "github.com/stretchr/testify/require" ) +type notifyPipeReadCloser struct { + *io.PipeReader + closed chan struct{} + once sync.Once +} + +func (r *notifyPipeReadCloser) Close() error { + r.once.Do(func() { + close(r.closed) + }) + return r.PipeReader.Close() +} + func init() { gin.SetMode(gin.TestMode) } @@ -42,6 +56,44 @@ func setupStreamTest(t *testing.T, body io.Reader) (*gin.Context, *http.Response return c, resp, info } +func TestStreamScannerHandler_ClientCancelClosesUpstreamBody(t *testing.T) { + pr, pw := io.Pipe() + t.Cleanup(func() { + _ = pw.Close() + }) + + body := ¬ifyPipeReadCloser{ + PipeReader: pr, + closed: make(chan struct{}), + } + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + ctx, cancel := context.WithCancel(context.Background()) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil).WithContext(ctx) + resp := &http.Response{Body: body} + info := &relaycommon.RelayInfo{ChannelMeta: &relaycommon.ChannelMeta{}} + + done := make(chan struct{}) + go func() { + defer close(done) + StreamScannerHandler(c, resp, info, func(data string, sr *StreamResult) {}) + }() + + cancel() + + select { + case <-body.closed: + case <-time.After(time.Second): + t.Fatal("upstream response body was not closed after client cancel") + } + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("StreamScannerHandler did not return after client cancel") + } + assert.Equal(t, relaycommon.StreamEndReasonClientGone, info.StreamStatus.EndReason) +} + func buildSSEBody(n int) string { var b strings.Builder for i := 0; i < n; i++ { diff --git a/relay/helper/valid_request.go b/relay/helper/valid_request.go index 6d07849..f0da5d4 100644 --- a/relay/helper/valid_request.go +++ b/relay/helper/valid_request.go @@ -4,6 +4,7 @@ import ( "errors" "fmt" "math" + "strconv" "strings" "github.com/MAX-API-Next/MAX-API/common" @@ -112,6 +113,17 @@ func GetAndValidateEmbeddingRequest(c *gin.Context, relayMode int) (*dto.Embeddi return embeddingRequest, nil } +const maxTokensLimit = math.MaxInt32 / 2 + +func exceedsMaxTokensLimit(values ...*uint) bool { + for _, v := range values { + if lo.FromPtrOr(v, uint(0)) > maxTokensLimit { + return true + } + } + return false +} + func GetAndValidateResponsesRequest(c *gin.Context) (*dto.OpenAIResponsesRequest, error) { request := &dto.OpenAIResponsesRequest{} err := common.UnmarshalBodyReusable(c, request) @@ -124,6 +136,9 @@ func GetAndValidateResponsesRequest(c *gin.Context) (*dto.OpenAIResponsesRequest if request.Input == nil { return nil, errors.New("input is required") } + if exceedsMaxTokensLimit(request.MaxOutputTokens) { + return nil, errors.New("max_output_tokens is invalid") + } return request, nil } @@ -151,7 +166,16 @@ func GetAndValidOpenAIImageRequest(c *gin.Context, relayMode int) (*dto.ImageReq formData := c.Request.PostForm imageRequest.Prompt = formData.Get("prompt") imageRequest.Model = formData.Get("model") - imageRequest.N = common.GetPointer(uint(common.String2Int(formData.Get("n")))) + if nValue := strings.TrimSpace(formData.Get("n")); nValue != "" { + n, err := strconv.Atoi(nValue) + if err != nil { + return nil, dto.ValidateImageN("n", -1) + } + if err := dto.ValidateImageN("n", n); err != nil { + return nil, err + } + imageRequest.N = common.GetPointer(uint(n)) + } imageRequest.Quality = formData.Get("quality") imageRequest.Size = formData.Get("size") if imageValue := formData.Get("image"); imageValue != "" { @@ -190,6 +214,12 @@ func GetAndValidOpenAIImageRequest(c *gin.Context, relayMode int) (*dto.ImageReq return nil, errors.New("size an unexpected error occurred in the parameter, please use 'x' instead of the multiplication sign '×'") } + if imageRequest.N != nil { + if err := dto.ValidateImageN("n", int(*imageRequest.N)); err != nil { + return nil, err + } + } + // Not "256x256", "512x512", or "1024x1024" if imageRequest.Model == "dall-e-2" || imageRequest.Model == "dall-e" { if imageRequest.Size != "" && imageRequest.Size != "256x256" && imageRequest.Size != "512x512" && imageRequest.Size != "1024x1024" { @@ -238,6 +268,9 @@ func GetAndValidateClaudeRequest(c *gin.Context) (textRequest *dto.ClaudeRequest if textRequest.Model == "" { return nil, errors.New("field model is required") } + if exceedsMaxTokensLimit(textRequest.MaxTokens, textRequest.MaxTokensToSample) { + return nil, errors.New("max_tokens is invalid") + } //if textRequest.Stream { // relayInfo.IsStream = true @@ -260,7 +293,7 @@ func GetAndValidateTextRequest(c *gin.Context, relayMode int) (*dto.GeneralOpenA textRequest.Model = c.Param("model") } - if lo.FromPtrOr(textRequest.MaxTokens, uint(0)) > math.MaxInt32/2 { + if exceedsMaxTokensLimit(textRequest.MaxTokens, textRequest.MaxCompletionTokens) { return nil, errors.New("max_tokens is invalid") } if textRequest.Model == "" { @@ -313,6 +346,9 @@ func GetAndValidateGeminiRequest(c *gin.Context) (*dto.GeminiChatRequest, error) if len(request.Contents) == 0 && len(request.Requests) == 0 { return nil, errors.New("contents is required") } + if exceedsMaxTokensLimit(request.GenerationConfig.MaxOutputTokens) { + return nil, errors.New("maxOutputTokens is invalid") + } //if c.Query("alt") == "sse" { // relayInfo.IsStream = true diff --git a/relay/relay_task.go b/relay/relay_task.go index b046ef6..6e892b5 100644 --- a/relay/relay_task.go +++ b/relay/relay_task.go @@ -126,14 +126,14 @@ func ResolveOriginTask(c *gin.Context, info *relaycommon.RelayInfo) *dto.TaskErr if seconds <= 0 { seconds = 4 } - sizeStr, _ := taskData["size"].(string) - if info.PriceData.OtherRatios == nil { - info.PriceData.OtherRatios = map[string]float64{} + if seconds > relaycommon.MaxTaskDurationSeconds { + seconds = relaycommon.MaxTaskDurationSeconds } - info.PriceData.OtherRatios["seconds"] = float64(seconds) - info.PriceData.OtherRatios["size"] = 1 + sizeStr, _ := taskData["size"].(string) + info.PriceData.AddOtherRatio("seconds", float64(seconds)) + info.PriceData.AddOtherRatio("size", 1) if sizeStr == "1792x1024" || sizeStr == "1024x1792" { - info.PriceData.OtherRatios["size"] = 1.666667 + info.PriceData.AddOtherRatio("size", 1.666667) } } } @@ -216,11 +216,13 @@ func RelayTaskSubmit(c *gin.Context, info *relaycommon.RelayInfo) (*TaskSubmitRe } if !common.StringsContains(constant.TaskPricePatches, modelName) { + quotaWithRatios := float64(info.PriceData.Quota) for _, ra := range info.PriceData.OtherRatios { if ra != 1.0 { - info.PriceData.Quota = int(float64(info.PriceData.Quota) * ra) + quotaWithRatios *= ra } } + info.PriceData.Quota = common.QuotaFromFloat(quotaWithRatios) } } @@ -357,21 +359,21 @@ func buildTaskSubmitRequestBody(c *gin.Context, info *relaycommon.RelayInfo, ada // 公式: baseQuota × ∏(ratio) — 其中 baseQuota 是不含 OtherRatios 的基础额度。 func recalcQuotaFromRatios(info *relaycommon.RelayInfo, ratios map[string]float64) int { // 从 PriceData 获取不含 OtherRatios 的基础价格 - baseQuota := info.PriceData.Quota + baseQuota := float64(info.PriceData.Quota) // 先除掉原有的 OtherRatios 恢复基础额度 for _, ra := range info.PriceData.OtherRatios { if ra != 1.0 && ra > 0 { - baseQuota = int(float64(baseQuota) / ra) + baseQuota /= ra } } // 应用新的 ratios - result := float64(baseQuota) + result := baseQuota for _, ra := range ratios { if ra != 1.0 { result *= ra } } - return int(result) + return common.QuotaFromFloat(result) } var fetchRespBuilders = map[int]func(c *gin.Context) (respBody []byte, taskResp *dto.TaskError){ diff --git a/service/http_client.go b/service/http_client.go index e31fa63..a04ad55 100644 --- a/service/http_client.go +++ b/service/http_client.go @@ -6,7 +6,6 @@ import ( "net" "net/http" "net/url" - "strconv" "sync" "time" @@ -107,38 +106,11 @@ func getCachedSSRFProtection() (*common.SSRFProtection, bool, error) { func ssrfProtectedDialContext(ctx context.Context, network, addr string) (net.Conn, error) { dialer := &net.Dialer{} - protection, enabled, err := getCachedSSRFProtection() - if err != nil { - return nil, fmt.Errorf("request reject - %v", err) - } - if !enabled { - return dialer.DialContext(ctx, network, addr) - } - - host, portStr, err := net.SplitHostPort(addr) - if err != nil { - return nil, err - } - port, err := strconv.Atoi(portStr) - if err != nil { - return nil, fmt.Errorf("invalid port: %s", portStr) - } - dialAddrs, err := protection.ResolveValidatedDialAddresses(ctx, host, port) - if err != nil { - return nil, err - } - if len(dialAddrs) == 0 { - return nil, fmt.Errorf("no validated dial addresses for %s", addr) - } - var lastErr error - for _, dialAddr := range dialAddrs { - conn, err := dialer.DialContext(ctx, network, dialAddr) - if err == nil { - return conn, nil - } - lastErr = err - } - return nil, lastErr + return (&protectedFetchDialer{ + resolver: net.DefaultResolver, + dialContext: dialer.DialContext, + getProtection: getCachedSSRFProtection, + }).DialContext(ctx, network, addr) } func newSSRFProtectedHTTPClient() *http.Client { @@ -151,15 +123,7 @@ func newSSRFProtectedHTTPClient() *http.Client { } func checkRedirect(req *http.Request, via []*http.Request) error { - fetchSetting := system_setting.GetFetchSetting() - urlStr := req.URL.String() - if err := common.ValidateURLWithFetchSetting(urlStr, fetchSetting.EnableSSRFProtection, fetchSetting.AllowPrivateIp, fetchSetting.DomainFilterMode, fetchSetting.IpFilterMode, fetchSetting.DomainList, fetchSetting.IpList, fetchSetting.AllowedPorts, fetchSetting.ApplyIPFilterForDomain); err != nil { - return fmt.Errorf("redirect to %s blocked: %v", urlStr, err) - } - if len(via) >= 10 { - return fmt.Errorf("stopped after 10 redirects") - } - return nil + return checkProtectedFetchRedirect(req, via) } func InitHttpClient() { @@ -172,6 +136,12 @@ func GetHttpClient() *http.Client { } func GetSSRFProtectedHttpClient() *http.Client { + if _, enabled, err := getCachedSSRFProtection(); err == nil && !enabled { + if httpClient != nil { + return httpClient + } + return http.DefaultClient + } if ssrfProtectedHTTPClient != nil { return ssrfProtectedHTTPClient } diff --git a/service/protected_fetch_client.go b/service/protected_fetch_client.go new file mode 100644 index 0000000..29feb78 --- /dev/null +++ b/service/protected_fetch_client.go @@ -0,0 +1,115 @@ +package service + +import ( + "context" + "fmt" + "net" + "net/http" + "strconv" + + "github.com/MAX-API-Next/MAX-API/common" +) + +type ssrfResolver interface { + LookupIPAddr(ctx context.Context, host string) ([]net.IPAddr, error) +} + +type protectedFetchDialer struct { + resolver ssrfResolver + dialContext func(ctx context.Context, network, address string) (net.Conn, error) + getProtection func() (*common.SSRFProtection, bool, error) +} + +func ValidateSSRFProtectedFetchURL(urlStr string) error { + protection, enabled, err := getCachedSSRFProtection() + if err != nil { + return fmt.Errorf("request reject - %v", err) + } + if !enabled { + return nil + } + return protection.ValidateURL(urlStr) +} + +func checkProtectedFetchRedirect(req *http.Request, via []*http.Request) error { + if req == nil || req.URL == nil { + return fmt.Errorf("invalid redirect request") + } + if err := ValidateSSRFProtectedFetchURL(req.URL.String()); err != nil { + return fmt.Errorf("redirect to %s blocked: %v", req.URL.String(), err) + } + if len(via) >= 10 { + return fmt.Errorf("stopped after 10 redirects") + } + return nil +} + +func (d *protectedFetchDialer) DialContext(ctx context.Context, network, addr string) (net.Conn, error) { + protection, enabled, err := d.getProtection() + if err != nil { + return nil, fmt.Errorf("request reject - %v", err) + } + if !enabled { + return d.dialContext(ctx, network, addr) + } + + host, portText, err := net.SplitHostPort(addr) + if err != nil { + return nil, fmt.Errorf("invalid dial address %s: %w", addr, err) + } + port, err := strconv.Atoi(portText) + if err != nil { + return nil, fmt.Errorf("invalid port: %s", portText) + } + if err := protection.ValidateNetworkTarget(host, port); err != nil { + return nil, err + } + + if ip := net.ParseIP(host); ip != nil { + return d.dialContext(ctx, network, net.JoinHostPort(ip.String(), portText)) + } + if !protection.ApplyIPFilterForDomain { + return d.dialContext(ctx, network, addr) + } + + resolved, err := d.resolver.LookupIPAddr(ctx, host) + if err != nil { + return nil, fmt.Errorf("DNS resolution failed for %s: %v", host, err) + } + + var candidateIPs []net.IP + for _, ipAddr := range resolved { + ip := ipAddr.IP + if ip == nil || !networkAllowsIP(network, ip) { + continue + } + if err := protection.ValidateResolvedIP(host, ip); err != nil { + return nil, err + } + candidateIPs = append(candidateIPs, ip) + } + + var lastDialErr error + for _, ip := range candidateIPs { + conn, err := d.dialContext(ctx, network, net.JoinHostPort(ip.String(), portText)) + if err == nil { + return conn, nil + } + lastDialErr = err + } + if lastDialErr != nil { + return nil, lastDialErr + } + return nil, fmt.Errorf("DNS resolution for %s returned no usable IP addresses", host) +} + +func networkAllowsIP(network string, ip net.IP) bool { + switch network { + case "tcp4": + return ip.To4() != nil + case "tcp6": + return ip.To4() == nil + default: + return true + } +} diff --git a/service/protected_fetch_client_test.go b/service/protected_fetch_client_test.go new file mode 100644 index 0000000..de5b1bb --- /dev/null +++ b/service/protected_fetch_client_test.go @@ -0,0 +1,225 @@ +package service + +import ( + "context" + "fmt" + "net" + "net/http" + "net/http/httptest" + "testing" + + "github.com/MAX-API-Next/MAX-API/common" + "github.com/MAX-API-Next/MAX-API/setting/system_setting" + "github.com/stretchr/testify/require" +) + +type staticSSRFResolver map[string][]net.IPAddr + +func (r staticSSRFResolver) LookupIPAddr(ctx context.Context, host string) ([]net.IPAddr, error) { + if ips, ok := r[host]; ok { + return ips, nil + } + return nil, fmt.Errorf("unexpected lookup for %s", host) +} + +func staticProtection(protection *common.SSRFProtection) func() (*common.SSRFProtection, bool, error) { + return func() (*common.SSRFProtection, bool, error) { + return protection, true, nil + } +} + +func testConn(t *testing.T) net.Conn { + t.Helper() + clientConn, serverConn := net.Pipe() + t.Cleanup(func() { + _ = clientConn.Close() + _ = serverConn.Close() + }) + return clientConn +} + +func configureSSRFProtectedFetchTest(t *testing.T) { + t.Helper() + fetchSetting := system_setting.GetFetchSetting() + original := *fetchSetting + t.Cleanup(func() { + *fetchSetting = original + }) + + fetchSetting.EnableSSRFProtection = true + fetchSetting.AllowPrivateIp = false + fetchSetting.DomainFilterMode = false + fetchSetting.IpFilterMode = false + fetchSetting.DomainList = nil + fetchSetting.IpList = nil + fetchSetting.AllowedPorts = []string{"80", "443"} + fetchSetting.ApplyIPFilterForDomain = true +} + +func TestProtectedFetchDialerRejectsPrivateReboundAddress(t *testing.T) { + dialer := &protectedFetchDialer{ + resolver: staticSSRFResolver{ + "safe.example": {{IP: net.ParseIP("127.0.0.1")}}, + }, + dialContext: func(ctx context.Context, network, address string) (net.Conn, error) { + t.Fatalf("dialContext should not be called for blocked address %s", address) + return nil, nil + }, + getProtection: staticProtection(&common.SSRFProtection{ + AllowPrivateIp: false, + DomainFilterMode: false, + IpFilterMode: false, + ApplyIPFilterForDomain: true, + }), + } + + conn, err := dialer.DialContext(context.Background(), "tcp", "safe.example:80") + + require.Error(t, err) + require.Nil(t, conn) + require.Contains(t, err.Error(), "private IP address not allowed") +} + +func TestProtectedFetchDialerRejectsMixedResolvedIPs(t *testing.T) { + var dialed []string + dialer := &protectedFetchDialer{ + resolver: staticSSRFResolver{ + "safe.example": { + {IP: net.ParseIP("10.0.0.1")}, + {IP: net.ParseIP("8.8.8.8")}, + }, + }, + dialContext: func(ctx context.Context, network, address string) (net.Conn, error) { + dialed = append(dialed, address) + return testConn(t), nil + }, + getProtection: staticProtection(&common.SSRFProtection{ + AllowPrivateIp: false, + DomainFilterMode: false, + IpFilterMode: false, + ApplyIPFilterForDomain: true, + }), + } + + conn, err := dialer.DialContext(context.Background(), "tcp", "safe.example:443") + + require.Error(t, err) + require.Nil(t, conn) + require.Empty(t, dialed) + require.Contains(t, err.Error(), "private IP address not allowed") +} + +func TestProtectedFetchDialerDialsWhenAllResolvedIPsAllowed(t *testing.T) { + var dialed []string + dialer := &protectedFetchDialer{ + resolver: staticSSRFResolver{ + "safe.example": { + {IP: net.ParseIP("8.8.8.8")}, + {IP: net.ParseIP("1.1.1.1")}, + }, + }, + dialContext: func(ctx context.Context, network, address string) (net.Conn, error) { + dialed = append(dialed, address) + return testConn(t), nil + }, + getProtection: staticProtection(&common.SSRFProtection{ + AllowPrivateIp: false, + DomainFilterMode: false, + IpFilterMode: false, + ApplyIPFilterForDomain: true, + }), + } + + conn, err := dialer.DialContext(context.Background(), "tcp", "safe.example:443") + + require.NoError(t, err) + require.NotNil(t, conn) + require.Equal(t, []string{"8.8.8.8:443"}, dialed) +} + +func TestProtectedFetchDialerAllowsPrivateIPWhenWhitelisted(t *testing.T) { + var dialed []string + dialer := &protectedFetchDialer{ + resolver: staticSSRFResolver{ + "internal.example": {{IP: net.ParseIP("10.1.2.3")}}, + }, + dialContext: func(ctx context.Context, network, address string) (net.Conn, error) { + dialed = append(dialed, address) + return testConn(t), nil + }, + getProtection: staticProtection(&common.SSRFProtection{ + AllowPrivateIp: true, + DomainFilterMode: false, + IpFilterMode: true, + IpList: []string{"10.0.0.0/8"}, + ApplyIPFilterForDomain: true, + }), + } + + conn, err := dialer.DialContext(context.Background(), "tcp", "internal.example:80") + + require.NoError(t, err) + require.NotNil(t, conn) + require.Equal(t, []string{"10.1.2.3:80"}, dialed) +} + +func TestProtectedFetchDialerSkipsResolvedIPCheckWhenDisabled(t *testing.T) { + var dialed []string + dialer := &protectedFetchDialer{ + resolver: staticSSRFResolver{}, + dialContext: func(ctx context.Context, network, address string) (net.Conn, error) { + dialed = append(dialed, address) + return testConn(t), nil + }, + getProtection: staticProtection(&common.SSRFProtection{ + AllowPrivateIp: false, + DomainFilterMode: false, + IpFilterMode: false, + ApplyIPFilterForDomain: false, + }), + } + + conn, err := dialer.DialContext(context.Background(), "tcp", "safe.example:80") + + require.NoError(t, err) + require.NotNil(t, conn) + require.Equal(t, []string{"safe.example:80"}, dialed) +} + +func TestValidateSSRFProtectedFetchURLRejectsPrivateIP(t *testing.T) { + configureSSRFProtectedFetchTest(t) + + err := ValidateSSRFProtectedFetchURL("http://127.0.0.1/resource") + + require.Error(t, err) + require.Contains(t, err.Error(), "private IP address not allowed") +} + +func TestProtectedFetchRedirectRejectsPrivateTarget(t *testing.T) { + configureSSRFProtectedFetchTest(t) + + req := httptest.NewRequest(http.MethodGet, "http://127.0.0.1/redirected", nil) + err := checkProtectedFetchRedirect(req, nil) + + require.Error(t, err) + require.Contains(t, err.Error(), "private IP address not allowed") +} + +func TestGetSSRFProtectedHTTPClientFallsBackWhenProtectionDisabled(t *testing.T) { + fetchSetting := system_setting.GetFetchSetting() + originalFetchSetting := *fetchSetting + originalHTTPClient := httpClient + originalProtectedClient := ssrfProtectedHTTPClient + t.Cleanup(func() { + *fetchSetting = originalFetchSetting + httpClient = originalHTTPClient + ssrfProtectedHTTPClient = originalProtectedClient + }) + + fetchSetting.EnableSSRFProtection = false + expected := &http.Client{} + httpClient = expected + ssrfProtectedHTTPClient = &http.Client{} + + require.Same(t, expected, GetSSRFProtectedHttpClient()) +} diff --git a/service/quota.go b/service/quota.go index 0e7cb2e..bdca2bb 100644 --- a/service/quota.go +++ b/service/quota.go @@ -54,7 +54,7 @@ func calculateAudioQuota(info QuotaInfo) int { groupRatio := decimal.NewFromFloat(info.GroupRatio) quota := modelPrice.Mul(quotaPerUnit).Mul(groupRatio) - return int(quota.IntPart()) + return decimalToQuota(quota) } completionRatio := decimal.NewFromFloat(ratio_setting.GetCompletionRatio(info.ModelName)) @@ -83,7 +83,7 @@ func calculateAudioQuota(info QuotaInfo) int { quota = decimal.NewFromInt(1) } - return int(quota.Round(0).IntPart()) + return decimalToQuota(quota) } func PreWssConsumeQuota(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, usage *dto.RealtimeUsage) error { diff --git a/service/task_billing.go b/service/task_billing.go index c0d977d..9ba0749 100644 --- a/service/task_billing.go +++ b/service/task_billing.go @@ -225,6 +225,9 @@ func RecalculateTaskQuota(ctx context.Context, task *model.Task, actualQuota int taskAdjustTokenQuota(ctx, task, quotaDelta) task.Quota = actualQuota + if err := task.UpdateQuota(); err != nil { + logger.LogError(ctx, fmt.Sprintf("差额结算回写 quota 失败 task %s: %s", task.TaskID, err.Error())) + } var logType int var logQuota int @@ -305,7 +308,7 @@ func RecalculateTaskQuotaByTokens(ctx context.Context, task *model.Task, totalTo } // 计算实际应扣费额度: totalTokens * modelRatio * groupRatio * otherMultiplier - actualQuota := int(float64(totalTokens) * modelRatio * finalGroupRatio * otherMultiplier) + actualQuota := common.QuotaFromFloat(float64(totalTokens) * modelRatio * finalGroupRatio * otherMultiplier) reason := fmt.Sprintf("token重算:tokens=%d, modelRatio=%.2f, groupRatio=%.2f, otherMultiplier=%.4f", totalTokens, modelRatio, finalGroupRatio, otherMultiplier) RecalculateTaskQuota(ctx, task, actualQuota, reason) diff --git a/service/text_quota.go b/service/text_quota.go index 273bb2f..3eced8e 100644 --- a/service/text_quota.go +++ b/service/text_quota.go @@ -145,15 +145,13 @@ func composeTieredTextQuota(relayInfo *relaycommon.RelayInfo, summary textQuotaS if tieredResult != nil { if snap := relayInfo.TieredBillingSnapshot; snap != nil { - return int(decimal.NewFromFloat(tieredResult.ActualQuotaBeforeGroup). + return decimalToQuota(decimal.NewFromFloat(tieredResult.ActualQuotaBeforeGroup). Mul(decimal.NewFromFloat(snap.GroupRatio)). - Add(summary.ToolCallSurchargeQuota). - Round(0). - IntPart()) + Add(summary.ToolCallSurchargeQuota)) } } - return tieredQuota + int(summary.ToolCallSurchargeQuota.Round(0).IntPart()) + return decimalToQuota(decimal.NewFromInt(int64(tieredQuota)).Add(summary.ToolCallSurchargeQuota)) } func calculateTextQuotaSummary(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, usage *dto.Usage) textQuotaSummary { @@ -287,7 +285,7 @@ func calculateTextQuotaSummary(ctx *gin.Context, relayInfo *relaycommon.RelayInf if !ratio.IsZero() && quotaCalculateDecimal.LessThanOrEqual(decimal.Zero) { quotaCalculateDecimal = decimal.NewFromInt(1) } - summary.Quota = int(quotaCalculateDecimal.Round(0).IntPart()) + summary.Quota = decimalToQuota(quotaCalculateDecimal) } else { quotaCalculateDecimal := dModelPrice.Mul(dQuotaPerUnit).Mul(dGroupRatio) quotaCalculateDecimal = quotaCalculateDecimal.Add(summary.ToolCallSurchargeQuota) @@ -297,7 +295,7 @@ func calculateTextQuotaSummary(ctx *gin.Context, relayInfo *relaycommon.RelayInf quotaCalculateDecimal = quotaCalculateDecimal.Mul(decimal.NewFromFloat(otherRatio)) } } - summary.Quota = int(quotaCalculateDecimal.Round(0).IntPart()) + summary.Quota = decimalToQuota(quotaCalculateDecimal) } if summary.TotalTokens == 0 { @@ -309,6 +307,11 @@ func calculateTextQuotaSummary(ctx *gin.Context, relayInfo *relaycommon.RelayInf return summary } +func decimalToQuota(d decimal.Decimal) int { + f, _ := d.Round(0).Float64() + return common.QuotaFromFloat(f) +} + func usageSemanticFromUsage(relayInfo *relaycommon.RelayInfo, usage *dto.Usage) string { if usage != nil && usage.UsageSemantic != "" { return usage.UsageSemantic diff --git a/service/text_quota_test.go b/service/text_quota_test.go index d1c3964..72e8266 100644 --- a/service/text_quota_test.go +++ b/service/text_quota_test.go @@ -1,6 +1,7 @@ package service import ( + "math" "net/http/httptest" "testing" "time" @@ -13,9 +14,17 @@ import ( "github.com/MAX-API-Next/MAX-API/types" "github.com/gin-gonic/gin" + "github.com/shopspring/decimal" "github.com/stretchr/testify/require" ) +func TestDecimalToQuotaSaturation(t *testing.T) { + overflowing := decimal.NewFromInt(2000).Mul(decimal.NewFromFloat(1.8446744073686647e19)) + require.Equal(t, math.MaxInt32, decimalToQuota(overflowing)) + require.Equal(t, math.MinInt32, decimalToQuota(overflowing.Neg())) + require.Equal(t, 42, decimalToQuota(decimal.NewFromFloat(41.7))) +} + func TestCalculateTextQuotaSummaryUnifiedForClaudeSemantic(t *testing.T) { gin.SetMode(gin.TestMode) w := httptest.NewRecorder() @@ -600,3 +609,13 @@ func TestComposeTieredTextQuotaErrorFallbackUsesPreConsumedQuota(t *testing.T) { require.Equal(t, int64(12500), summary.ToolCallSurchargeQuota.Round(0).IntPart()) require.Equal(t, 14500, quota) } + +func TestComposeTieredTextQuotaFallbackSaturatesFinalQuota(t *testing.T) { + summary := textQuotaSummary{ + ToolCallSurchargeQuota: decimal.NewFromInt(100), + } + + quota := composeTieredTextQuota(&relaycommon.RelayInfo{}, summary, math.MaxInt32-50, nil) + + require.Equal(t, math.MaxInt32, quota) +} diff --git a/service/token_counter.go b/service/token_counter.go index d3c30c7..78bd898 100644 --- a/service/token_counter.go +++ b/service/token_counter.go @@ -208,8 +208,12 @@ func EstimateRequestToken(c *gin.Context, meta *types.TokenCountMeta, info *rela if err != nil { return 0, fmt.Errorf("error getting audio duration: %v", err) } - // 一分钟 1000 token,与 $price / minute 对齐 - totalAudioToken += int(math.Round(math.Ceil(duration) / 60.0 * 1000)) + // 一分钟 1000 token,与 $price / minute 对齐。 + audioTokens := common.QuotaFromFloat(math.Round(math.Ceil(duration) / 60.0 * 1000)) + if audioTokens < 0 { + audioTokens = 0 + } + totalAudioToken += audioTokens } return totalAudioToken, nil } @@ -377,7 +381,7 @@ func CountAudioTokenInput(audioBase64 string, audioFormat string) (int, error) { if err != nil { return 0, err } - return int(duration / 60 * 100 / 0.06), nil + return common.QuotaFromFloat(duration / 60 * 100 / 0.06), nil } func CountAudioTokenOutput(audioBase64 string, audioFormat string) (int, error) { @@ -388,7 +392,7 @@ func CountAudioTokenOutput(audioBase64 string, audioFormat string) (int, error) if err != nil { return 0, err } - return int(duration / 60 * 200 / 0.24), nil + return common.QuotaFromFloat(duration / 60 * 200 / 0.24), nil } // CountTextToken 统计文本的token数量,仅OpenAI模型使用tokenizer,其余模型使用估算 diff --git a/service/tool_billing.go b/service/tool_billing.go index a476f56..af9f5f2 100644 --- a/service/tool_billing.go +++ b/service/tool_billing.go @@ -49,7 +49,7 @@ func ComputeToolCallQuota(usage ToolCallUsage, groupRatio float64) ToolCallResul return } totalPrice := pricePer1K * float64(count) / 1000 - quota := int(math.Round(totalPrice * common.QuotaPerUnit * groupRatio)) + quota := common.QuotaFromFloat(math.Round(totalPrice * common.QuotaPerUnit * groupRatio)) items = append(items, ToolCallItem{ Name: toolName, CallCount: count, @@ -70,7 +70,7 @@ func ComputeToolCallQuota(usage ToolCallUsage, groupRatio float64) ToolCallResul if usage.ImageGenerationCall { price := operation_setting.GetGPTImage1PriceOnceCall(usage.ImageGenerationQuality, usage.ImageGenerationSize) - quota := int(math.Round(price * common.QuotaPerUnit * groupRatio)) + quota := common.QuotaFromFloat(math.Round(price * common.QuotaPerUnit * groupRatio)) items = append(items, ToolCallItem{ Name: "image_generation", CallCount: 1, diff --git a/setting/ratio_setting/group_ratio.go b/setting/ratio_setting/group_ratio.go index cbfdfcc..dc1d1f3 100644 --- a/setting/ratio_setting/group_ratio.go +++ b/setting/ratio_setting/group_ratio.go @@ -70,24 +70,41 @@ func GetGroupRatioSetting() *GroupRatioSetting { } func GetGroupRatioCopy() map[string]float64 { - return groupRatioMap.ReadAll() + return filterReservedAutoRouteGroupRatios(groupRatioMap.ReadAll()) } func ContainsGroupRatio(name string) bool { - _, ok := groupRatioMap.Get(name) + trimmedName := strings.TrimSpace(name) + if isReservedAutoRouteGroupName(trimmedName) { + return false + } + _, ok := groupRatioMap.Get(trimmedName) return ok } func GroupRatio2JSONString() string { - return groupRatioMap.MarshalJSONString() + jsonBytes, err := common.Marshal(GetGroupRatioCopy()) + if err != nil { + return "{}" + } + return string(jsonBytes) } func UpdateGroupRatioByJSONString(jsonStr string) error { - return types.LoadFromJsonString(groupRatioMap, jsonStr) + normalized, err := NormalizeGroupRatioJSONString(jsonStr) + if err != nil { + return err + } + return types.LoadFromJsonString(groupRatioMap, normalized) } func GetGroupRatio(name string) float64 { - ratio, ok := groupRatioMap.Get(name) + trimmedName := strings.TrimSpace(name) + if isReservedAutoRouteGroupName(trimmedName) { + common.SysLog("group ratio not found: " + name) + return 1 + } + ratio, ok := groupRatioMap.Get(trimmedName) if !ok { common.SysLog("group ratio not found: " + name) return 1 @@ -116,21 +133,51 @@ func UpdateGroupGroupRatioByJSONString(jsonStr string) error { } func CheckGroupRatio(jsonStr string) error { + _, err := normalizeGroupRatioMap(jsonStr) + return err +} + +func NormalizeGroupRatioJSONString(jsonStr string) (string, error) { + normalized, err := normalizeGroupRatioMap(jsonStr) + if err != nil { + return "", err + } + jsonBytes, err := common.Marshal(normalized) + if err != nil { + return "", err + } + return string(jsonBytes), nil +} + +func normalizeGroupRatioMap(jsonStr string) (map[string]float64, error) { checkGroupRatio := make(map[string]float64) err := common.Unmarshal([]byte(jsonStr), &checkGroupRatio) if err != nil { - return err + return nil, err } + normalized := make(map[string]float64, len(checkGroupRatio)) for name, ratio := range checkGroupRatio { trimmedName := strings.TrimSpace(name) if isReservedAutoRouteGroupName(trimmedName) { - return errors.New("group name conflicts with auto route namespace: " + trimmedName) + continue } if ratio < 0 { - return errors.New("group ratio must be not less than 0: " + name) + return nil, errors.New("group ratio must be not less than 0: " + trimmedName) + } + normalized[trimmedName] = ratio + } + return normalized, nil +} + +func filterReservedAutoRouteGroupRatios(ratios map[string]float64) map[string]float64 { + filtered := make(map[string]float64, len(ratios)) + for name, ratio := range ratios { + if isReservedAutoRouteGroupName(strings.TrimSpace(name)) { + continue } + filtered[name] = ratio } - return nil + return filtered } func isReservedAutoRouteGroupName(name string) bool { diff --git a/setting/ratio_setting/group_ratio_test.go b/setting/ratio_setting/group_ratio_test.go index 79b5b08..ecf2894 100644 --- a/setting/ratio_setting/group_ratio_test.go +++ b/setting/ratio_setting/group_ratio_test.go @@ -6,18 +6,45 @@ import ( "github.com/stretchr/testify/require" ) -func TestCheckGroupRatioRejectsAutoRouteNamespace(t *testing.T) { - for _, jsonStr := range []string{ - `{"auto":1}`, - `{"auto:fast":1}`, - `{" auto:fast ":1}`, - } { - err := CheckGroupRatio(jsonStr) - require.Error(t, err) - require.Contains(t, err.Error(), "auto route namespace") - } +func TestGroupRatioFiltersAutoRouteNamespace(t *testing.T) { + original := GroupRatio2JSONString() + t.Cleanup(func() { + require.NoError(t, UpdateGroupRatioByJSONString(original)) + }) + + normalized, err := NormalizeGroupRatioJSONString(`{"auto":1,"auto:fast":2," auto:cheap ":3,"default":1.25,"vip":0.5}`) + require.NoError(t, err) + require.NotContains(t, normalized, `"auto"`) + require.NotContains(t, normalized, `"auto:fast"`) + require.NotContains(t, normalized, `" auto:cheap "`) + require.Contains(t, normalized, `"default":1.25`) + require.Contains(t, normalized, `"vip":0.5`) + + require.NoError(t, UpdateGroupRatioByJSONString(`{"auto":1,"default":1.25}`)) + require.NotContains(t, GetGroupRatioCopy(), "auto") + require.False(t, ContainsGroupRatio("auto")) + require.Equal(t, 1.0, GetGroupRatio("auto")) + require.Equal(t, 1.25, GetGroupRatio("default")) } func TestCheckGroupRatioAcceptsNormalGroups(t *testing.T) { require.NoError(t, CheckGroupRatio(`{"default":1,"vip":0}`)) } + +func TestGroupRatioTrimsNormalGroupNames(t *testing.T) { + original := GroupRatio2JSONString() + t.Cleanup(func() { + require.NoError(t, UpdateGroupRatioByJSONString(original)) + }) + + normalized, err := NormalizeGroupRatioJSONString(`{" vip ":0.5}`) + require.NoError(t, err) + require.Contains(t, normalized, `"vip":0.5`) + require.NotContains(t, normalized, `" vip "`) + + require.NoError(t, UpdateGroupRatioByJSONString(`{" vip ":0.5}`)) + require.True(t, ContainsGroupRatio("vip")) + require.True(t, ContainsGroupRatio(" vip ")) + require.Equal(t, 0.5, GetGroupRatio("vip")) + require.Equal(t, 0.5, GetGroupRatio(" vip ")) +} diff --git a/setting/task_billing_setting/rate_card.go b/setting/task_billing_setting/rate_card.go index afbae82..2820ee8 100644 --- a/setting/task_billing_setting/rate_card.go +++ b/setting/task_billing_setting/rate_card.go @@ -77,7 +77,7 @@ func Calculate(input types.TaskBillingInput, groupRatio float64) (*types.TaskBil return nil, nil } totalPrice := row.UnitPrice * quantity - quota := int(totalPrice * common.QuotaPerUnit * groupRatio) + quota := common.QuotaFromFloat(totalPrice * common.QuotaPerUnit * groupRatio) if totalPrice > 0 && quota <= 0 { quota = 1 } @@ -121,8 +121,8 @@ func validateRateCards(rateCards map[string]RateCard) error { return fmt.Errorf("rate card %s has no rows", key) } for i, row := range card.Rows { - if row.UnitPrice < 0 { - return fmt.Errorf("rate card %s row %d has negative unit_price", key, i) + if row.UnitPrice < 0 || math.IsNaN(row.UnitPrice) || math.IsInf(row.UnitPrice, 0) { + return fmt.Errorf("rate card %s row %d has invalid unit_price", key, i) } if len(row.Match) == 0 { return fmt.Errorf("rate card %s row %d has empty match", key, i) @@ -132,6 +132,10 @@ func validateRateCards(rateCards map[string]RateCard) error { return nil } +func validQuantity(value float64) bool { + return value > 0 && value <= math.MaxInt32 && !math.IsNaN(value) && !math.IsInf(value, 0) +} + func findRateCard(models ...string) (*RateCard, string) { rateCards := taskBillingSetting.RateCards for _, model := range models { @@ -212,23 +216,23 @@ func mergeFields(card RateCard, input types.TaskBillingInput) map[string]string func resolveQuantity(card RateCard, input types.TaskBillingInput, fields map[string]string) (float64, error) { if card.QuantityField == "" { - if card.DefaultQuantity > 0 { + if validQuantity(card.DefaultQuantity) { return card.DefaultQuantity, nil } return 1, nil } if input.Numbers != nil { - if value, ok := input.Numbers[card.QuantityField]; ok && value > 0 { + if value, ok := input.Numbers[card.QuantityField]; ok && validQuantity(value) { return value, nil } } if raw := fields[card.QuantityField]; raw != "" { value, err := strconv.ParseFloat(raw, 64) - if err == nil && value > 0 { + if err == nil && validQuantity(value) { return value, nil } } - if card.DefaultQuantity > 0 { + if validQuantity(card.DefaultQuantity) { return card.DefaultQuantity, nil } return 0, fmt.Errorf("missing positive quantity field %q", card.QuantityField) diff --git a/tools/jsonwrapcheck/allowlist.txt b/tools/jsonwrapcheck/allowlist.txt index 5958326..80c27c7 100644 --- a/tools/jsonwrapcheck/allowlist.txt +++ b/tools/jsonwrapcheck/allowlist.txt @@ -41,7 +41,6 @@ controller/midjourney.go|UpdateMidjourneyTaskBulk|Marshal|a94b7820053b86d12899b1 controller/midjourney.go|UpdateMidjourneyTaskBulk|Marshal|bcab14b3ab8142cf9c370ab4fa67ba523db802318c0238baaf97e88d98e24c59 controller/midjourney.go|UpdateMidjourneyTaskBulk|Marshal|3bfaa1b054790e95bd8bd83e0ad1ac0241147d6178e0a08467e03f4f29db6a6c controller/midjourney.go|checkMjTaskNeedUpdate|Marshal|20a71e5a83bf926ecfe69d5861a4dbbd2f976b5997798ef857aec390fa9a15e3 -controller/misc.go|ResetPassword|NewDecoder|dc2d92e0dd4c66591ccceefe086dca58c0af182c9e2c61f92c80cda6833d2f9a controller/model_meta.go|enrichModels|Marshal|a2da6b346ff93776c463c1afd5a6ad7a4e117499a61427d2b1bc5a014827e24a controller/model_meta.go|enrichModels|Marshal|a2da6b346ff93776c463c1afd5a6ad7a4e117499a61427d2b1bc5a014827e24a controller/model_sync.go|fetchJSON|Unmarshal|f56c52bcda72d3e86474db3392395083226b9fec13223b5f4489771613b9902b diff --git a/types/price_data.go b/types/price_data.go index 93bc6ae..fe5d63a 100644 --- a/types/price_data.go +++ b/types/price_data.go @@ -1,6 +1,9 @@ package types -import "fmt" +import ( + "fmt" + "math" +) type GroupRatioInfo struct { GroupRatio float64 @@ -31,7 +34,7 @@ func (p *PriceData) AddOtherRatio(key string, ratio float64) { if p.OtherRatios == nil { p.OtherRatios = make(map[string]float64) } - if ratio <= 0 { + if !(ratio > 0) || math.IsInf(ratio, 1) { return } p.OtherRatios[key] = ratio diff --git a/web/default/src/components/model-group-selector.tsx b/web/default/src/components/model-group-selector.tsx index 3a912bc..4c1e463 100644 --- a/web/default/src/components/model-group-selector.tsx +++ b/web/default/src/components/model-group-selector.tsx @@ -59,7 +59,9 @@ interface GroupOption { } function shouldShowGroupRatio(ratio: GroupOption['ratio']) { - return ratio !== undefined && ratio !== 0 && ratio !== '0' + if (ratio === undefined || ratio === '') return false + const numeric = typeof ratio === 'number' ? ratio : Number(ratio) + return !Number.isNaN(numeric) && numeric !== 0 } interface ModelSelectorProps { diff --git a/web/default/src/features/usage-logs/components/common-logs-filter-bar.tsx b/web/default/src/features/usage-logs/components/common-logs-filter-bar.tsx index cefaba3..b0875e6 100644 --- a/web/default/src/features/usage-logs/components/common-logs-filter-bar.tsx +++ b/web/default/src/features/usage-logs/components/common-logs-filter-bar.tsx @@ -43,6 +43,7 @@ import { LOG_TYPE_ALL_VALUE, LOG_TYPE_FILTER_VALUES, LOG_TYPE_FILTERS, + QUOTA_FILTER_ALL_VALUE, } from '../constants' import { buildSearchParams } from '../lib/filter' import { getDefaultTimeRange } from '../lib/utils' @@ -54,6 +55,11 @@ import { LogsFilterInput, LogsFilterToolbar, } from './logs-filter-toolbar' +import { + isQuotaFilterValue, + QuotaFilterSelect, + type QuotaFilterValue, +} from './quota-filter-select' import { useUsageLogsContext } from './usage-logs-provider' const route = getRouteApi('/_authenticated/usage-logs/$section') @@ -104,6 +110,10 @@ export function CommonLogsFilterBar