From 54e63756b97bebf8b4cab2a9f6aeaa9e89378365 Mon Sep 17 00:00:00 2001 From: CSCITech Date: Tue, 7 Jul 2026 00:20:23 +0800 Subject: [PATCH 1/6] v1.0.4-preview.2 --- model/option.go | 30 +++++++++++- model/option_test.go | 53 ++++++++++++++------ setting/ratio_setting/group_ratio.go | 59 ++++++++++++++++++++--- setting/ratio_setting/group_ratio_test.go | 29 +++++++---- 4 files changed, 136 insertions(+), 35 deletions(-) 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/setting/ratio_setting/group_ratio.go b/setting/ratio_setting/group_ratio.go index cbfdfcc..fdd40f4 100644 --- a/setting/ratio_setting/group_ratio.go +++ b/setting/ratio_setting/group_ratio.go @@ -70,23 +70,38 @@ func GetGroupRatioSetting() *GroupRatioSetting { } func GetGroupRatioCopy() map[string]float64 { - return groupRatioMap.ReadAll() + return filterReservedAutoRouteGroupRatios(groupRatioMap.ReadAll()) } func ContainsGroupRatio(name string) bool { + if isReservedAutoRouteGroupName(strings.TrimSpace(name)) { + return false + } _, ok := groupRatioMap.Get(name) 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 { + if isReservedAutoRouteGroupName(strings.TrimSpace(name)) { + common.SysLog("group ratio not found: " + name) + return 1 + } ratio, ok := groupRatioMap.Get(name) if !ok { common.SysLog("group ratio not found: " + name) @@ -116,21 +131,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: " + name) + } + normalized[name] = 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..c1769e9 100644 --- a/setting/ratio_setting/group_ratio_test.go +++ b/setting/ratio_setting/group_ratio_test.go @@ -6,16 +6,25 @@ 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) { From 6df7ad2ac767aa399bc00ad84cf11340954962db Mon Sep 17 00:00:00 2001 From: CSCITech Date: Tue, 7 Jul 2026 12:43:52 +0800 Subject: [PATCH 2/6] v1.0.4-preview.2 --- common/quota_math.go | 21 ++ common/quota_math_test.go | 18 ++ common/ssrf_protection.go | 13 + controller/misc.go | 47 ++- controller/model_list_test.go | 81 +++++ controller/oauth.go | 21 +- controller/task_video.go | 2 +- controller/token_test.go | 4 +- controller/user.go | 65 +++- dto/openai_image.go | 3 + go.mod | 12 +- go.sum | 32 +- i18n/keys.go | 3 + i18n/locales/en.yaml | 3 + i18n/locales/zh-CN.yaml | 5 +- i18n/locales/zh-TW.yaml | 3 + model/errors.go | 3 + model/task.go | 4 + model/user.go | 347 ++++++++++++++++++--- model/user_update_test.go | 162 ++++++++++ pkg/billingexpr/billingexpr_test.go | 4 + pkg/billingexpr/round.go | 15 +- relay/channel/ali/image.go | 3 + relay/channel/api_request.go | 42 +-- relay/channel/gemini/relay_responses.go | 5 +- relay/channel/openai/audio.go | 4 +- relay/channel/openai/helper.go | 2 +- relay/channel/openai/relay_image.go | 20 +- relay/channel/openai/responses_via_chat.go | 5 +- relay/channel/task/ali/adaptor.go | 8 +- relay/channel/task/gemini/billing.go | 7 +- relay/channel/task/kling/adaptor.go | 2 +- relay/common/relay_utils.go | 23 ++ relay/common/relay_utils_test.go | 62 ++++ relay/helper/common.go | 26 +- relay/helper/max_tokens_bounds_test.go | 67 ++++ relay/helper/openai_image_request_test.go | 82 +++++ relay/helper/price.go | 8 +- relay/helper/stream_scanner.go | 128 ++++---- relay/helper/stream_scanner_test.go | 52 +++ relay/helper/valid_request.go | 35 ++- relay/relay_task.go | 24 +- service/http_client.go | 54 +--- service/protected_fetch_client.go | 115 +++++++ service/protected_fetch_client_test.go | 225 +++++++++++++ service/quota.go | 4 +- service/task_billing.go | 5 +- service/text_quota.go | 17 +- service/text_quota_test.go | 9 + service/token_counter.go | 12 +- service/tool_billing.go | 4 +- setting/task_billing_setting/rate_card.go | 18 +- tools/jsonwrapcheck/allowlist.txt | 1 - types/price_data.go | 7 +- 54 files changed, 1622 insertions(+), 322 deletions(-) create mode 100644 common/quota_math.go create mode 100644 common/quota_math_test.go create mode 100644 relay/helper/max_tokens_bounds_test.go create mode 100644 relay/helper/openai_image_request_test.go create mode 100644 service/protected_fetch_client.go create mode 100644 service/protected_fetch_client_test.go 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/misc.go b/controller/misc.go index 53d7eee..8324572 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) @@ -330,6 +322,8 @@ func SendPasswordResetEmail(c *gin.Context) { if err != nil { logger.LogError(c.Request.Context(), fmt.Sprintf("failed to send password reset email to %s: %s", email, err.Error())) } + } else if err != nil && !errors.Is(err, model.ErrEmailNotFound) { + logger.LogWarn(c.Request.Context(), fmt.Sprintf("skip password reset email for %s: %s", 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..47f245e 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 @@ -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/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/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..3b692be 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 + } } - exist, err := model.CheckUserExistOrDeleted(user.Username, user.Email) + emailForExistCheck := "" + if common.EmailVerificationEnabled { + emailForExistCheck = 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 } @@ -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 :) } @@ -777,6 +809,14 @@ func UpdateSelf(c *gin.Context) { } 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 } @@ -793,19 +833,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 @@ -1084,7 +1126,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 +1142,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/dto/openai_image.go b/dto/openai_image.go index 1991075..7930a68 100644 --- a/dto/openai_image.go +++ b/dto/openai_image.go @@ -11,6 +11,9 @@ import ( "github.com/gin-gonic/gin" ) +// MaxImageN caps image generation count before it becomes a billing multiplier. +const MaxImageN = 128 + type ImageRequest struct { Model string `json:"model"` Prompt string `json:"prompt" binding:"required"` diff --git a/go.mod b/go.mod index e59b5cf..eed20e9 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/crypto v0.51.0 + golang.org/x/image v0.41.0 + golang.org/x/net v0.55.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/sys v0.45.0 + golang.org/x/text v0.37.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..a0bc3a5 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,18 +327,18 @@ 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.51.0 h1:IBPXwPfKxY7cWQZ38ZCIRPI50YLeevDLlLnyC5wRGTI= +golang.org/x/crypto v0.51.0/go.mod h1:8AdwkbraGNABw2kOX6YFPs3WM22XqI4EXEd8g+x7Oc8= 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.41.0 h1:8wS72eGJMJaBxK6okTzd4WaXumUlTVlb753MlsSvTCo= +golang.org/x/image v0.41.0/go.mod h1:uIc348UZMSvS5Z65CVZ7iDPaNobNFEPeJ4kbqTOszmA= +golang.org/x/mod v0.35.0 h1:Ww1D637e6Pg+Zb2KrWfHQUnH2dQRLBQyAtpr/haaJeM= +golang.org/x/mod v0.35.0/go.mod h1:+GwiRhIInF8wPm+4AoT6L0FA1QWAad3OMdTRx4tFYlU= 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/net v0.55.0 h1:bcvxaJn3e1U6InsFWt1JUq1aSjnRxLzT2rtD2KfkDF8= +golang.org/x/net v0.55.0/go.mod h1:L5U2KuzuOe1lY7Z+aWVIKK6qEeJXnXV9yzGA+WCHJww= 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/sys v0.0.0-20190726091711-fc99dfbffb4e/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.37.0 h1:Cqjiwd9eSg8e0QAkyCaQTNHFIIzWtidPahFWR83rTrc= +golang.org/x/text v0.37.0/go.mod h1:a5sjxXGs9hsn/AJVwuElvCAo9v8QYLzvavO5z2PiM38= 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.44.0 h1:UP4ajHPIcuMjT1GqzDWRlalUEoY+uzoZKnhOjbIPD2c= +golang.org/x/tools v0.44.0/go.mod h1:KA0AfVErSdxRZIsOVipbv3rQhVXTnlU6UhKxHd1seDI= 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/task.go b/model/task.go index 9c1929f..8a092ed 100644 --- a/model/task.go +++ b/model/task.go @@ -402,6 +402,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. diff --git a/model/user.go b/model/user.go index 572d3aa..0ee6791 100644 --- a/model/user.go +++ b/model/user.go @@ -187,10 +187,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 LOWER(email) = ?", username, email).Error } if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { @@ -204,6 +205,76 @@ func CheckUserExistOrDeleted(username string, email string) (bool, error) { return true, nil } +func NormalizeEmail(email string) string { + return strings.ToLower(strings.TrimSpace(email)) +} + +func emailQuery(tx *gorm.DB, email string) *gorm.DB { + if tx == nil { + tx = DB + } + return tx.Unscoped().Model(&User{}).Where("LOWER(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 + } + case common.UsingMySQL: + var ids []int + if err := tx.Raw("SELECT id FROM users WHERE LOWER(email) = ? FOR UPDATE", email).Scan(&ids).Error; err != nil { + return err + } + } + return fn(tx) +} + func GetMaxUserId() int { var user User DB.Unscoped().Last(&user) @@ -430,28 +501,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.Email = NormalizeEmail(user.Email) + 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 { - return err - } + 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 } - user.Quota = common.QuotaForNewUser - //user.SetAccessToken(common.GetUUID()) - user.AffCode = common.GetRandomString(4) + 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 +} - // 初始化用户设置,包括默认的边栏配置 - if user.Setting == "" { - defaultSetting := dto.UserSetting{} - // 这里暂时不设置SidebarModules,因为需要在用户创建后根据角色设置 - user.SetSetting(defaultSetting) +// 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 := DB.Transaction(func(tx *gorm.DB) error { + return withNormalizedEmailLock(tx, 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).Update("email", email).Error; err != nil { + return err + } + user.Email = email + return tx.First(user, user.Id).Error + }) + }); err != nil { + return err } + return updateUserCache(*user) +} - result := DB.Create(user) - if result.Error != nil { - return result.Error +func (user *User) Insert(inviterId int) error { + if err := DB.Transaction(func(tx *gorm.DB) error { + 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{} + // 这里暂时不设置SidebarModules,因为需要在用户创建后根据角色设置 + user.SetSetting(defaultSetting) + } + + return tx.Create(user).Error + }) + }); err != nil { + return err } // 用户创建成功后,根据角色初始化边栏配置 @@ -490,28 +612,21 @@ func (user *User) Insert(inviterId int) error { // 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) + user.Quota = common.QuotaForNewUser + user.AffCode = common.GetRandomString(4) - // 初始化用户设置 - if user.Setting == "" { - defaultSetting := dto.UserSetting{} - user.SetSetting(defaultSetting) - } - - 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. @@ -546,6 +661,16 @@ func (user *User) FinalizeOAuthUserCreation(inviterId int) { } func (user *User) Update(updatePassword bool) error { + if err := user.UpdateWithTx(DB, updatePassword); 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 { var err error if updatePassword { user.Password, err = common.Password2Hash(user.Password) @@ -555,21 +680,124 @@ 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)) + 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) map[string]interface{} { + fullUser := newUser.CreatedAt != 0 + updates := map[string]interface{}{} + + if fullUser || newUser.Username != "" { + updates["username"] = newUser.Username + } + if fullUser || newUser.DisplayName != "" { + updates["display_name"] = newUser.DisplayName + } + if fullUser || newUser.Role != 0 { + updates["role"] = newUser.Role + } + if fullUser || newUser.Status != 0 { + updates["status"] = newUser.Status + } + if fullUser || newUser.Email != "" { + updates["email"] = newUser.Email + } + if fullUser || newUser.GitHubId != "" { + updates["github_id"] = newUser.GitHubId + } + if fullUser || newUser.DiscordId != "" { + updates["discord_id"] = newUser.DiscordId + } + if fullUser || newUser.OidcId != "" { + updates["oidc_id"] = newUser.OidcId + } + if fullUser || newUser.WeChatId != "" { + updates["wechat_id"] = newUser.WeChatId + } + if fullUser || newUser.TelegramId != "" { + updates["telegram_id"] = newUser.TelegramId + } + if fullUser || newUser.AccessToken != nil { + updates["access_token"] = newUser.AccessToken + } + if fullUser || newUser.Group != "" { + updates["group"] = newUser.Group + } + if fullUser || newUser.AffCode != "" { + updates["aff_code"] = newUser.AffCode + } + if fullUser || newUser.AffCount != 0 { + updates["aff_count"] = newUser.AffCount + } + if fullUser || newUser.AffQuota != 0 { + updates["aff_quota"] = newUser.AffQuota + } + if fullUser || newUser.AffHistoryQuota != 0 { + updates["aff_history"] = newUser.AffHistoryQuota + } + if fullUser || newUser.InviterId != 0 { + updates["inviter_id"] = newUser.InviterId + } + if fullUser || newUser.LinuxDOId != "" { + updates["linux_do_id"] = newUser.LinuxDOId + } + if fullUser || newUser.Setting != "" { + updates["setting"] = newUser.Setting + } + if fullUser || newUser.Remark != "" { + updates["remark"] = newUser.Remark + } + if fullUser || newUser.StripeCustomer != "" { + updates["stripe_customer"] = newUser.StripeCustomer + } + if fullUser || newUser.LastLoginAt != 0 { + updates["last_login_at"] = newUser.LastLoginAt + } + 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)) + if !fullUser { + copyUnspecifiedUserUpdateValues(updates, current) + } + return updates +} + +func copyUnspecifiedUserUpdateValues(updates map[string]interface{}, current User) { + defaults := map[string]interface{}{ + "role": current.Role, + "status": current.Status, + "email": 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 { @@ -671,6 +899,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 +977,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("LOWER(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 +1024,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..62264e4 100644 --- a/model/user_update_test.go +++ b/model/user_update_test.go @@ -84,6 +84,41 @@ 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.Update(false)) + + 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 TestUserUpdateIgnoresCacheWriteFailure(t *testing.T) { setupUserUpdateTestState(t) @@ -164,3 +199,130 @@ 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 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..2375548 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 imageRequest.Parameters.N < 0 || imageRequest.Parameters.N > dto.MaxImageN { + return nil, fmt.Errorf("parameters.n must be an integer between 1 and %d", dto.MaxImageN) + } 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..d6e1d1f 100644 --- a/relay/channel/api_request.go +++ b/relay/channel/api_request.go @@ -400,10 +400,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,36 +455,22 @@ 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() + mutex.Lock() + defer mutex.Unlock() - err := helper.PingData(c) - if err != nil { - logger.LogError(c, "SSE ping error: "+err.Error()) - done <- err - return - } - - logger.LogDebug(c, "SSE ping data sent") - done <- nil - }() - - // 设置发送ping数据的超时时间 - select { - case err := <-done: + helper.ExtendWriteDeadline(c) + err := helper.PingData(c) + if err != nil { + logger.LogError(c, "SSE ping error: "+err.Error()) 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") } + + logger.LogDebug(c, "SSE ping data sent") + return nil } func DoRequest(c *gin.Context, req *http.Request, info *common.RelayInfo) (*http.Response, error) { @@ -501,17 +489,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") } }() 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..5432ef7 100644 --- a/relay/channel/task/ali/adaptor.go +++ b/relay/channel/task/ali/adaptor.go @@ -520,7 +520,8 @@ func (a *TaskAdaptor) convertToAliRequest(info *relaycommon.RelayInfo, req relay } else { aliReq.Parameters.Duration = seconds } - } else { + } + if aliReq.Parameters.Duration <= 0 { aliReq.Parameters.Duration = 5 // 默认5秒 } @@ -565,6 +566,9 @@ func (a *TaskAdaptor) convertToAliKlingRequest(upstreamModel string, req relayco } aliReq.Parameters.Duration = seconds } + if aliReq.Parameters.Duration <= 0 { + aliReq.Parameters.Duration = 5 + } if imageURL := firstNonEmpty(req.InputReference, req.Image); imageURL != "" { aliReq.Input.Media = []map[string]interface{}{ { @@ -773,7 +777,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/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_utils.go b/relay/common/relay_utils.go index ec19987..e1d66d9 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 := req.DurationValue() + if seconds == 0 && req.Seconds != "" { + seconds, _ = strconv.Atoi(req.Seconds) + } + if seconds < 0 || seconds > MaxTaskDurationSeconds { + return createTaskError(fmt.Errorf("seconds must be between 1 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 { @@ -190,6 +205,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 +270,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..e3934ed 100644 --- a/relay/common/relay_utils_test.go +++ b/relay/common/relay_utils_test.go @@ -53,6 +53,68 @@ 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 + }{ + { + name: "huge duration is rejected", + body: `{"model":"sora-2","prompt":"a cat","duration":9999999999}`, + wantErr: true, + }, + { + name: "huge seconds string is rejected", + body: `{"model":"sora-2","prompt":"a cat","seconds":"9999999999"}`, + wantErr: true, + }, + { + name: "negative duration is rejected", + body: `{"model":"sora-2","prompt":"a cat","duration":-8}`, + wantErr: true, + }, + { + 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) + 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) + 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..f0da06a 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,10 @@ func SetEventStreamHeaders(c *gin.Context) { } func ClaudeData(c *gin.Context, resp dto.ClaudeResponse) error { + if requestContextDone(c) { + return nil + } + jsonData, err := common.Marshal(resp) if err != nil { common.SysError("error marshalling stream response: " + err.Error()) @@ -67,15 +75,23 @@ func ClaudeData(c *gin.Context, resp dto.ClaudeResponse) error { } func ClaudeChunkData(c *gin.Context, resp dto.ClaudeResponse, data string) { + 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 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 +99,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 +112,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/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..7dc2e1b --- /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 1 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..ffeb7fe 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,13 @@ 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 || n < 0 || n > dto.MaxImageN { + return nil, fmt.Errorf("n must be an integer between 1 and %d", dto.MaxImageN) + } + imageRequest.N = common.GetPointer(uint(n)) + } imageRequest.Quality = formData.Get("quality") imageRequest.Size = formData.Get("size") if imageValue := formData.Get("image"); imageValue != "" { @@ -190,6 +211,10 @@ 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 && *imageRequest.N > dto.MaxImageN { + return nil, fmt.Errorf("n must be an integer between 1 and %d", dto.MaxImageN) + } + // 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 +263,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 +288,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 +341,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..f8f60b9 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 tieredQuota + decimalToQuota(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..e52a1f1 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() 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/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 From 3e487b161362483506c2c31388715d7e17c7b90e Mon Sep 17 00:00:00 2001 From: CSCITech Date: Tue, 7 Jul 2026 13:55:02 +0800 Subject: [PATCH 3/6] v1.0.4-preview.2 --- controller/log.go | 12 +- controller/midjourney.go | 2 + controller/misc.go | 6 +- controller/oauth.go | 4 +- controller/task.go | 2 + model/log.go | 44 +++- model/log_test.go | 124 +++++++++-- model/main.go | 19 ++ model/midjourney.go | 5 + model/payment_method_guard_test.go | 33 +++ model/task.go | 5 + model/topup.go | 41 ++-- model/user.go | 202 +++++++++++++----- model/user_update_test.go | 195 +++++++++++++++++ relay/channel/api_request.go | 47 +++- relay/channel/api_request_test.go | 99 +++++++++ service/text_quota.go | 2 +- service/text_quota_test.go | 10 + setting/ratio_setting/group_ratio.go | 4 +- setting/ratio_setting/group_ratio_test.go | 16 ++ .../src/components/model-group-selector.tsx | 4 +- .../components/common-logs-filter-bar.tsx | 44 +++- .../components/quota-filter-select.tsx | 86 ++++++++ .../components/task-logs-filter-bar.tsx | 40 +++- .../src/features/usage-logs/constants.ts | 26 +++ .../src/features/usage-logs/lib/filter.ts | 1 + .../src/features/usage-logs/lib/utils.test.ts | 45 +++- .../src/features/usage-logs/lib/utils.ts | 10 + web/default/src/features/usage-logs/types.ts | 5 + web/default/src/i18n/locales/en.json | 4 + web/default/src/i18n/locales/fr.json | 4 + web/default/src/i18n/locales/ja.json | 4 + web/default/src/i18n/locales/ru.json | 4 + web/default/src/i18n/locales/vi.json | 4 + web/default/src/i18n/locales/zh.json | 4 + .../_authenticated/usage-logs/$section.tsx | 6 +- 36 files changed, 1040 insertions(+), 123 deletions(-) create mode 100644 web/default/src/features/usage-logs/components/quota-filter-select.tsx diff --git a/controller/log.go b/controller/log.go index 9839641..fffc414 100644 --- a/controller/log.go +++ b/controller/log.go @@ -25,7 +25,8 @@ 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(logType, logFilter, startTimestamp, endTimestamp, modelName, username, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), channel, group, requestId, upstreamRequestId, quotaFilter) if err != nil { common.ApiError(c, err) return @@ -77,7 +78,8 @@ 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(userId, logType, logFilter, startTimestamp, endTimestamp, modelName, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), group, requestId, upstreamRequestId, quotaFilter) if err != nil { common.ApiError(c, err) return @@ -165,7 +167,8 @@ 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(logType, logFilter, startTimestamp, endTimestamp, modelName, username, tokenName, channel, group, quotaFilter) if err != nil { common.ApiError(c, err) return @@ -193,7 +196,8 @@ 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(logType, logFilter, startTimestamp, endTimestamp, modelName, username, tokenName, channel, group, 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 8324572..4adc489 100644 --- a/controller/misc.go +++ b/controller/misc.go @@ -320,10 +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 err != nil && !errors.Is(err, model.ErrEmailNotFound) { - logger.LogWarn(c.Request.Context(), fmt.Sprintf("skip password reset email for %s: %s", 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, diff --git a/controller/oauth.go b/controller/oauth.go index 47f245e..7db4337 100644 --- a/controller/oauth.go +++ b/controller/oauth.go @@ -285,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 @@ -311,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 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/model/log.go b/model/log.go index 8908317..19f2bec 100644 --- a/model/log.go +++ b/model/log.go @@ -91,6 +91,38 @@ const ( LogFilterEmptyRetry = "empty_retry" ) +const ( + LogQuotaFilterAbnormal = "abnormal" + LogQuotaFilterZero = "zero" + LogQuotaFilterNegative = "negative" +) + +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,11 +636,12 @@ 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) { +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, quotaFilter string) (logs []*Log, total int64, err error) { tx, err := applyLogFilter(applyLogTypeFilter(LOG_DB, logType), logFilter) if err != nil { return nil, 0, err } + tx = applyQuotaFilter(tx, "logs.quota", quotaFilter) if tx, err = applyExplicitLogTextFilter(tx, "logs.model_name", modelName); err != nil { return nil, 0, err @@ -692,11 +725,12 @@ 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) { +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, quotaFilter string) (logs []*Log, total int64, err error) { tx, err := applyLogFilter(applyLogTypeFilter(LOG_DB.Where("logs.user_id = ?", userId), logType), logFilter) if err != nil { return nil, 0, err } + tx = applyQuotaFilter(tx, "logs.quota", quotaFilter) if tx, err = applyExplicitLogTextFilter(tx, "logs.model_name", modelName); err != nil { return nil, 0, err @@ -771,7 +805,7 @@ 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(logType int, logFilter string, startTimestamp int64, endTimestamp int64, modelName string, username string, tokenName string, channel int, group string, quotaFilter string) (stat Stat, err error) { tx := LOG_DB.Table("logs").Select("sum(quota) quota") // 为rpm和tpm创建单独的查询 @@ -785,6 +819,8 @@ func SumUsedQuota(logType int, logFilter string, startTimestamp int64, endTimest if err != nil { return stat, err } + tx = applyQuotaFilter(tx, "logs.quota", quotaFilter) + rpmTpmQuery = applyQuotaFilter(rpmTpmQuery, "logs.quota", quotaFilter) if tx, err = applyExplicitLogTextFilter(tx, "username", username); err != nil { return stat, err @@ -820,7 +856,7 @@ func SumUsedQuota(logType int, logFilter string, startTimestamp int64, endTimest if logType != LogTypeUnknown { tx = tx.Where("logs.type = ?", logType) rpmTpmQuery = rpmTpmQuery.Where("logs.type = ?", logType) - } else if !isRetryLogFilter(logFilter) { + } 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..c1a83e2 100644 --- a/model/log_test.go +++ b/model/log_test.go @@ -32,7 +32,7 @@ func TestGetAllLogsRetryFilter(t *testing.T) { logs := createRetryFilterLogs(t) - got, total, err := GetAllLogs(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "") + got, total, err := GetAllLogs(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "", "") require.NoError(t, err) require.EqualValues(t, 5, total) require.Len(t, got, 5) @@ -47,7 +47,7 @@ func TestGetUserLogsRetryFilter(t *testing.T) { logs := createRetryFilterLogs(t) - got, total, err := GetUserLogs(1, LogTypeUnknown, LogFilterRetry, 0, 0, "", "", 0, 10, "", "", "") + got, total, err := GetUserLogs(1, LogTypeUnknown, LogFilterRetry, 0, 0, "", "", 0, 10, "", "", "", "") require.NoError(t, err) require.EqualValues(t, 4, total) require.Len(t, got, 4) @@ -67,10 +67,36 @@ func TestSumUsedQuotaRetryFilter(t *testing.T) { createRetryFilterLogs(t) - stat, err := SumUsedQuota(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, "") + stat, err := SumUsedQuota(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, "", "") 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) +} + +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(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, "", "") + require.NoError(t, err) + require.Equal(t, 1550, stat.Quota) + require.Equal(t, 4, stat.Rpm) require.Equal(t, 96, stat.Tpm) } @@ -82,13 +108,13 @@ func TestGetAllLogsRetrySubtypeFilters(t *testing.T) { logs := createRetryFilterLogs(t) - errorLogs, total, err := GetAllLogs(LogTypeUnknown, LogFilterErrorRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "") + errorLogs, total, err := GetAllLogs(LogTypeUnknown, LogFilterErrorRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "", "") 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(LogTypeUnknown, LogFilterEmptyRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "", "") require.NoError(t, err) require.EqualValues(t, 2, total) require.Len(t, emptyLogs, 2) @@ -103,13 +129,13 @@ func TestSumUsedQuotaRetrySubtypeFilters(t *testing.T) { createRetryFilterLogs(t) - errorStat, err := SumUsedQuota(LogTypeUnknown, LogFilterErrorRetry, 0, 0, "", "", "", 0, "") + errorStat, err := SumUsedQuota(LogTypeUnknown, LogFilterErrorRetry, 0, 0, "", "", "", 0, "", "") 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(LogTypeUnknown, LogFilterEmptyRetry, 0, 0, "", "", "", 0, "", "") require.NoError(t, err) require.Equal(t, 550, emptyStat.Quota) require.Equal(t, 2, emptyStat.Rpm) @@ -124,7 +150,7 @@ func TestSumUsedQuotaAppliesExplicitLogType(t *testing.T) { createRetryFilterLogs(t) - stat, err := SumUsedQuota(LogTypeError, "", 0, 0, "", "", "", 0, "") + stat, err := SumUsedQuota(LogTypeError, "", 0, 0, "", "", "", 0, "", "") require.NoError(t, err) require.Equal(t, 0, stat.Quota) require.Equal(t, 1, stat.Rpm) @@ -158,7 +184,7 @@ 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(LogTypeUnknown, "", now-90000, now-80000, "", "", "", 0, "", "") require.NoError(t, err) require.Equal(t, 200, stat.Quota) @@ -166,6 +192,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(LogTypeConsume, "", 0, 0, "", "", "", 0, 10, 0, "", "", "", 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(1, LogTypeConsume, "", 0, 0, "", "", 0, 10, "", "", "", 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(LogTypeUnknown, "", 0, 0, "", "", "", 0, "", 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 +289,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(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "", "") 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(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, "", "") require.NoError(t, err) require.Equal(t, 200, stat.Quota) require.Equal(t, 1, stat.Rpm) @@ -250,7 +338,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(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "", "") require.NoError(t, err) require.EqualValues(t, 2, total) require.Len(t, got, 2) @@ -289,7 +377,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(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "", "") require.NoError(t, err) require.EqualValues(t, 0, total) require.Empty(t, got) @@ -309,17 +397,17 @@ func TestRetryFilterReadPathsReturnReadinessError(t *testing.T) { ensureLogRetryMarkerBackfillCompletedForRead = originalEnsure }) - got, total, err := GetAllLogs(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "") + got, total, err := GetAllLogs(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, 10, 0, "", "", "", "") 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(1, LogTypeUnknown, LogFilterRetry, 0, 0, "", "", 0, 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(LogTypeUnknown, LogFilterRetry, 0, 0, "", "", "", 0, "", "") 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/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 8a092ed..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 != "" { @@ -455,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) } @@ -490,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/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 0ee6791..56679a2 100644 --- a/model/user.go +++ b/model/user.go @@ -29,6 +29,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 +55,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, @@ -191,7 +202,7 @@ func CheckUserExistOrDeleted(username string, email string) (bool, error) { if email == "" { err = DB.Unscoped().First(&user, "username = ?", username).Error } else { - err = DB.Unscoped().First(&user, "username = ? or LOWER(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) { @@ -209,11 +220,15 @@ 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("LOWER(email) = ?", NormalizeEmail(email)) + return tx.Unscoped().Model(&User{}).Where("normalized_email = ?", NormalizeEmail(email)) } func CountUsersByEmail(email string) (int64, error) { @@ -266,13 +281,79 @@ func withNormalizedEmailLock(tx *gorm.DB, email string, fn func(tx *gorm.DB) err if err := tx.Exec("SELECT pg_advisory_xact_lock(hashtext(?))", email).Error; err != nil { return err } + return fn(tx) case common.UsingMySQL: - var ids []int - if err := tx.Raw("SELECT id FROM users WHERE LOWER(email) = ? FOR UPDATE", email).Scan(&ids).Error; err != nil { + 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) } - 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 { @@ -502,7 +583,7 @@ func (user *User) TransferAffQuotaToQuota(quota int) error { } func (user *User) prepareForInsert(tx *gorm.DB) error { - user.Email = NormalizeEmail(user.Email) + user.normalizeEmailFields() if err := ensureEmailAvailableWithTx(tx, user.Email, 0); err != nil { return err } @@ -537,17 +618,19 @@ func ensureEmailAvailableWithTx(tx *gorm.DB, email string, excludeUserID int) er // user, preventing concurrent binds from sharing the same normalized address. func BindEmailToUser(user *User, email string) error { email = NormalizeEmail(email) - if err := DB.Transaction(func(tx *gorm.DB) error { - return withNormalizedEmailLock(tx, 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).Update("email", email).Error; err != nil { - return err - } - user.Email = email - return tx.First(user, user.Id).Error - }) + 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 } @@ -555,23 +638,21 @@ func BindEmailToUser(user *User, email string) error { } func (user *User) Insert(inviterId int) error { - if err := DB.Transaction(func(tx *gorm.DB) error { - 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{} - // 这里暂时不设置SidebarModules,因为需要在用户创建后根据角色设置 - user.SetSetting(defaultSetting) - } + 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) - return tx.Create(user).Error - }) + // 初始化用户设置,包括默认的边栏配置 + if user.Setting == "" { + defaultSetting := dto.UserSetting{} + // 这里暂时不设置SidebarModules,因为需要在用户创建后根据角色设置 + user.SetSetting(defaultSetting) + } + + return tx.Create(user).Error }); err != nil { return err } @@ -609,6 +690,8 @@ 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 { @@ -707,7 +790,9 @@ func buildUserUpdateValues(current User, newUser User, updatePassword bool) map[ updates["status"] = newUser.Status } if fullUser || newUser.Email != "" { - updates["email"] = newUser.Email + email := NormalizeEmail(newUser.Email) + updates["email"] = email + updates["normalized_email"] = email } if fullUser || newUser.GitHubId != "" { updates["github_id"] = newUser.GitHubId @@ -772,26 +857,27 @@ func buildUserUpdateValues(current User, newUser User, updatePassword bool) map[ func copyUnspecifiedUserUpdateValues(updates map[string]interface{}, current User) { defaults := map[string]interface{}{ - "role": current.Role, - "status": current.Status, - "email": 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, + "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 { @@ -848,7 +934,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 } @@ -987,7 +1077,7 @@ func GetUniqueUserByEmail(email string) (*User, error) { return nil, ErrEmailNotFound } var users []User - if err := DB.Where("LOWER(email) = ?", email).Limit(2).Find(&users).Error; err != nil { + if err := DB.Where("normalized_email = ?", email).Limit(2).Find(&users).Error; err != nil { return nil, err } switch len(users) { diff --git a/model/user_update_test.go b/model/user_update_test.go index 62264e4..f55a345 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) @@ -172,6 +256,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) @@ -220,6 +392,29 @@ func TestEnsureEmailAvailableRejectsExistingEmailCaseInsensitive(t *testing.T) { 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) diff --git a/relay/channel/api_request.go b/relay/channel/api_request.go index d6e1d1f..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": {}, @@ -459,18 +461,42 @@ func startPingKeepAlive(c *gin.Context, pingInterval time.Duration) (context.Can } func sendPingData(c *gin.Context, mutex *sync.Mutex) error { - mutex.Lock() - defer mutex.Unlock() + 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()) - return err + helper.ExtendWriteDeadline(c) + err := helper.PingData(c) + if err != nil { + logger.LogError(c, "SSE ping error: "+err.Error()) + done <- err + return + } + + logger.LogDebug(c, "SSE ping data sent") + done <- nil + }() + + timer := time.NewTimer(sendPingDataTimeout) + defer timer.Stop() + + var requestDone <-chan struct{} + if c != nil && c.Request != nil { + requestDone = c.Request.Context().Done() } - logger.LogDebug(c, "SSE ping data sent") - return nil + select { + case err := <-done: + return err + 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) + } } func DoRequest(c *gin.Context, req *http.Request, info *common.RelayInfo) (*http.Response, error) { @@ -570,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/service/text_quota.go b/service/text_quota.go index f8f60b9..3eced8e 100644 --- a/service/text_quota.go +++ b/service/text_quota.go @@ -151,7 +151,7 @@ func composeTieredTextQuota(relayInfo *relaycommon.RelayInfo, summary textQuotaS } } - return tieredQuota + decimalToQuota(summary.ToolCallSurchargeQuota) + return decimalToQuota(decimal.NewFromInt(int64(tieredQuota)).Add(summary.ToolCallSurchargeQuota)) } func calculateTextQuotaSummary(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, usage *dto.Usage) textQuotaSummary { diff --git a/service/text_quota_test.go b/service/text_quota_test.go index e52a1f1..72e8266 100644 --- a/service/text_quota_test.go +++ b/service/text_quota_test.go @@ -609,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/setting/ratio_setting/group_ratio.go b/setting/ratio_setting/group_ratio.go index fdd40f4..e9d80b9 100644 --- a/setting/ratio_setting/group_ratio.go +++ b/setting/ratio_setting/group_ratio.go @@ -160,9 +160,9 @@ func normalizeGroupRatioMap(jsonStr string) (map[string]float64, error) { continue } if ratio < 0 { - return nil, 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[name] = ratio + normalized[trimmedName] = ratio } return normalized, nil } diff --git a/setting/ratio_setting/group_ratio_test.go b/setting/ratio_setting/group_ratio_test.go index c1769e9..fb8a476 100644 --- a/setting/ratio_setting/group_ratio_test.go +++ b/setting/ratio_setting/group_ratio_test.go @@ -30,3 +30,19 @@ func TestGroupRatioFiltersAutoRouteNamespace(t *testing.T) { 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.Equal(t, 0.5, GetGroupRatio("vip")) +} 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( username: searchParams.username || undefined, requestId: searchParams.requestId || undefined, upstreamRequestId: searchParams.upstreamRequestId || undefined, + quotaFilter: + searchParams.quotaFilter && isQuotaFilterValue(searchParams.quotaFilter) + ? searchParams.quotaFilter + : QUOTA_FILTER_ALL_VALUE, }) const typeArr = searchParams.type @@ -124,6 +134,7 @@ export function CommonLogsFilterBar( searchParams.username, searchParams.requestId, searchParams.upstreamRequestId, + searchParams.quotaFilter, searchParams.type, ]) @@ -151,7 +162,11 @@ export function CommonLogsFilterBar( const handleReset = useCallback(() => { const { start, end } = getDefaultTimeRange() - const resetFilters: CommonLogFilters = { startTime: start, endTime: end } + const resetFilters: CommonLogFilters = { + startTime: start, + endTime: end, + quotaFilter: QUOTA_FILTER_ALL_VALUE, + } setFilters(resetFilters) setLogType(LOG_TYPE_ALL_VALUE) @@ -162,6 +177,7 @@ export function CommonLogsFilterBar( page: 1, pageSize: 100, type: [LOG_TYPE_ALL_VALUE], + quotaFilter: undefined, startTime: start.getTime(), endTime: end.getTime(), searchVersion: undefined, @@ -184,8 +200,15 @@ export function CommonLogsFilterBar( !!filters.upstreamRequestId const hasTypeFilter = logType !== LOG_TYPE_ALL_VALUE + const hasQuotaFilter = + filters.quotaFilter != null && + filters.quotaFilter !== QUOTA_FILTER_ALL_VALUE const hasAdditionalFilters = - !!filters.model || !!filters.group || hasTypeFilter || hasExpandedFilters + !!filters.model || + !!filters.group || + hasTypeFilter || + hasQuotaFilter || + hasExpandedFilters const expandedFilterCount = [ filters.token, @@ -299,6 +322,16 @@ export function CommonLogsFilterBar( ) + const quotaFilter = ( + + handleChange('quotaFilter', value)} + /> + + ) const advancedFilters = ( <> @@ -360,6 +393,7 @@ export function CommonLogsFilterBar( {modelFilter} {groupFilter} {typeFilter} + {quotaFilter} } advancedFilters={advancedFilters} @@ -369,12 +403,14 @@ export function CommonLogsFilterBar( {modelFilter} {groupFilter} {typeFilter} + {quotaFilter} {advancedFilters} } mobileFilterCount={ - [filters.model, filters.group, hasTypeFilter].filter(Boolean).length + - expandedFilterCount + [filters.model, filters.group, hasTypeFilter, hasQuotaFilter].filter( + Boolean + ).length + expandedFilterCount } hasAdvancedActiveFilters={hasExpandedFilters} advancedFilterCount={expandedFilterCount} diff --git a/web/default/src/features/usage-logs/components/quota-filter-select.tsx b/web/default/src/features/usage-logs/components/quota-filter-select.tsx new file mode 100644 index 0000000..e79c93b --- /dev/null +++ b/web/default/src/features/usage-logs/components/quota-filter-select.tsx @@ -0,0 +1,86 @@ +/* +Copyright (C) 2023-2026 MAX-API-Next + +This program is free software: you can redistribute it and/or modify +it under the terms of the GNU Affero General Public License as +published by the Free Software Foundation, either version 3 of the +License, or (at your option) any later version. + +This program is distributed in the hope that it will be useful, +but WITHOUT ANY WARRANTY; without even the implied warranty of +MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +GNU Affero General Public License for more details. + +You should have received a copy of the GNU Affero General Public License +along with this program. If not, see . + +For commercial licensing, please contact https://github.com/MAX-API-Next/MAX-API/issues +*/ +import { useMemo } from 'react' +import { useTranslation } from 'react-i18next' +import { + Select, + SelectContent, + SelectGroup, + SelectItem, + SelectTrigger, + SelectValue, +} from '@/components/ui/select' +import { + QUOTA_FILTER_ALL_VALUE, + QUOTA_FILTER_VALUES, + QUOTA_FILTERS, +} from '../constants' + +export type QuotaFilterValue = (typeof QUOTA_FILTER_VALUES)[number] + +export function isQuotaFilterValue(value: string): value is QuotaFilterValue { + return (QUOTA_FILTER_VALUES as readonly string[]).includes(value) +} + +interface QuotaFilterSelectProps { + value: QuotaFilterValue + onValueChange: (value: QuotaFilterValue) => void +} + +export function QuotaFilterSelect(props: QuotaFilterSelectProps) { + const { t } = useTranslation() + const items = useMemo( + () => + QUOTA_FILTERS.map((filter) => ({ + value: filter.value, + label: t(filter.label), + })), + [t] + ) + const label = + items.find((filter) => filter.value === props.value)?.label ?? + t('All Billing') + + return ( + + ) +} diff --git a/web/default/src/features/usage-logs/components/task-logs-filter-bar.tsx b/web/default/src/features/usage-logs/components/task-logs-filter-bar.tsx index 8e79029..ad81ca6 100644 --- a/web/default/src/features/usage-logs/components/task-logs-filter-bar.tsx +++ b/web/default/src/features/usage-logs/components/task-logs-filter-bar.tsx @@ -22,6 +22,7 @@ import { useNavigate, getRouteApi } from '@tanstack/react-router' import { type Table } from '@tanstack/react-table' import { useTranslation } from 'react-i18next' import { useIsAdmin } from '@/hooks/use-admin' +import { QUOTA_FILTER_ALL_VALUE } from '../constants' import { buildSearchParams } from '../lib/filter' import { getDefaultTimeRange } from '../lib/utils' import type { DrawingLogFilters, LogCategory, TaskLogFilters } from '../types' @@ -31,6 +32,11 @@ import { LogsFilterInput, LogsFilterToolbar, } from './logs-filter-toolbar' +import { + isQuotaFilterValue, + QuotaFilterSelect, + type QuotaFilterValue, +} from './quota-filter-select' const route = getRouteApi('/_authenticated/usage-logs/$section') @@ -85,6 +91,10 @@ export function TaskLogsFilterBar(props: TaskLogsFilterBarProps) { ...(searchParams.channel ? { channel: String(searchParams.channel) } : {}), + quotaFilter: + searchParams.quotaFilter && isQuotaFilterValue(searchParams.quotaFilter) + ? searchParams.quotaFilter + : QUOTA_FILTER_ALL_VALUE, } const next: TaskLogsFilters = props.logCategory === 'drawing' @@ -105,6 +115,7 @@ export function TaskLogsFilterBar(props: TaskLogsFilterBarProps) { searchParams.endTime, searchParams.channel, searchParams.filter, + searchParams.quotaFilter, ]) const handleChange = useCallback( @@ -130,7 +141,11 @@ export function TaskLogsFilterBar(props: TaskLogsFilterBarProps) { const handleReset = useCallback(() => { const { start, end } = getDefaultTimeRange() - const resetFilters: TaskLogsFilters = { startTime: start, endTime: end } + const resetFilters: TaskLogsFilters = { + startTime: start, + endTime: end, + quotaFilter: QUOTA_FILTER_ALL_VALUE, + } setFilters(resetFilters) navigate({ @@ -139,6 +154,7 @@ export function TaskLogsFilterBar(props: TaskLogsFilterBarProps) { search: { page: 1, pageSize: 100, + quotaFilter: undefined, startTime: start.getTime(), endTime: end.getTime(), searchVersion: undefined, @@ -161,11 +177,15 @@ export function TaskLogsFilterBar(props: TaskLogsFilterBarProps) { ) const filterValue = getFilterValue(filters, props.logCategory) + const hasQuotaFilter = + filters.quotaFilter != null && + filters.quotaFilter !== QUOTA_FILTER_ALL_VALUE const placeholder = props.logCategory === 'drawing' ? t('Filter by Midjourney task ID') : t('Filter by task ID') - const hasAdditionalFilters = !!filterValue || !!filters.channel + const hasAdditionalFilters = + !!filterValue || !!filters.channel || hasQuotaFilter const dateRangeFilter = ( (props: TaskLogsFilterBarProps) { /> ) : null + const quotaFilter = ( + + handleChange('quotaFilter', value)} + /> + + ) return ( (props: TaskLogsFilterBarProps) { {dateRangeFilter} {taskIdFilter} {channelFilter} + {quotaFilter} } mobilePinnedFilters={dateRangeFilter} @@ -215,9 +246,12 @@ export function TaskLogsFilterBar(props: TaskLogsFilterBarProps) { <> {taskIdFilter} {channelFilter} + {quotaFilter} } - mobileFilterCount={[filterValue, filters.channel].filter(Boolean).length} + mobileFilterCount={ + [filterValue, filters.channel, hasQuotaFilter].filter(Boolean).length + } hasActiveFilters={hasAdditionalFilters} onSearch={handleApply} searchLoading={fetchingLogs > 0} diff --git a/web/default/src/features/usage-logs/constants.ts b/web/default/src/features/usage-logs/constants.ts index 54a8423..988bec8 100644 --- a/web/default/src/features/usage-logs/constants.ts +++ b/web/default/src/features/usage-logs/constants.ts @@ -137,6 +137,32 @@ export const LOG_TYPE_SEARCH_VALUES = [ LOG_TYPE_RETRY_VALUE, ] as [string, ...string[]] +export const QUOTA_FILTER_ALL_VALUE = 'all' as const +export const QUOTA_FILTER_ABNORMAL_VALUE = 'abnormal' as const +export const QUOTA_FILTER_ZERO_VALUE = 'zero' as const +export const QUOTA_FILTER_NEGATIVE_VALUE = 'negative' as const + +export const QUOTA_FILTERS = [ + { label: 'All Billing', value: QUOTA_FILTER_ALL_VALUE }, + { label: 'Abnormal Billing', value: QUOTA_FILTER_ABNORMAL_VALUE }, + { label: 'Zero Billing', value: QUOTA_FILTER_ZERO_VALUE }, + { label: 'Negative Billing', value: QUOTA_FILTER_NEGATIVE_VALUE }, +] as const + +export const QUOTA_FILTER_VALUES = [ + QUOTA_FILTER_ALL_VALUE, + QUOTA_FILTER_ABNORMAL_VALUE, + QUOTA_FILTER_ZERO_VALUE, + QUOTA_FILTER_NEGATIVE_VALUE, +] as const + +export const QUOTA_FILTER_SEARCH_VALUES = [ + QUOTA_FILTER_ALL_VALUE, + QUOTA_FILTER_ABNORMAL_VALUE, + QUOTA_FILTER_ZERO_VALUE, + QUOTA_FILTER_NEGATIVE_VALUE, +] satisfies [string, ...string[]] + // ============================================================================ // Drawing Logs (Midjourney) Constants // ============================================================================ diff --git a/web/default/src/features/usage-logs/lib/filter.ts b/web/default/src/features/usage-logs/lib/filter.ts index d63e42a..6ceaf18 100644 --- a/web/default/src/features/usage-logs/lib/filter.ts +++ b/web/default/src/features/usage-logs/lib/filter.ts @@ -43,6 +43,7 @@ export function buildSearchParams( ...(filters.startTime && { startTime: filters.startTime.getTime() }), ...(filters.endTime && { endTime: filters.endTime.getTime() }), ...(filters.channel && { channel: filters.channel }), + ...(filters.quotaFilter && { quotaFilter: filters.quotaFilter }), } switch (logCategory) { diff --git a/web/default/src/features/usage-logs/lib/utils.test.ts b/web/default/src/features/usage-logs/lib/utils.test.ts index 1cdb813..85a311c 100644 --- a/web/default/src/features/usage-logs/lib/utils.test.ts +++ b/web/default/src/features/usage-logs/lib/utils.test.ts @@ -18,7 +18,11 @@ For commercial licensing, please contact https://github.com/MAX-API-Next/MAX-API */ import assert from 'node:assert/strict' import { describe, test } from 'node:test' -import { buildApiParams, matchesCommonLogTypeFilter } from './utils' +import { + buildApiParams, + buildBaseParams, + matchesCommonLogTypeFilter, +} from './utils' describe('buildApiParams', () => { test('treats mixed retry and numeric type filters as retry filters', () => { @@ -102,6 +106,40 @@ describe('buildApiParams', () => { assert.equal(params.log_filter, undefined) } }) + + test('maps quota filters to API params', () => { + const params = buildApiParams({ + page: 1, + pageSize: 100, + searchParams: { quotaFilter: 'abnormal' }, + isAdmin: true, + }) + + assert.equal(params.quota_filter, 'abnormal') + }) + + test('does not send all quota filter to API', () => { + const params = buildApiParams({ + page: 1, + pageSize: 100, + searchParams: { quotaFilter: 'all' }, + isAdmin: true, + }) + + assert.equal(params.quota_filter, undefined) + }) +}) + +describe('buildBaseParams', () => { + test('maps quota filters for task-like logs', () => { + const params = buildBaseParams({ + page: 2, + pageSize: 50, + searchParams: { quotaFilter: 'negative' }, + }) + + assert.equal(params.quota_filter, 'negative') + }) }) describe('matchesCommonLogTypeFilter', () => { @@ -128,7 +166,10 @@ describe('matchesCommonLogTypeFilter', () => { ) assert.equal( matchesCommonLogTypeFilter( - { type: 2, other: JSON.stringify({ admin_info: { use_channel: ['1'] } }) }, + { + type: 2, + other: JSON.stringify({ admin_info: { use_channel: ['1'] } }), + }, ['retry'] ), false diff --git a/web/default/src/features/usage-logs/lib/utils.ts b/web/default/src/features/usage-logs/lib/utils.ts index 4ea1f41..9030af9 100644 --- a/web/default/src/features/usage-logs/lib/utils.ts +++ b/web/default/src/features/usage-logs/lib/utils.ts @@ -35,6 +35,7 @@ import { LOG_TYPE_RETRY_VALUE, LOG_TYPE_ERROR_RETRY_VALUE, LOG_TYPE_EMPTY_RETRY_VALUE, + QUOTA_FILTER_ALL_VALUE, } from '../constants' import type { GetLogsParams, @@ -214,6 +215,7 @@ export function buildBaseParams(config: { channel_id?: string start_timestamp?: number end_timestamp?: number + quota_filter?: string } { const { page, pageSize, searchParams, useMilliseconds = false } = config @@ -225,6 +227,10 @@ export function buildBaseParams(config: { channel_id: String(searchParams.channel), } : {}), + ...(searchParams.quotaFilter && + searchParams.quotaFilter !== QUOTA_FILTER_ALL_VALUE + ? { quota_filter: String(searchParams.quotaFilter) } + : {}), ...buildTimeRangeParams(searchParams, useMilliseconds), } } @@ -303,6 +309,10 @@ export function buildApiParams(config: { ...(searchParams.upstreamRequestId ? { upstream_request_id: String(searchParams.upstreamRequestId) } : {}), + ...(searchParams.quotaFilter && + searchParams.quotaFilter !== QUOTA_FILTER_ALL_VALUE + ? { quota_filter: String(searchParams.quotaFilter) } + : {}), ...buildTimeRangeParams(searchParams, false), } if (searchParams.type) { diff --git a/web/default/src/features/usage-logs/types.ts b/web/default/src/features/usage-logs/types.ts index 677cf76..3f7d49b 100644 --- a/web/default/src/features/usage-logs/types.ts +++ b/web/default/src/features/usage-logs/types.ts @@ -72,6 +72,7 @@ export interface CommonFilters { startTime?: Date endTime?: Date channel?: string + quotaFilter?: string } /** @@ -315,6 +316,7 @@ export interface GetLogsParams { group?: string request_id?: string upstream_request_id?: string + quota_filter?: string } export interface GetLogsResponse { @@ -346,6 +348,7 @@ export interface GetLogStatsParams { group?: string request_id?: string upstream_request_id?: string + quota_filter?: string } export interface GetLogStatsResponse { @@ -365,6 +368,7 @@ export interface GetMidjourneyLogsParams { mj_id?: string start_timestamp?: number end_timestamp?: number + quota_filter?: string } // ============================================================================ @@ -378,6 +382,7 @@ export interface GetTaskLogsParams { task_id?: string start_timestamp?: number end_timestamp?: number + quota_filter?: string } // ============================================================================ diff --git a/web/default/src/i18n/locales/en.json b/web/default/src/i18n/locales/en.json index 38ccfa2..ef4f8b9 100644 --- a/web/default/src/i18n/locales/en.json +++ b/web/default/src/i18n/locales/en.json @@ -113,6 +113,7 @@ "A live control layer for high-volume AI operations": "A live control layer for high-volume AI operations", "A production Agent may call multiple models, tools, knowledge bases, image or video tasks, and search systems inside one user intent.": "A production Agent may call multiple models, tools, knowledge bases, image or video tasks, and search systems inside one user intent.", "A technological nerve center for every AI request": "A technological nerve center for every AI request", + "Abnormal Billing": "Abnormal Billing", "About": "About", "About {{days}} days left": "About {{days}} days left", "Accept Unpriced Models": "Accept Unpriced Models", @@ -315,6 +316,7 @@ "Alibaba Cloud Bailian": "Alibaba Cloud Bailian", "Alipay": "Alipay", "All": "All", + "All Billing": "All Billing", "All categories": "All categories", "All conditions must match before this tier is used.": "All conditions must match before this tier is used.", "All edits are overwrite operations. Leave fields empty to keep current values unchanged.": "All edits are overwrite operations. Leave fields empty to keep current values unchanged.", @@ -2826,6 +2828,7 @@ "Needs API key": "Needs API key", "Negative": "Negative", "Negative Balance": "Negative Balance", + "Negative Billing": "Negative Billing", "Nested JSON defining per-group rules for adding (+:), removing (-:), or appending usable groups.": "Nested JSON defining per-group rules for adding (+:), removing (-:), or appending usable groups.", "Nested JSON: source group →": "Nested JSON: source group →", "Network proxy for this channel (supports socks5 protocol)": "Network proxy for this channel (supports socks5 protocol)", @@ -5166,6 +5169,7 @@ "Your Turnstile secret key": "Your Turnstile secret key", "Your Turnstile site key": "Your Turnstile site key", "Zero Balance": "Zero Balance", + "Zero Billing": "Zero Billing", "Zero retention": "Zero retention", "Zhipu": "Zhipu", "Zhipu V4": "Zhipu V4", diff --git a/web/default/src/i18n/locales/fr.json b/web/default/src/i18n/locales/fr.json index aa60bff..3e9bc86 100644 --- a/web/default/src/i18n/locales/fr.json +++ b/web/default/src/i18n/locales/fr.json @@ -113,6 +113,7 @@ "A live control layer for high-volume AI operations": "Une couche de contrôle en direct pour les opérations IA à fort volume", "A production Agent may call multiple models, tools, knowledge bases, image or video tasks, and search systems inside one user intent.": "Un Agent en production peut appeler plusieurs modèles, outils, bases de connaissances, tâches image ou vidéo et systèmes de recherche dans une seule intention utilisateur.", "A technological nerve center for every AI request": "Un centre nerveux technologique pour chaque requête IA", + "Abnormal Billing": "Facturation anormale", "About": "À propos", "About {{days}} days left": "Environ {{days}} jours restants", "Accept Unpriced Models": "Accepter les modèles non tarifés", @@ -315,6 +316,7 @@ "Alibaba Cloud Bailian": "Alibaba Cloud Bailian", "Alipay": "Alipay", "All": "Tout", + "All Billing": "Toutes les facturations", "All categories": "Toutes catégories", "All conditions must match before this tier is used.": "Toutes les conditions doivent correspondre avant que ce palier soit utilisé.", "All edits are overwrite operations. Leave fields empty to keep current values unchanged.": "Toutes les modifications sont des opérations d'écrasement. Laissez les champs vides pour conserver les valeurs actuelles inchangées.", @@ -2826,6 +2828,7 @@ "Needs API key": "Clé API requise", "Negative": "Négatif", "Negative Balance": "Solde négatif", + "Negative Billing": "Facturation négative", "Nested JSON defining per-group rules for adding (+:), removing (-:), or appending usable groups.": "JSON imbriqué définissant des règles par groupe pour ajouter (+:), supprimer (-:), ou ajouter des groupes utilisables.", "Nested JSON: source group →": "JSON imbriqué : groupe source →", "Network proxy for this channel (supports socks5 protocol)": "Proxy réseau pour ce canal (supporte le protocole socks5)", @@ -5166,6 +5169,7 @@ "Your Turnstile secret key": "Votre clé secrète Turnstile", "Your Turnstile site key": "Votre clé de site Turnstile", "Zero Balance": "Solde nul", + "Zero Billing": "Facturation à zéro", "Zero retention": "Aucune rétention", "Zhipu": "Zhipu", "Zhipu V4": "Zhipu V4", diff --git a/web/default/src/i18n/locales/ja.json b/web/default/src/i18n/locales/ja.json index 2aa7928..069843d 100644 --- a/web/default/src/i18n/locales/ja.json +++ b/web/default/src/i18n/locales/ja.json @@ -113,6 +113,7 @@ "A live control layer for high-volume AI operations": "大規模AI運用のためのライブ制御層", "A production Agent may call multiple models, tools, knowledge bases, image or video tasks, and search systems inside one user intent.": "本番環境の Agent は、1 つのユーザー意図の中で複数のモデル、ツール、ナレッジベース、画像や動画タスク、検索システムを呼び出すことがあります。", "A technological nerve center for every AI request": "すべてのAIリクエストのための技術的な中枢", + "Abnormal Billing": "異常課金", "About": "このサービスについて", "About {{days}} days left": "約 {{days}} 日分", "Accept Unpriced Models": "価格設定されていないモデルを許可", @@ -315,6 +316,7 @@ "Alibaba Cloud Bailian": "Alibaba Cloud Bailian", "Alipay": "Alipay", "All": "すべて", + "All Billing": "すべての課金", "All categories": "すべてのカテゴリ", "All conditions must match before this tier is used.": "この段階を使用するには、すべての条件に一致する必要があります。", "All edits are overwrite operations. Leave fields empty to keep current values unchanged.": "すべての編集は上書き操作です。現在の値を変更しないままにするには、フィールドを空のままにしてください。", @@ -2826,6 +2828,7 @@ "Needs API key": "API キーが必要", "Negative": "マイナス", "Negative Balance": "マイナス残高", + "Negative Billing": "負の課金", "Nested JSON defining per-group rules for adding (+:), removing (-:), or appending usable groups.": "追加 (+:)、削除 (-:)、または使用可能なグループの追加を行うグループごとのルールを定義するネストされたJSON。", "Nested JSON: source group →": "ネストされたJSON: ソースグループ →", "Network proxy for this channel (supports socks5 protocol)": "このチャネルのネットワークプロキシ (socks5プロトコルをサポート)", @@ -5166,6 +5169,7 @@ "Your Turnstile secret key": "あなたのTurnstileシークレットキー", "Your Turnstile site key": "あなたのTurnstileサイトキー", "Zero Balance": "ゼロ残高", + "Zero Billing": "ゼロ課金", "Zero retention": "データ保持なし", "Zhipu": "Zhipu", "Zhipu V4": "Zhipu V 4", diff --git a/web/default/src/i18n/locales/ru.json b/web/default/src/i18n/locales/ru.json index 2764b5a..8ce11da 100644 --- a/web/default/src/i18n/locales/ru.json +++ b/web/default/src/i18n/locales/ru.json @@ -113,6 +113,7 @@ "A live control layer for high-volume AI operations": "Рабочий слой контроля для высоконагруженных AI-операций", "A production Agent may call multiple models, tools, knowledge bases, image or video tasks, and search systems inside one user intent.": "Производственный Agent может в рамках одного намерения пользователя вызывать несколько моделей, инструментов, баз знаний, задач изображений или видео и поисковых систем.", "A technological nerve center for every AI request": "Технологический нервный центр для каждого AI-запроса", + "Abnormal Billing": "Аномальная тарификация", "About": "О проекте", "About {{days}} days left": "Примерно {{days}} дней", "Accept Unpriced Models": "Принимать модели без цены", @@ -315,6 +316,7 @@ "Alibaba Cloud Bailian": "Alibaba Cloud Bailian", "Alipay": "Alipay", "All": "Все", + "All Billing": "Вся тарификация", "All categories": "Все категории", "All conditions must match before this tier is used.": "Все условия должны совпасть, прежде чем будет использован этот уровень.", "All edits are overwrite operations. Leave fields empty to keep current values unchanged.": "Все изменения являются операциями перезаписи. Оставьте поля пустыми, чтобы сохранить текущие значения без изменений.", @@ -2826,6 +2828,7 @@ "Needs API key": "Нужен API-ключ", "Negative": "Отрицательный", "Negative Balance": "Отрицательный баланс", + "Negative Billing": "Отрицательная тарификация", "Nested JSON defining per-group rules for adding (+:), removing (-:), or appending usable groups.": "Вложенный JSON, определяющий правила для каждой группы для добавления (+:), удаления (-:) или добавления используемых групп.", "Nested JSON: source group →": "Вложенный JSON: исходная группа →", "Network proxy for this channel (supports socks5 protocol)": "Сетевой прокси для этого канала (поддерживает протокол socks5)", @@ -5166,6 +5169,7 @@ "Your Turnstile secret key": "Секретный ключ Turnstile", "Your Turnstile site key": "Ключ сайта Turnstile", "Zero Balance": "Нулевой баланс", + "Zero Billing": "Нулевая тарификация", "Zero retention": "Без хранения данных", "Zhipu": "Zhipu", "Zhipu V4": "Zhipu V4", diff --git a/web/default/src/i18n/locales/vi.json b/web/default/src/i18n/locales/vi.json index c41eef6..b11e6d1 100644 --- a/web/default/src/i18n/locales/vi.json +++ b/web/default/src/i18n/locales/vi.json @@ -113,6 +113,7 @@ "A live control layer for high-volume AI operations": "Lớp điều khiển trực tiếp cho vận hành AI lưu lượng lớn", "A production Agent may call multiple models, tools, knowledge bases, image or video tasks, and search systems inside one user intent.": "Một Agent trong production có thể gọi nhiều mô hình, công cụ, kho tri thức, tác vụ hình ảnh hoặc video và hệ thống tìm kiếm trong cùng một ý định người dùng.", "A technological nerve center for every AI request": "Trung tâm thần kinh công nghệ cho mọi yêu cầu AI", + "Abnormal Billing": "Tính phí bất thường", "About": "Giới thiệu", "About {{days}} days left": "Còn khoảng {{days}} ngày", "Accept Unpriced Models": "Chấp nhận các Mô hình chưa định giá", @@ -315,6 +316,7 @@ "Alibaba Cloud Bailian": "Alibaba Cloud Bailian", "Alipay": "Alipay", "All": "All", + "All Billing": "Tất cả tính phí", "All categories": "Tất cả danh mục", "All conditions must match before this tier is used.": "Tất cả điều kiện phải khớp trước khi tầng này được sử dụng.", "All edits are overwrite operations. Leave fields empty to keep current values unchanged.": "Tất cả các chỉnh sửa đều là thao tác ghi đè. Để trống các trường để giữ nguyên giá trị hiện tại.", @@ -2826,6 +2828,7 @@ "Needs API key": "Cần khóa API", "Negative": "Âm", "Negative Balance": "Số dư âm", + "Negative Billing": "Tính phí âm", "Nested JSON defining per-group rules for adding (+:), removing (-:), or appending usable groups.": "JSON lồng nhau xác định quy tắc theo nhóm để thêm (+:), xóa (-:), hoặc nối các nhóm có thể sử dụng.", "Nested JSON: source group →": "JSON lồng nhau: nhóm nguồn →", "Network proxy for this channel (supports socks5 protocol)": "Proxy mạng cho kênh này (hỗ trợ giao thức socks5)", @@ -5166,6 +5169,7 @@ "Your Turnstile secret key": "Khóa bí mật Turnstile của bạn", "Your Turnstile site key": "Khóa site Turnstile của bạn", "Zero Balance": "Số dư bằng không", + "Zero Billing": "Tính phí bằng 0", "Zero retention": "Không lưu dữ liệu", "Zhipu": "Zhipu", "Zhipu V4": "Zhipu V4", diff --git a/web/default/src/i18n/locales/zh.json b/web/default/src/i18n/locales/zh.json index 6f245c0..31f9253 100644 --- a/web/default/src/i18n/locales/zh.json +++ b/web/default/src/i18n/locales/zh.json @@ -113,6 +113,7 @@ "A live control layer for high-volume AI operations": "面向高流量 AI 运营的实时控制层", "A production Agent may call multiple models, tools, knowledge bases, image or video tasks, and search systems inside one user intent.": "生产环境中的 Agent 可能会在一次用户意图中连续调用多个模型、工具、知识库、图像任务、视频任务和搜索系统。", "A technological nerve center for every AI request": "每一次 AI 请求的技术神经中枢", + "Abnormal Billing": "异常计费", "About": "关于", "About {{days}} days left": "约剩 {{days}} 天", "Accept Unpriced Models": "接受未定价模型", @@ -315,6 +316,7 @@ "Alibaba Cloud Bailian": "阿里云百炼 / 通义千问", "Alipay": "支付宝", "All": "全部", + "All Billing": "全部计费", "All categories": "全部分类", "All conditions must match before this tier is used.": "所有条件都匹配后才会使用此阶梯。", "All edits are overwrite operations. Leave fields empty to keep current values unchanged.": "所有编辑都是覆盖操作。留空字段将保持当前值不变。", @@ -2826,6 +2828,7 @@ "Needs API key": "需要 API 密钥", "Negative": "负余额", "Negative Balance": "负余额", + "Negative Billing": "负计费", "Nested JSON defining per-group rules for adding (+:), removing (-:), or appending usable groups.": "嵌套 JSON,定义按分组添加(+:)、移除(-:)或追加可用分组的规则。", "Nested JSON: source group →": "嵌套 JSON:源分组 →", "Network proxy for this channel (supports socks5 protocol)": "此渠道的网络代理(支持 socks5 协议)", @@ -5166,6 +5169,7 @@ "Your Turnstile secret key": "您的 Turnstile 密钥", "Your Turnstile site key": "您的 Turnstile 站点密钥", "Zero Balance": "零余额", + "Zero Billing": "0 计费", "Zero retention": "零数据保留", "Zhipu": "智谱", "Zhipu V4": "智谱 V4", diff --git a/web/default/src/routes/_authenticated/usage-logs/$section.tsx b/web/default/src/routes/_authenticated/usage-logs/$section.tsx index 726671c..ddabb29 100644 --- a/web/default/src/routes/_authenticated/usage-logs/$section.tsx +++ b/web/default/src/routes/_authenticated/usage-logs/$section.tsx @@ -19,7 +19,10 @@ For commercial licensing, please contact https://github.com/MAX-API-Next/MAX-API import z from 'zod' import { createFileRoute, redirect } from '@tanstack/react-router' import { UsageLogs } from '@/features/usage-logs' -import { LOG_TYPE_SEARCH_VALUES } from '@/features/usage-logs/constants' +import { + LOG_TYPE_SEARCH_VALUES, + QUOTA_FILTER_SEARCH_VALUES, +} from '@/features/usage-logs/constants' import { isUsageLogsSectionId, USAGE_LOGS_DEFAULT_SECTION, @@ -48,6 +51,7 @@ const usageLogsSearchSchema = z.object({ username: z.string().optional().catch(''), requestId: z.string().optional().catch(''), upstreamRequestId: z.string().optional().catch(''), + quotaFilter: z.enum(QUOTA_FILTER_SEARCH_VALUES).optional().catch(undefined), startTime: z.number().optional(), endTime: z.number().optional(), }) From 9fd1a8e9b7bdcdca7c41b17204e4a731aa90574f Mon Sep 17 00:00:00 2001 From: CSCITech Date: Tue, 7 Jul 2026 15:09:20 +0800 Subject: [PATCH 4/6] v1.0.4-preview.2 --- controller/discord.go | 2 +- controller/github.go | 2 +- controller/linuxdo.go | 2 +- controller/log.go | 59 ++++++- controller/oidc.go | 2 +- controller/telegram.go | 2 +- controller/user.go | 24 ++- controller/wechat.go | 2 +- dto/openai_image.go | 11 ++ dto/openai_image_test.go | 21 +++ go.mod | 2 +- go.sum | 4 +- model/log.go | 150 +++++++++------- model/log_test.go | 51 +++--- model/task_cas_test.go | 79 +++++++++ model/user.go | 171 +++++++++++++++---- model/user_update_test.go | 30 +++- relay/channel/ali/image.go | 4 +- relay/channel/task/ali/adaptor.go | 32 +--- relay/channel/task/ali/adaptor_kling_test.go | 52 ++++++ relay/channel/task/ali/adaptor_wan27_test.go | 48 ++++++ relay/common/relay_info.go | 26 +++ relay/common/relay_info_test.go | 31 ++++ relay/common/relay_utils.go | 15 +- relay/common/relay_utils_test.go | 36 ++-- relay/helper/common.go | 12 ++ relay/helper/common_test.go | 22 +++ relay/helper/valid_request.go | 13 +- setting/ratio_setting/group_ratio.go | 10 +- setting/ratio_setting/group_ratio_test.go | 2 + 30 files changed, 725 insertions(+), 192 deletions(-) create mode 100644 dto/openai_image_test.go create mode 100644 relay/helper/common_test.go 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 fffc414..ff550fd 100644 --- a/controller/log.go +++ b/controller/log.go @@ -26,7 +26,22 @@ func GetAllLogs(c *gin.Context) { requestId := c.Query("request_id") upstreamRequestId := c.Query("upstream_request_id") quotaFilter := c.Query("quota_filter") - logs, total, err := model.GetAllLogs(logType, logFilter, startTimestamp, endTimestamp, modelName, username, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), channel, group, requestId, upstreamRequestId, quotaFilter) + 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 @@ -79,7 +94,21 @@ func GetUserLogs(c *gin.Context) { requestId := c.Query("request_id") upstreamRequestId := c.Query("upstream_request_id") quotaFilter := c.Query("quota_filter") - logs, total, err := model.GetUserLogs(userId, logType, logFilter, startTimestamp, endTimestamp, modelName, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), group, requestId, upstreamRequestId, quotaFilter) + 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 @@ -168,7 +197,18 @@ func GetLogsStat(c *gin.Context) { channel, _ := strconv.Atoi(c.Query("channel")) group := c.Query("group") quotaFilter := c.Query("quota_filter") - stat, err := model.SumUsedQuota(logType, logFilter, startTimestamp, endTimestamp, modelName, username, tokenName, channel, group, quotaFilter) + 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 @@ -197,7 +237,18 @@ func GetLogsSelfStat(c *gin.Context) { channel, _ := strconv.Atoi(c.Query("channel")) group := c.Query("group") quotaFilter := c.Query("quota_filter") - quotaNum, err := model.SumUsedQuota(logType, logFilter, startTimestamp, endTimestamp, modelName, username, tokenName, channel, group, quotaFilter) + 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/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/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/user.go b/controller/user.go index 3b692be..2388f1a 100644 --- a/controller/user.go +++ b/controller/user.go @@ -388,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 } @@ -438,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(), @@ -799,10 +799,17 @@ 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 = "" @@ -820,7 +827,7 @@ func UpdateSelf(c *gin.Context) { 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 } @@ -1082,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 } 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 7930a68..9b81c07 100644 --- a/dto/openai_image.go +++ b/dto/openai_image.go @@ -2,6 +2,7 @@ package dto import ( "encoding/json" + "fmt" "reflect" "strings" @@ -14,6 +15,16 @@ import ( // 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 1 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..3339721 --- /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 1 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 1 and %d", MaxImageN), err.Error()) +} diff --git a/go.mod b/go.mod index eed20e9..25ed01a 100644 --- a/go.mod +++ b/go.mod @@ -49,7 +49,7 @@ 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.51.0 + golang.org/x/crypto v0.52.0 golang.org/x/image v0.41.0 golang.org/x/net v0.55.0 golang.org/x/sync v0.20.0 diff --git a/go.sum b/go.sum index a0bc3a5..00cc117 100644 --- a/go.sum +++ b/go.sum @@ -327,8 +327,8 @@ 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.51.0 h1:IBPXwPfKxY7cWQZ38ZCIRPI50YLeevDLlLnyC5wRGTI= -golang.org/x/crypto v0.51.0/go.mod h1:8AdwkbraGNABw2kOX6YFPs3WM22XqI4EXEd8g+x7Oc8= +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.41.0 h1:8wS72eGJMJaBxK6okTzd4WaXumUlTVlb753MlsSvTCo= diff --git a/model/log.go b/model/log.go index 19f2bec..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"` @@ -97,6 +97,24 @@ const ( 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: @@ -636,45 +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, quotaFilter 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", quotaFilter) + 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 } @@ -725,46 +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, quotaFilter 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", quotaFilter) + 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 } @@ -805,57 +823,57 @@ 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, quotaFilter 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", quotaFilter) - rpmTpmQuery = applyQuotaFilter(rpmTpmQuery, "logs.quota", quotaFilter) + 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) + 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 c1a83e2..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,7 +74,7 @@ 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) @@ -93,7 +100,7 @@ func TestSumUsedQuotaRetryFilterIgnoresNonConsumeQuota(t *testing.T) { }), }).Error) - 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) @@ -108,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) @@ -129,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, 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) @@ -150,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) @@ -184,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) @@ -235,19 +246,19 @@ func TestLogQuotaFilters(t *testing.T) { } require.NoError(t, LOG_DB.Create(&logs).Error) - zeroLogs, total, err := GetAllLogs(LogTypeConsume, "", 0, 0, "", "", "", 0, 10, 0, "", "", "", LogQuotaFilterZero) + 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(1, LogTypeConsume, "", 0, 0, "", "", 0, 10, "", "", "", LogQuotaFilterNegative) + 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(LogTypeUnknown, "", 0, 0, "", "", "", 0, "", LogQuotaFilterAbnormal) + abnormalStat, err := SumUsedQuota(LogQueryParams{LogType: LogTypeUnknown, QuotaFilter: LogQuotaFilterAbnormal}) require.NoError(t, err) require.Equal(t, -125, abnormalStat.Quota) require.Equal(t, 3, abnormalStat.Rpm) @@ -289,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) @@ -338,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) @@ -377,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) @@ -397,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/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/user.go b/model/user.go index 56679a2..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 { @@ -667,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)) } } @@ -723,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)) } } @@ -744,7 +771,17 @@ func (user *User) FinalizeOAuthUserCreation(inviterId int) { } func (user *User) Update(updatePassword bool) error { - if err := user.UpdateWithTx(DB, updatePassword); err != nil { + 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 { @@ -754,6 +791,14 @@ func (user *User) Update(updatePassword bool) error { } 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) @@ -766,93 +811,151 @@ func (user *User) UpdateWithTx(tx *gorm.DB, updatePassword bool) error { if err = tx.First(¤t, user.Id).Error; err != nil { return err } - result := tx.Model(¤t).Updates(buildUserUpdateValues(current, newUser, updatePassword)) + result := tx.Model(¤t).Updates(buildUserUpdateValues(current, newUser, updatePassword, fields...)) if err = ensureUserUpdateMatchedTx(tx, result, user.Id, errors.New("用户不存在")); err != nil { return err } return tx.First(user, user.Id).Error } -func buildUserUpdateValues(current User, newUser User, updatePassword bool) map[string]interface{} { - fullUser := newUser.CreatedAt != 0 +func buildUserUpdateValues(current User, newUser User, updatePassword bool, fields ...UserUpdateField) map[string]interface{} { updates := map[string]interface{}{} - if fullUser || newUser.Username != "" { + if len(fields) > 0 { + for _, field := range fields { + applyUserUpdateField(updates, newUser, field) + } + } else { + applyNonZeroUserUpdateValues(updates, newUser) + } + if updatePassword { + updates["password"] = newUser.Password + } + + copyUnspecifiedUserUpdateValues(updates, current) + return updates +} + +func applyNonZeroUserUpdateValues(updates map[string]interface{}, newUser User) { + if newUser.Username != "" { updates["username"] = newUser.Username } - if fullUser || newUser.DisplayName != "" { + if newUser.DisplayName != "" { updates["display_name"] = newUser.DisplayName } - if fullUser || newUser.Role != 0 { + if newUser.Role != 0 { updates["role"] = newUser.Role } - if fullUser || newUser.Status != 0 { + if newUser.Status != 0 { updates["status"] = newUser.Status } - if fullUser || newUser.Email != "" { + if newUser.Email != "" { email := NormalizeEmail(newUser.Email) updates["email"] = email updates["normalized_email"] = email } - if fullUser || newUser.GitHubId != "" { + if newUser.GitHubId != "" { updates["github_id"] = newUser.GitHubId } - if fullUser || newUser.DiscordId != "" { + if newUser.DiscordId != "" { updates["discord_id"] = newUser.DiscordId } - if fullUser || newUser.OidcId != "" { + if newUser.OidcId != "" { updates["oidc_id"] = newUser.OidcId } - if fullUser || newUser.WeChatId != "" { + if newUser.WeChatId != "" { updates["wechat_id"] = newUser.WeChatId } - if fullUser || newUser.TelegramId != "" { + if newUser.TelegramId != "" { updates["telegram_id"] = newUser.TelegramId } - if fullUser || newUser.AccessToken != nil { + if newUser.AccessToken != nil { updates["access_token"] = newUser.AccessToken } - if fullUser || newUser.Group != "" { + if newUser.Group != "" { updates["group"] = newUser.Group } - if fullUser || newUser.AffCode != "" { + if newUser.AffCode != "" { updates["aff_code"] = newUser.AffCode } - if fullUser || newUser.AffCount != 0 { + if newUser.AffCount != 0 { updates["aff_count"] = newUser.AffCount } - if fullUser || newUser.AffQuota != 0 { + if newUser.AffQuota != 0 { updates["aff_quota"] = newUser.AffQuota } - if fullUser || newUser.AffHistoryQuota != 0 { + if newUser.AffHistoryQuota != 0 { updates["aff_history"] = newUser.AffHistoryQuota } - if fullUser || newUser.InviterId != 0 { + if newUser.InviterId != 0 { updates["inviter_id"] = newUser.InviterId } - if fullUser || newUser.LinuxDOId != "" { + if newUser.LinuxDOId != "" { updates["linux_do_id"] = newUser.LinuxDOId } - if fullUser || newUser.Setting != "" { + if newUser.Setting != "" { updates["setting"] = newUser.Setting } - if fullUser || newUser.Remark != "" { + if newUser.Remark != "" { updates["remark"] = newUser.Remark } - if fullUser || newUser.StripeCustomer != "" { + if newUser.StripeCustomer != "" { updates["stripe_customer"] = newUser.StripeCustomer } - if fullUser || newUser.LastLoginAt != 0 { + if newUser.LastLoginAt != 0 { updates["last_login_at"] = newUser.LastLoginAt } - if updatePassword { - updates["password"] = newUser.Password - } +} - if !fullUser { - copyUnspecifiedUserUpdateValues(updates, current) +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 } - return updates } func copyUnspecifiedUserUpdateValues(updates map[string]interface{}, current User) { diff --git a/model/user_update_test.go b/model/user_update_test.go index f55a345..4bc68c5 100644 --- a/model/user_update_test.go +++ b/model/user_update_test.go @@ -192,7 +192,7 @@ func TestUserUpdatePersistsZeroValueProfileFields(t *testing.T) { loaded.UsedQuota = 1 loaded.RequestCount = 1 - require.NoError(t, loaded.Update(false)) + require.NoError(t, loaded.UpdateFields(false, UserUpdateFieldDisplayName, UserUpdateFieldAffCount)) var got User require.NoError(t, DB.First(&got, user.Id).Error) @@ -203,6 +203,34 @@ func TestUserUpdatePersistsZeroValueProfileFields(t *testing.T) { 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) diff --git a/relay/channel/ali/image.go b/relay/channel/ali/image.go index 2375548..f681c82 100644 --- a/relay/channel/ali/image.go +++ b/relay/channel/ali/image.go @@ -54,8 +54,8 @@ func oaiImage2AliImageRequest(info *relaycommon.RelayInfo, request dto.ImageRequ } } - if imageRequest.Parameters.N < 0 || imageRequest.Parameters.N > dto.MaxImageN { - return nil, fmt.Errorf("parameters.n must be an integer between 1 and %d", dto.MaxImageN) + 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/task/ali/adaptor.go b/relay/channel/task/ali/adaptor.go index 5432ef7..15438ca 100644 --- a/relay/channel/task/ali/adaptor.go +++ b/relay/channel/task/ali/adaptor.go @@ -510,20 +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 - } - } - if aliReq.Parameters.Duration <= 0 { - 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 { @@ -557,18 +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 - } - if aliReq.Parameters.Duration <= 0 { - aliReq.Parameters.Duration = 5 + 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{}{ { 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/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 e1d66d9..1cd1cd0 100644 --- a/relay/common/relay_utils.go +++ b/relay/common/relay_utils.go @@ -117,12 +117,12 @@ func validatePrompt(prompt string) *dto.TaskError { const MaxTaskDurationSeconds = 3600 func validateTaskDurationBounds(req TaskSubmitReq) *dto.TaskError { - seconds := req.DurationValue() - if seconds == 0 && req.Seconds != "" { - seconds, _ = strconv.Atoi(req.Seconds) + 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 1 and %d", MaxTaskDurationSeconds), "invalid_seconds", http.StatusBadRequest, true) + return createTaskError(fmt.Errorf("seconds must be between 0 and %d", MaxTaskDurationSeconds), "invalid_seconds", http.StatusBadRequest, true) } return nil } @@ -183,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) != "" { diff --git a/relay/common/relay_utils_test.go b/relay/common/relay_utils_test.go index e3934ed..7db18de 100644 --- a/relay/common/relay_utils_test.go +++ b/relay/common/relay_utils_test.go @@ -65,24 +65,34 @@ func TestTaskDurationBounds(t *testing.T) { } tests := []struct { - name string - body string - wantErr bool + name string + body string + wantErr bool + wantMessage string }{ { - name: "huge duration is rejected", - body: `{"model":"sora-2","prompt":"a cat","duration":9999999999}`, - wantErr: true, + 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, + 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, + 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", @@ -97,6 +107,7 @@ func TestTaskDurationBounds(t *testing.T) { 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) @@ -108,6 +119,7 @@ func TestTaskDurationBounds(t *testing.T) { 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) diff --git a/relay/helper/common.go b/relay/helper/common.go index f0da06a..72fb876 100644 --- a/relay/helper/common.go +++ b/relay/helper/common.go @@ -59,6 +59,10 @@ 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 } @@ -75,6 +79,10 @@ 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 } @@ -85,6 +93,10 @@ func ClaudeChunkData(c *gin.Context, resp dto.ClaudeResponse, 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()) } 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/valid_request.go b/relay/helper/valid_request.go index ffeb7fe..f0da5d4 100644 --- a/relay/helper/valid_request.go +++ b/relay/helper/valid_request.go @@ -168,8 +168,11 @@ func GetAndValidOpenAIImageRequest(c *gin.Context, relayMode int) (*dto.ImageReq imageRequest.Model = formData.Get("model") if nValue := strings.TrimSpace(formData.Get("n")); nValue != "" { n, err := strconv.Atoi(nValue) - if err != nil || n < 0 || n > dto.MaxImageN { - return nil, fmt.Errorf("n must be an integer between 1 and %d", dto.MaxImageN) + 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)) } @@ -211,8 +214,10 @@ 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 && *imageRequest.N > dto.MaxImageN { - return nil, fmt.Errorf("n must be an integer between 1 and %d", dto.MaxImageN) + if imageRequest.N != nil { + if err := dto.ValidateImageN("n", int(*imageRequest.N)); err != nil { + return nil, err + } } // Not "256x256", "512x512", or "1024x1024" diff --git a/setting/ratio_setting/group_ratio.go b/setting/ratio_setting/group_ratio.go index e9d80b9..dc1d1f3 100644 --- a/setting/ratio_setting/group_ratio.go +++ b/setting/ratio_setting/group_ratio.go @@ -74,10 +74,11 @@ func GetGroupRatioCopy() map[string]float64 { } func ContainsGroupRatio(name string) bool { - if isReservedAutoRouteGroupName(strings.TrimSpace(name)) { + trimmedName := strings.TrimSpace(name) + if isReservedAutoRouteGroupName(trimmedName) { return false } - _, ok := groupRatioMap.Get(name) + _, ok := groupRatioMap.Get(trimmedName) return ok } @@ -98,11 +99,12 @@ func UpdateGroupRatioByJSONString(jsonStr string) error { } func GetGroupRatio(name string) float64 { - if isReservedAutoRouteGroupName(strings.TrimSpace(name)) { + trimmedName := strings.TrimSpace(name) + if isReservedAutoRouteGroupName(trimmedName) { common.SysLog("group ratio not found: " + name) return 1 } - ratio, ok := groupRatioMap.Get(name) + ratio, ok := groupRatioMap.Get(trimmedName) if !ok { common.SysLog("group ratio not found: " + name) return 1 diff --git a/setting/ratio_setting/group_ratio_test.go b/setting/ratio_setting/group_ratio_test.go index fb8a476..ecf2894 100644 --- a/setting/ratio_setting/group_ratio_test.go +++ b/setting/ratio_setting/group_ratio_test.go @@ -44,5 +44,7 @@ func TestGroupRatioTrimsNormalGroupNames(t *testing.T) { 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 ")) } From fba173551e5dee83c7327c2ce422acd53c51b347 Mon Sep 17 00:00:00 2001 From: CSCITech Date: Tue, 7 Jul 2026 15:11:46 +0800 Subject: [PATCH 5/6] v1.0.4-preview.2 --- web/default/src/features/usage-logs/constants.ts | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/web/default/src/features/usage-logs/constants.ts b/web/default/src/features/usage-logs/constants.ts index 988bec8..087e0e5 100644 --- a/web/default/src/features/usage-logs/constants.ts +++ b/web/default/src/features/usage-logs/constants.ts @@ -157,10 +157,7 @@ export const QUOTA_FILTER_VALUES = [ ] as const export const QUOTA_FILTER_SEARCH_VALUES = [ - QUOTA_FILTER_ALL_VALUE, - QUOTA_FILTER_ABNORMAL_VALUE, - QUOTA_FILTER_ZERO_VALUE, - QUOTA_FILTER_NEGATIVE_VALUE, + ...QUOTA_FILTER_VALUES, ] satisfies [string, ...string[]] // ============================================================================ From f26d2f6135af94f7791ca616390f891719b1a165 Mon Sep 17 00:00:00 2001 From: CSCITech Date: Tue, 7 Jul 2026 15:29:50 +0800 Subject: [PATCH 6/6] v1.0.4-preview.2 --- dto/openai_image.go | 2 +- dto/openai_image_test.go | 4 ++-- go.mod | 6 +++--- go.sum | 20 ++++++++++---------- relay/helper/openai_image_request_test.go | 2 +- 5 files changed, 17 insertions(+), 17 deletions(-) diff --git a/dto/openai_image.go b/dto/openai_image.go index 9b81c07..0ff5b2b 100644 --- a/dto/openai_image.go +++ b/dto/openai_image.go @@ -20,7 +20,7 @@ func ValidateImageN(field string, n int) error { field = "n" } if n < 0 || n > MaxImageN { - return fmt.Errorf("%s must be an integer between 1 and %d", field, MaxImageN) + return fmt.Errorf("%s must be an integer between 0 and %d", field, MaxImageN) } return nil } diff --git a/dto/openai_image_test.go b/dto/openai_image_test.go index 3339721..663d409 100644 --- a/dto/openai_image_test.go +++ b/dto/openai_image_test.go @@ -13,9 +13,9 @@ func TestValidateImageN(t *testing.T) { err := ValidateImageN("", -1) require.Error(t, err) - require.Equal(t, fmt.Sprintf("n must be an integer between 1 and %d", MaxImageN), err.Error()) + 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 1 and %d", MaxImageN), err.Error()) + 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 25ed01a..f33a45d 100644 --- a/go.mod +++ b/go.mod @@ -50,11 +50,11 @@ require ( github.com/waffo-com/waffo-go v1.3.1 github.com/yapingcat/gomedia v0.0.0-20240906162731-17feea57090c golang.org/x/crypto v0.52.0 - golang.org/x/image v0.41.0 + golang.org/x/image v0.43.0 golang.org/x/net v0.55.0 - golang.org/x/sync v0.20.0 + golang.org/x/sync v0.21.0 golang.org/x/sys v0.45.0 - golang.org/x/text v0.37.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 00cc117..ff5e8e6 100644 --- a/go.sum +++ b/go.sum @@ -331,16 +331,16 @@ 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.41.0 h1:8wS72eGJMJaBxK6okTzd4WaXumUlTVlb753MlsSvTCo= -golang.org/x/image v0.41.0/go.mod h1:uIc348UZMSvS5Z65CVZ7iDPaNobNFEPeJ4kbqTOszmA= -golang.org/x/mod v0.35.0 h1:Ww1D637e6Pg+Zb2KrWfHQUnH2dQRLBQyAtpr/haaJeM= -golang.org/x/mod v0.35.0/go.mod h1:+GwiRhIInF8wPm+4AoT6L0FA1QWAad3OMdTRx4tFYlU= +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.55.0 h1:bcvxaJn3e1U6InsFWt1JUq1aSjnRxLzT2rtD2KfkDF8= golang.org/x/net v0.55.0/go.mod h1:L5U2KuzuOe1lY7Z+aWVIKK6qEeJXnXV9yzGA+WCHJww= -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/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= @@ -359,11 +359,11 @@ golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuX 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.37.0 h1:Cqjiwd9eSg8e0QAkyCaQTNHFIIzWtidPahFWR83rTrc= -golang.org/x/text v0.37.0/go.mod h1:a5sjxXGs9hsn/AJVwuElvCAo9v8QYLzvavO5z2PiM38= +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.44.0 h1:UP4ajHPIcuMjT1GqzDWRlalUEoY+uzoZKnhOjbIPD2c= -golang.org/x/tools v0.44.0/go.mod h1:KA0AfVErSdxRZIsOVipbv3rQhVXTnlU6UhKxHd1seDI= +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/relay/helper/openai_image_request_test.go b/relay/helper/openai_image_request_test.go index 7dc2e1b..f88a48e 100644 --- a/relay/helper/openai_image_request_test.go +++ b/relay/helper/openai_image_request_test.go @@ -24,7 +24,7 @@ func TestGetAndValidOpenAIImageRequestNBounds(t *testing.T) { return c } - boundErr := fmt.Sprintf("n must be an integer between 1 and %d", dto.MaxImageN) + boundErr := fmt.Sprintf("n must be an integer between 0 and %d", dto.MaxImageN) tests := []struct { name string