diff --git a/controller/channel-test.go b/controller/channel-test.go index b294979d5877..fee7b02b3409 100644 --- a/controller/channel-test.go +++ b/controller/channel-test.go @@ -438,7 +438,7 @@ func testChannel(ctx context.Context, channel *model.Channel, testUserID int, te if resp != nil { httpResp = resp.(*http.Response) if httpResp.StatusCode != http.StatusOK { - err := service.RelayErrorHandler(c.Request.Context(), httpResp, true) + upstreamErr := service.RelayErrorHandler(c.Request.Context(), httpResp, true) common.SysError(fmt.Sprintf( "channel test bad response: channel_id=%d name=%s type=%d model=%s endpoint_type=%s status=%d err=%v", channel.Id, @@ -447,12 +447,13 @@ func testChannel(ctx context.Context, channel *model.Channel, testUserID int, te testModel, endpointType, httpResp.StatusCode, - err, + upstreamErr, )) + // newAPIError 保留上游真实状态码,供自动禁用规则(状态码/关键词)判断 return testResult{ context: c, - localErr: err, - newAPIError: types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError), + localErr: upstreamErr, + newAPIError: upstreamErr, } } } @@ -867,6 +868,10 @@ func TestChannel(c *gin.Context) { requestCtx = c.Request.Context() } result := testChannel(requestCtx, channel, testUserID, testModel, endpointType, isStream) + // 手动测试失败时与正常转发、定时巡检走同一套错误处理,满足条件即禁用/冷却命中的密钥。 + // 必须放在 localErr 分支之前:上游返回错误时 localErr 与 newAPIError 同时非空, + // 提前 return 会让错误处理永远执行不到。 + processManualTestChannelError(channel, result) if result.localErr != nil { resp := gin.H{ "success": false, @@ -899,6 +904,24 @@ func TestChannel(c *gin.Context) { }) } +// processManualTestChannelError 让手动「测试连接」失败与正常转发、定时巡检走同一套 +// 错误处理:先按多 Key 限额模式给命中的密钥设置冷却;未命中且满足自动禁用条件 +// (全局开关 + 渠道 AutoBan + 状态码/关键词规则)时禁用该密钥。渠道未启用时不处理。 +// 注意上游返回错误时 testResult 的 localErr 与 newAPIError 同时非空,不能用 localErr +// 区分"测试未发出",因此仅以 newAPIError 是否存在为准。 +func processManualTestChannelError(channel *model.Channel, result testResult) { + if result.newAPIError == nil { + return + } + if channel.Status != common.ChannelStatusEnabled { + return + } + processChannelError(result.context, *types.NewChannelError( + channel.Id, channel.Type, channel.Name, channel.ChannelInfo.IsMultiKey, + common.GetContextKeyString(result.context, constant.ContextKeyChannelKey), + channel.GetAutoBan()), result.newAPIError) +} + // channelTestSummary records the outcome of one channel test cycle so the // system task can persist a per-run result for history. type channelTestSummary struct { diff --git a/controller/channel.go b/controller/channel.go index 4b2950615598..69c1561f2f3d 100644 --- a/controller/channel.go +++ b/controller/channel.go @@ -1699,6 +1699,7 @@ func ManageMultiKeys(c *gin.Context) { channel.ChannelInfo.MultiKeyStatusList[keyIndex] = 2 // disabled + channel.SyncMultiKeyChannelStatus() err = channel.Update() if err != nil { common.ApiError(c, err) @@ -1744,6 +1745,7 @@ func ManageMultiKeys(c *gin.Context) { delete(channel.ChannelInfo.MultiKeyCooldownUntil, keyIndex) } + channel.SyncMultiKeyChannelStatus() err = channel.Update() if err != nil { common.ApiError(c, err) @@ -1769,6 +1771,7 @@ func ManageMultiKeys(c *gin.Context) { channel.ChannelInfo.MultiKeyDisabledReason = make(map[int]string) channel.ChannelInfo.MultiKeyCooldownUntil = make(map[int]int64) + channel.SyncMultiKeyChannelStatus() err = channel.Update() if err != nil { common.ApiError(c, err) @@ -1816,6 +1819,7 @@ func ManageMultiKeys(c *gin.Context) { return } + channel.SyncMultiKeyChannelStatus() err = channel.Update() if err != nil { common.ApiError(c, err) @@ -1896,6 +1900,7 @@ func ManageMultiKeys(c *gin.Context) { channel.ChannelInfo.MultiKeyDisabledTime = newDisabledTime channel.ChannelInfo.MultiKeyDisabledReason = newDisabledReason + channel.SyncMultiKeyChannelStatus() err = channel.Update() if err != nil { common.ApiError(c, err) @@ -1964,6 +1969,7 @@ func ManageMultiKeys(c *gin.Context) { channel.ChannelInfo.MultiKeyDisabledTime = newDisabledTime channel.ChannelInfo.MultiKeyDisabledReason = newDisabledReason + channel.SyncMultiKeyChannelStatus() err = channel.Update() if err != nil { common.ApiError(c, err) diff --git a/controller/channel_test_internal_test.go b/controller/channel_test_internal_test.go index 85da7f7bab5e..d55dc5fbc51e 100644 --- a/controller/channel_test_internal_test.go +++ b/controller/channel_test_internal_test.go @@ -3,24 +3,30 @@ package controller import ( "bytes" "context" + "errors" "fmt" "net/http" "net/http/httptest" + "strconv" "sync/atomic" "testing" + "time" "github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/constant" "github.com/QuantumNous/new-api/model" "github.com/QuantumNous/new-api/pkg/billingexpr" relaycommon "github.com/QuantumNous/new-api/relay/common" + relaytypes "github.com/QuantumNous/new-api/relaykit/types" "github.com/QuantumNous/new-api/relaykit/dto" "github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/setting/operation_setting" "github.com/QuantumNous/new-api/types" "github.com/gin-gonic/gin" + "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "gorm.io/gorm" ) func TestValidateChannelProxy(t *testing.T) { @@ -463,3 +469,303 @@ func TestTestAllChannelsRejectsExistingActiveTask(t *testing.T) { require.Contains(t, recorder.Body.String(), existing.TaskID) require.Contains(t, recorder.Body.String(), "已有通道测试任务正在运行或等待中") } + +func setupManualChannelTestErrorTest(t *testing.T) { + t.Helper() + previousDB := model.DB + previousType := common.MainDatabaseType() + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + require.NoError(t, err) + // :memory: 下每个连接是独立库,单连接保证异步禁用协程与断言读到同一份数据 + sqlDB, err := db.DB() + require.NoError(t, err) + sqlDB.SetMaxOpenConns(1) + require.NoError(t, db.AutoMigrate(&model.Channel{}, &model.User{}, &model.Ability{})) + model.DB = db + common.SetMainDatabaseType(common.DatabaseTypeSQLite) + // 禁用渠道后的通知协程会查询 root 用户;Setting 允许未配置价格的测试模型 + require.NoError(t, db.Create(&model.User{ + Username: "root", Role: common.RoleRootUser, Status: common.UserStatusEnabled, + Setting: `{"accept_unset_model_ratio_model":true}`, + }).Error) + + memoryCacheEnabled := common.MemoryCacheEnabled + common.MemoryCacheEnabled = false + // 通知路径按 RedisEnabled 分流;测试进程未初始化 Redis,须显式走内存限流 + redisEnabled := common.RedisEnabled + common.RedisEnabled = false + autoDisableEnabled := common.AutomaticDisableChannelEnabled + common.AutomaticDisableChannelEnabled = true + keywords := operation_setting.AutomaticDisableKeywords + operation_setting.AutomaticDisableKeywordsFromString("insufficient balance") + errorLogEnabled := constant.ErrorLogEnabled + constant.ErrorLogEnabled = false + + t.Cleanup(func() { + model.DB = previousDB + common.SetMainDatabaseType(previousType) + common.MemoryCacheEnabled = memoryCacheEnabled + common.RedisEnabled = redisEnabled + common.AutomaticDisableChannelEnabled = autoDisableEnabled + operation_setting.AutomaticDisableKeywords = keywords + constant.ErrorLogEnabled = errorLogEnabled + }) +} + +func newManualTestContext(usingKey string) *gin.Context { + gin.SetMode(gin.TestMode) + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil) + common.SetContextKey(c, constant.ContextKeyChannelKey, usingKey) + return c +} + +// TestTestChannelHandlerCoolsKeyOnUpstreamQuotaError drives the full manual +// test handler against a stub upstream returning 429 insufficient quota, and +// protects the wiring: the used key must get the limit-pattern cooldown even +// though the upstream error also sets localErr (which makes TestChannel return +// early to the client). +func TestTestChannelHandlerCoolsKeyOnUpstreamQuotaError(t *testing.T) { + setupManualChannelTestErrorTest(t) + // 限额模式冷却不应依赖全局自动禁用开关 + common.AutomaticDisableChannelEnabled = false + + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusTooManyRequests) + _, _ = w.Write([]byte(`{"error":{"message":"Your token-plan quota has been exhausted.","type":"insufficient_quota","code":"insufficient_quota"}}`)) + })) + t.Cleanup(upstream.Close) + upstreamURL := upstream.URL + + channel := &model.Channel{ + Name: "aliyun-quota-e2e", + Type: constant.ChannelTypeOpenAI, + Key: "upk1", + Status: common.ChannelStatusEnabled, + Models: "qwen-e2e", + Group: "default", + BaseURL: &upstreamURL, + Priority: func() *int64 { p := int64(0); return &p }(), + ChannelInfo: model.ChannelInfo{ + IsMultiKey: true, + MultiKeySize: 1, + MultiKeyMode: constant.MultiKeyModeRandom, + MultiKeyLimitPatterns: []model.LimitPattern{ + {Name: "aliyun", Regex: "Your token-plan quota has been exhausted", ResetCycle: "monthly:1"}, + }, + }, + } + require.NoError(t, model.DB.Create(channel).Error) + + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/api/channel/test/"+strconv.Itoa(channel.Id), nil) + c.Params = gin.Params{{Key: "id", Value: strconv.Itoa(channel.Id)}} + + TestChannel(c) + + require.Equal(t, http.StatusOK, recorder.Code) + assert.Contains(t, recorder.Body.String(), "token-plan quota has been exhausted") + + // 冷却同步落库,无需等待异步任务 + var stored model.Channel + require.NoError(t, model.DB.First(&stored, channel.Id).Error) + assert.Equal(t, common.ChannelStatusTempDisabled, stored.ChannelInfo.MultiKeyStatusList[0], + "used key should be cooled down by the limit pattern") + assert.Greater(t, stored.ChannelInfo.MultiKeyCooldownUntil[0], common.GetTimestamp()) + assert.Equal(t, common.ChannelStatusEnabled, stored.Status, "channel stays enabled while cooling") +} + +func newInsufficientBalanceError() *relaytypes.NewAPIError { + return relaytypes.NewOpenAIError( + errors.New("insufficient balance"), + relaytypes.ErrorCodeBadResponseStatusCode, + http.StatusPaymentRequired, + ) +} + +// TestProcessManualTestChannelErrorDisablesFailingKey protects the reported +// behavior: a manual "test connection" failure (e.g. insufficient balance) +// must auto-disable the used key just like the relay path, when the channel +// has auto-ban enabled and the error matches the disable rules. Upstream +// errors produce a testResult with BOTH localErr and newAPIError set, so the +// case mirrors that real shape. +func TestProcessManualTestChannelErrorDisablesFailingKey(t *testing.T) { + setupManualChannelTestErrorTest(t) + + autoBan := 1 + channel := &model.Channel{ + Name: "manual-test-multi-key", + Type: constant.ChannelTypeOpenAI, + Key: "k1\nk2", + Status: common.ChannelStatusEnabled, + Models: "gpt-4o-mini", + Group: "default", + AutoBan: &autoBan, + ChannelInfo: model.ChannelInfo{ + IsMultiKey: true, + MultiKeySize: 2, + MultiKeyMode: constant.MultiKeyModeRandom, + }, + } + require.NoError(t, model.DB.Create(channel).Error) + + upstreamErr := newInsufficientBalanceError() + processManualTestChannelError(channel, testResult{ + context: newManualTestContext("k1"), + localErr: upstreamErr, + newAPIError: upstreamErr, + }) + + require.Eventually(t, func() bool { + var stored model.Channel + if err := model.DB.First(&stored, channel.Id).Error; err != nil { + return false + } + return stored.ChannelInfo.MultiKeyStatusList[0] == common.ChannelStatusAutoDisabled + }, 5*time.Second, 20*time.Millisecond, "used key should be auto-disabled") + + var stored model.Channel + require.NoError(t, model.DB.First(&stored, channel.Id).Error) + assert.Equal(t, common.ChannelStatusAutoDisabled, stored.ChannelInfo.MultiKeyStatusList[0]) + assert.Contains(t, stored.ChannelInfo.MultiKeyDisabledReason[0], "insufficient balance") + // 另一个 key 仍然可用,渠道不应被整体禁用 + assert.NotContains(t, stored.ChannelInfo.MultiKeyStatusList, 1) + assert.Equal(t, common.ChannelStatusEnabled, stored.Status) +} + +// TestProcessManualTestChannelErrorDisablesByStatusCode covers the disable +// rule driven by the upstream status code (402) rather than keywords, which +// only works when the test keeps the real upstream status code. +func TestProcessManualTestChannelErrorDisablesByStatusCode(t *testing.T) { + setupManualChannelTestErrorTest(t) + operation_setting.AutomaticDisableKeywordsFromString("") + previousRanges := operation_setting.AutomaticDisableStatusCodeRanges + operation_setting.AutomaticDisableStatusCodeRanges = []operation_setting.StatusCodeRange{{Start: 402, End: 402}} + t.Cleanup(func() { + operation_setting.AutomaticDisableStatusCodeRanges = previousRanges + }) + + autoBan := 1 + channel := &model.Channel{ + Name: "manual-test-status-code", + Type: constant.ChannelTypeOpenAI, + Key: "k1\nk2", + Status: common.ChannelStatusEnabled, + Models: "gpt-4o-mini", + Group: "default", + AutoBan: &autoBan, + ChannelInfo: model.ChannelInfo{ + IsMultiKey: true, + MultiKeySize: 2, + MultiKeyMode: constant.MultiKeyModeRandom, + }, + } + require.NoError(t, model.DB.Create(channel).Error) + + upstreamErr := newInsufficientBalanceError() + processManualTestChannelError(channel, testResult{ + context: newManualTestContext("k1"), + localErr: upstreamErr, + newAPIError: upstreamErr, + }) + + require.Eventually(t, func() bool { + var stored model.Channel + if err := model.DB.First(&stored, channel.Id).Error; err != nil { + return false + } + return stored.ChannelInfo.MultiKeyStatusList[0] == common.ChannelStatusAutoDisabled + }, 5*time.Second, 20*time.Millisecond, "used key should be auto-disabled by status code rule") +} + +// TestProcessManualTestChannelErrorKeyLimitCooldownWins mirrors the relay +// semantics: when a multi-key limit pattern matches, the key gets a cooldown +// (temp disabled) instead of a hard disable, even with auto-ban enabled. +func TestProcessManualTestChannelErrorKeyLimitCooldownWins(t *testing.T) { + setupManualChannelTestErrorTest(t) + + autoBan := 1 + channel := &model.Channel{ + Name: "manual-test-limit-pattern", + Type: constant.ChannelTypeOpenAI, + Key: "k1\nk2", + Status: common.ChannelStatusEnabled, + Models: "gpt-4o-mini", + Group: "default", + AutoBan: &autoBan, + ChannelInfo: model.ChannelInfo{ + IsMultiKey: true, + MultiKeySize: 2, + MultiKeyMode: constant.MultiKeyModeRandom, + MultiKeyLimitPatterns: []model.LimitPattern{ + {Name: "balance", Regex: "insufficient balance", ResetCycle: "monthly:1"}, + }, + }, + } + require.NoError(t, model.DB.Create(channel).Error) + + processManualTestChannelError(channel, testResult{ + context: newManualTestContext("k1"), + newAPIError: newInsufficientBalanceError(), + }) + + var stored model.Channel + require.NoError(t, model.DB.First(&stored, channel.Id).Error) + assert.Equal(t, common.ChannelStatusTempDisabled, stored.ChannelInfo.MultiKeyStatusList[0]) + assert.Greater(t, stored.ChannelInfo.MultiKeyCooldownUntil[0], common.GetTimestamp()) + assert.Contains(t, stored.ChannelInfo.MultiKeyDisabledReason[0], "balance") +} + +func TestProcessManualTestChannelErrorSkipsInapplicableCases(t *testing.T) { + setupManualChannelTestErrorTest(t) + + autoBan := 1 + disabledChannel := &model.Channel{ + Name: "already-disabled", Type: constant.ChannelTypeOpenAI, Key: "k1", + Status: common.ChannelStatusManuallyDisabled, Models: "gpt-4o-mini", Group: "default", + AutoBan: &autoBan, + ChannelInfo: model.ChannelInfo{IsMultiKey: true, MultiKeySize: 1, MultiKeyMode: constant.MultiKeyModeRandom}, + } + require.NoError(t, model.DB.Create(disabledChannel).Error) + processManualTestChannelError(disabledChannel, testResult{ + context: newManualTestContext("k1"), + newAPIError: newInsufficientBalanceError(), + }) + + // AutoBan 带 gorm:"default:1",必须显式置 0 才能构造"未开启自动禁用"的渠道 + autoBanOff := 0 + enabledNoAutoBan := &model.Channel{ + Name: "no-autoban", Type: constant.ChannelTypeOpenAI, Key: "k1", + Status: common.ChannelStatusEnabled, Models: "gpt-4o-mini", Group: "default", + AutoBan: &autoBanOff, + ChannelInfo: model.ChannelInfo{IsMultiKey: true, MultiKeySize: 1, MultiKeyMode: constant.MultiKeyModeRandom}, + } + require.NoError(t, model.DB.Create(enabledNoAutoBan).Error) + processManualTestChannelError(enabledNoAutoBan, testResult{ + context: newManualTestContext("k1"), + newAPIError: newInsufficientBalanceError(), + }) + + enabled := &model.Channel{ + Name: "local-error", Type: constant.ChannelTypeOpenAI, Key: "k1", + Status: common.ChannelStatusEnabled, Models: "gpt-4o-mini", Group: "default", + AutoBan: &autoBan, + ChannelInfo: model.ChannelInfo{IsMultiKey: true, MultiKeySize: 1, MultiKeyMode: constant.MultiKeyModeRandom}, + } + require.NoError(t, model.DB.Create(enabled).Error) + processManualTestChannelError(enabled, testResult{ + context: newManualTestContext("k1"), + localErr: errors.New("request build failed"), + }) + processManualTestChannelError(enabled, testResult{context: newManualTestContext("k1")}) + + for _, ch := range []*model.Channel{disabledChannel, enabledNoAutoBan, enabled} { + var stored model.Channel + require.NoError(t, model.DB.First(&stored, ch.Id).Error) + assert.Empty(t, stored.ChannelInfo.MultiKeyStatusList, "channel %s should stay untouched", ch.Name) + assert.Equal(t, ch.Status, stored.Status) + } +} diff --git a/i18n/keys.go b/i18n/keys.go index 64a835e1a942..5e1e61812459 100644 --- a/i18n/keys.go +++ b/i18n/keys.go @@ -322,6 +322,7 @@ const ( MsgDistributorGroupAccessDenied = "distributor.group_access_denied" MsgDistributorGetChannelFailed = "distributor.get_channel_failed" MsgDistributorNoAvailableChannel = "distributor.no_available_channel" + MsgDistributorChannelNoAvailableKey = "distributor.channel_no_available_key" MsgDistributorInvalidMidjourney = "distributor.invalid_midjourney_request" MsgDistributorInvalidParseModel = "distributor.invalid_request_parse_model" ) diff --git a/i18n/locales/en.yaml b/i18n/locales/en.yaml index c533daecc32d..adba073664e6 100644 --- a/i18n/locales/en.yaml +++ b/i18n/locales/en.yaml @@ -272,6 +272,7 @@ distributor.invalid_playground_request: "Invalid playground request: {{.Error}}" distributor.group_access_denied: "No permission to access this group" distributor.get_channel_failed: "Failed to get available channel for model {{.Model}} under group {{.Group}} (distributor): {{.Error}}" distributor.no_available_channel: "No available channel for model {{.Model}} under group {{.Group}} (distributor)" +distributor.channel_no_available_key: "No channel with an available key for model {{.Model}} under group {{.Group}}: all keys are disabled or cooling down (distributor)" distributor.invalid_midjourney_request: "Invalid Midjourney request: {{.Error}}" distributor.invalid_request_parse_model: "Invalid request, unable to parse model" diff --git a/i18n/locales/zh-CN.yaml b/i18n/locales/zh-CN.yaml index a2f5275be9a8..db9b79f3baa6 100644 --- a/i18n/locales/zh-CN.yaml +++ b/i18n/locales/zh-CN.yaml @@ -273,6 +273,7 @@ distributor.invalid_playground_request: "无效的playground请求,{{.Error}}" distributor.group_access_denied: "无权访问该分组" distributor.get_channel_failed: "获取分组 {{.Group}} 下模型 {{.Model}} 的可用渠道失败(distributor):{{.Error}}" distributor.no_available_channel: "分组 {{.Group}} 下模型 {{.Model}} 无可用渠道(distributor)" +distributor.channel_no_available_key: "分组 {{.Group}} 下模型 {{.Model}} 的渠道暂无可用密钥(密钥全部被禁用或冷却中)(distributor)" distributor.invalid_midjourney_request: "无效的midjourney请求,{{.Error}}" distributor.invalid_request_parse_model: "无效的请求,无法解析模型" diff --git a/i18n/locales/zh-TW.yaml b/i18n/locales/zh-TW.yaml index 84ebd57ed587..78c5bfe1c9fe 100644 --- a/i18n/locales/zh-TW.yaml +++ b/i18n/locales/zh-TW.yaml @@ -273,6 +273,7 @@ distributor.invalid_playground_request: "無效的playground請求,{{.Error}}" distributor.group_access_denied: "無權存取該分組" distributor.get_channel_failed: "獲取分組 {{.Group}} 下模型 {{.Model}} 的可用管道失敗(distributor):{{.Error}}" distributor.no_available_channel: "分組 {{.Group}} 下模型 {{.Model}} 無可用管道(distributor)" +distributor.channel_no_available_key: "分組 {{.Group}} 下模型 {{.Model}} 的管道暫無可用密鑰(密鑰全部被停用或冷卻中)(distributor)" distributor.invalid_midjourney_request: "無效的midjourney請求,{{.Error}}" distributor.invalid_request_parse_model: "無效的請求,無法解析模型" diff --git a/middleware/distributor.go b/middleware/distributor.go index 3f53aa350349..a17f9c50f569 100644 --- a/middleware/distributor.go +++ b/middleware/distributor.go @@ -162,7 +162,30 @@ func Distribute() func(c *gin.Context) { } } common.SetContextKey(c, constant.ContextKeyRequestStartTime, time.Now()) - SetupContextForSelectedChannel(c, channel, modelRequest.Model) + setupErr := SetupContextForSelectedChannel(c, channel, modelRequest.Model) + if setupErr != nil && channel != nil && setupErr.GetErrorCode() == types.ErrorCodeChannelNoAvailableKey && !ok && shouldSelectChannel { + // 渠道本身启用但暂时没有可用密钥(例如多 key 渠道的密钥全部被禁用或处于冷却期)时, + // 逐级降低优先级重选渠道,而不是携带空密钥把请求发给上游。 + // 令牌绑定了指定渠道(ok)时不降级,直接返回错误。 + if next := selectChannelWithAvailableKey(c, channel.Id, modelRequest.Model); next != nil { + channel = next + setupErr = nil + } + } + if setupErr != nil && channel != nil { + if setupErr.GetErrorCode() == types.ErrorCodeChannelNoAvailableKey { + showGroup := common.GetContextKeyString(c, constant.ContextKeyUsingGroup) + if selectGroup, exists := common.GetContextKey(c, constant.ContextKeyAutoGroup); exists { + if g, isStr := selectGroup.(string); isStr { + showGroup = fmt.Sprintf("auto(%s)", g) + } + } + abortWithOpenAiMessage(c, http.StatusServiceUnavailable, i18n.T(c, i18n.MsgDistributorChannelNoAvailableKey, map[string]any{"Group": showGroup, "Model": modelRequest.Model}), types.ErrorCodeChannelNoAvailableKey) + } else { + abortWithOpenAiMessage(c, setupErr.StatusCode, setupErr.Error(), setupErr.GetErrorCode()) + } + return + } c.Next() if channel != nil && c.Writer != nil && c.Writer.Status() < http.StatusBadRequest { service.RecordChannelAffinity(c, channel.Id) @@ -170,6 +193,40 @@ func Distribute() func(c *gin.Context) { } } +// selectChannelWithAvailableKey 在原选中渠道因 channel:no_available_key 失败后, +// 按与 relay 重试一致的语义(retry 序号 = 优先级层级)逐级降低优先级重选渠道, +// 直到某个渠道真正取到可用密钥。同层内的选择是随机的,重选到已尝试的渠道时 +// 继续重试(同层可能还有未尝试的可用渠道),仅当连续尝试明显超过已见渠道数 +// 仍未出现新渠道时视为候选耗尽,返回 nil 由调用方返回错误。 +func selectChannelWithAvailableKey(c *gin.Context, failedChannelId int, modelName string) *model.Channel { + usingGroup := common.GetContextKeyString(c, constant.ContextKeyUsingGroup) + tried := map[int]bool{failedChannelId: true} + attempts := 0 + for retry := 1; ; retry++ { + attempts++ + next, _, err := service.CacheGetRandomSatisfiedChannel(&service.RetryParam{ + Ctx: c, + ModelName: modelName, + TokenGroup: usingGroup, + RequestPath: c.Request.URL.Path, + Retry: common.GetPointer(retry), + }) + if err != nil || next == nil { + return nil + } + if tried[next.Id] { + if attempts > len(tried)*2+2 { + return nil + } + continue + } + tried[next.Id] = true + if SetupContextForSelectedChannel(c, next, modelName) == nil { + return next + } + } +} + // channelSupportsRequestPath reports whether a channel can serve the request path. // Only Advanced Custom (type 58) channels are path-checked; all other channel types // always pass. A type-58 channel is usable only when one of its routes matches. diff --git a/middleware/distributor_fallback_test.go b/middleware/distributor_fallback_test.go new file mode 100644 index 000000000000..84527f016430 --- /dev/null +++ b/middleware/distributor_fallback_test.go @@ -0,0 +1,193 @@ +package middleware + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/constant" + "github.com/QuantumNous/new-api/model" + "github.com/QuantumNous/new-api/relaykit/types" + "github.com/QuantumNous/new-api/service" + "github.com/gin-gonic/gin" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +// setupDistributorFallbackDB lazily initializes a SQLite database backing the +// channel selection cache for the fallback tests. Middleware package tests do +// not bootstrap model.DB by default, and package tests run sequentially. +func setupDistributorFallbackDB(t *testing.T) { + t.Helper() + if model.DB == nil { + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + panic("failed to open test db: " + err.Error()) + } + // :memory: 下每个连接是独立库,单连接保证所有读写落在同一份数据上 + sqlDB, err := db.DB() + if err != nil { + panic("failed to get sql.DB: " + err.Error()) + } + sqlDB.SetMaxOpenConns(1) + model.DB = db + common.SetDatabaseTypes(common.DatabaseTypeSQLite, common.DatabaseTypeSQLite) + if err := db.AutoMigrate(&model.Channel{}, &model.Ability{}); err != nil { + panic("failed to migrate test db: " + err.Error()) + } + } + require.NoError(t, model.DB.Exec("DELETE FROM abilities").Error) + require.NoError(t, model.DB.Exec("DELETE FROM channels").Error) + + memoryCacheEnabled := common.MemoryCacheEnabled + common.MemoryCacheEnabled = true + t.Cleanup(func() { + common.MemoryCacheEnabled = memoryCacheEnabled + model.InitChannelCache() + }) +} + +// TestSelectChannelWithAvailableKeyFallsBackToLowerPriority protects the +// scenario where the highest-priority channel stays enabled but has no usable +// key (all keys disabled or cooling down): the request must fall through to the +// lower-priority channel instead of being sent upstream with an empty key. +func TestSelectChannelWithAvailableKeyFallsBackToLowerPriority(t *testing.T) { + setupDistributorFallbackDB(t) + + highPriority := int64(10) + high := &model.Channel{ + Name: "high-multi-key", + Type: constant.ChannelTypeOpenAI, + Key: "hk1\nhk2", + Status: common.ChannelStatusEnabled, + Models: "fallback-model", + Group: "default", + Priority: &highPriority, + ChannelInfo: model.ChannelInfo{ + IsMultiKey: true, + MultiKeySize: 2, + MultiKeyMode: constant.MultiKeyModeRandom, + MultiKeyStatusList: map[int]int{0: common.ChannelStatusTempDisabled, 1: common.ChannelStatusTempDisabled}, + MultiKeyCooldownUntil: map[int]int64{0: common.GetTimestamp() + 3600, 1: common.GetTimestamp() + 3600}, + }, + } + require.NoError(t, model.DB.Create(high).Error) + require.NoError(t, high.AddAbilities(nil)) + + lowPriority := int64(0) + low := &model.Channel{ + Name: "low-single-key", + Type: constant.ChannelTypeOpenAI, + Key: "low-key", + Status: common.ChannelStatusEnabled, + Models: "fallback-model", + Group: "default", + Priority: &lowPriority, + } + require.NoError(t, model.DB.Create(low).Error) + require.NoError(t, low.AddAbilities(nil)) + + model.InitChannelCache() + + gin.SetMode(gin.TestMode) + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil) + common.SetContextKey(c, constant.ContextKeyUsingGroup, "default") + + selected, _, err := service.CacheGetRandomSatisfiedChannel(&service.RetryParam{ + Ctx: c, + ModelName: "fallback-model", + TokenGroup: "default", + RequestPath: "/v1/chat/completions", + Retry: common.GetPointer(0), + }) + require.NoError(t, err) + require.NotNil(t, selected) + require.Equal(t, high.Id, selected.Id) + + setupErr := SetupContextForSelectedChannel(c, selected, "fallback-model") + require.NotNil(t, setupErr) + assert.Equal(t, types.ErrorCodeChannelNoAvailableKey, setupErr.GetErrorCode()) + + next := selectChannelWithAvailableKey(c, selected.Id, "fallback-model") + require.NotNil(t, next) + assert.Equal(t, low.Id, next.Id) + assert.Equal(t, "low-key", common.GetContextKeyString(c, constant.ContextKeyChannelKey)) +} + +// TestSelectChannelWithAvailableKeyStopsWhenAllChannelsRepeat ensures the +// fallback terminates (returns nil) when every candidate channel has already +// been tried, instead of looping forever on a saturated lowest priority tier. +func TestSelectChannelWithAvailableKeyStopsWhenAllChannelsRepeat(t *testing.T) { + setupDistributorFallbackDB(t) + + only := &model.Channel{ + Name: "only-cooling", + Type: constant.ChannelTypeOpenAI, + Key: "k1", + Status: common.ChannelStatusEnabled, + Models: "fallback-model", + Group: "default", + ChannelInfo: model.ChannelInfo{ + IsMultiKey: true, + MultiKeySize: 1, + MultiKeyMode: constant.MultiKeyModeRandom, + MultiKeyStatusList: map[int]int{0: common.ChannelStatusTempDisabled}, + MultiKeyCooldownUntil: map[int]int64{0: common.GetTimestamp() + 3600}, + }, + } + require.NoError(t, model.DB.Create(only).Error) + require.NoError(t, only.AddAbilities(nil)) + + model.InitChannelCache() + + gin.SetMode(gin.TestMode) + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil) + common.SetContextKey(c, constant.ContextKeyUsingGroup, "default") + + assert.Nil(t, selectChannelWithAvailableKey(c, only.Id, "fallback-model")) +} + +// TestSelectChannelWithAvailableKeyTerminatesAcrossTiers drives the fallback +// through multiple priority tiers whose channels all lack a usable key; it +// must exhaust the random same-tier retries and return nil deterministically. +func TestSelectChannelWithAvailableKeyTerminatesAcrossTiers(t *testing.T) { + setupDistributorFallbackDB(t) + + highPriority, lowPriority := int64(10), int64(0) + newCoolingChannel := func(name string, priority *int64) *model.Channel { + return &model.Channel{ + Name: name, Type: constant.ChannelTypeOpenAI, + Key: "k1\nk2", Status: common.ChannelStatusEnabled, + Models: "fallback-model", Group: "default", Priority: priority, + ChannelInfo: model.ChannelInfo{ + IsMultiKey: true, + MultiKeySize: 2, + MultiKeyMode: constant.MultiKeyModeRandom, + MultiKeyStatusList: map[int]int{0: common.ChannelStatusTempDisabled, 1: common.ChannelStatusTempDisabled}, + MultiKeyCooldownUntil: map[int]int64{0: common.GetTimestamp() + 3600, 1: common.GetTimestamp() + 3600}, + }, + } + } + for _, def := range []struct { + name string + priority *int64 + }{{"cool-high", &highPriority}, {"cool-low-a", &lowPriority}, {"cool-low-b", &lowPriority}} { + ch := newCoolingChannel(def.name, def.priority) + require.NoError(t, model.DB.Create(ch).Error) + require.NoError(t, ch.AddAbilities(nil)) + } + + model.InitChannelCache() + + gin.SetMode(gin.TestMode) + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil) + common.SetContextKey(c, constant.ContextKeyUsingGroup, "default") + + assert.Nil(t, selectChannelWithAvailableKey(c, 1, "fallback-model")) +} diff --git a/model/channel.go b/model/channel.go index 0cf4e43b4bc9..2dd03fc4b895 100644 --- a/model/channel.go +++ b/model/channel.go @@ -64,6 +64,10 @@ type LimitPattern struct { Regex string `json:"regex"` DateLayout string `json:"date_layout"` DefaultMinutes int `json:"default_minutes"` + // ResetCycle 为错误信息不含重置时间时的周期重置规则: + // "daily"(每日 0 点)、"weekly:N"(N=1 周一 .. 7 周日)、"monthly:N"(每月 N 号, + // 当月无 N 号时顺延到下月 1 号)。为空时回退到 DefaultMinutes。 + ResetCycle string `json:"reset_cycle,omitempty"` } type ChannelInfo struct { @@ -744,7 +748,7 @@ func handlerMultiKeyUpdate(channel *Channel, usingKey string, status int, reason if !hasEnabledMultiKey(keys, channel.ChannelInfo.MultiKeyStatusList) && !hasCoolingMultiKey(keys, channel.ChannelInfo.MultiKeyStatusList) { channel.Status = common.ChannelStatusAutoDisabled info := channel.GetOtherInfo() - info["status_reason"] = "All keys are disabled" + info["status_reason"] = allKeysDisabledReason info["status_time"] = common.GetTimestamp() channel.SetOtherInfo(info) } else if status == common.ChannelStatusEnabled { @@ -779,6 +783,40 @@ func hasCoolingMultiKey(keys []string, statusList map[int]int) bool { return false } +// allKeysDisabledReason marks channels whose status was flipped to disabled +// because every key is disabled (none cooling down). It is the re-enable +// criterion in SyncMultiKeyChannelStatus, so the marker must stay identical for +// the automatic path (handlerMultiKeyUpdate) and the manual key-management path. +const allKeysDisabledReason = "All keys are disabled" + +// SyncMultiKeyChannelStatus aligns the channel status with per-key availability +// after key management operations (disable/enable/delete in ManageMultiKeys). +// A still-enabled channel whose keys are all disabled — and none cooling down — +// is marked auto-disabled, so the distributor stops selecting it; a channel that +// was disabled by that rule is re-enabled once any key becomes usable again. +// Channels disabled manually or auto-disabled for other reasons are left +// untouched. The caller persists the change (Channel.Update + abilities sync). +func (channel *Channel) SyncMultiKeyChannelStatus() { + keys := channel.GetKeys() + hasUsableKey := hasEnabledMultiKey(keys, channel.ChannelInfo.MultiKeyStatusList) || + hasCoolingMultiKey(keys, channel.ChannelInfo.MultiKeyStatusList) + switch { + case !hasUsableKey && channel.Status == common.ChannelStatusEnabled: + channel.Status = common.ChannelStatusAutoDisabled + info := channel.GetOtherInfo() + info["status_reason"] = allKeysDisabledReason + info["status_time"] = common.GetTimestamp() + channel.SetOtherInfo(info) + case hasUsableKey && channel.Status == common.ChannelStatusAutoDisabled && + channel.GetOtherInfo()["status_reason"] == allKeysDisabledReason: + channel.Status = common.ChannelStatusEnabled + info := channel.GetOtherInfo() + delete(info, "status_reason") + delete(info, "status_time") + channel.SetOtherInfo(info) + } +} + func UpdateChannelStatus(channelId int, usingKey string, status int, reason string) bool { return updateChannelStatusWithCooldown(channelId, usingKey, status, reason, nil) } diff --git a/model/channel_status_test.go b/model/channel_status_test.go index e4ad86f8133c..6d6d5afc964a 100644 --- a/model/channel_status_test.go +++ b/model/channel_status_test.go @@ -51,6 +51,147 @@ func TestUpdateChannelStatusPersistsMultiKeyState(t *testing.T) { assert.Equal(t, 1, stored.ChannelInfo.MultiKeyPollingIndex) } +func TestSyncMultiKeyChannelStatusDisablesWhenAllKeysDisabled(t *testing.T) { + channel := &Channel{ + Id: 1, + Key: "k1\nk2", + Status: common.ChannelStatusEnabled, + ChannelInfo: ChannelInfo{ + IsMultiKey: true, + MultiKeySize: 2, + MultiKeyStatusList: map[int]int{0: common.ChannelStatusManuallyDisabled, 1: common.ChannelStatusManuallyDisabled}, + }, + } + + channel.SyncMultiKeyChannelStatus() + + assert.Equal(t, common.ChannelStatusAutoDisabled, channel.Status) + assert.Equal(t, allKeysDisabledReason, channel.GetOtherInfo()["status_reason"]) +} + +func TestSyncMultiKeyChannelStatusDisablesWhenKeyListEmpty(t *testing.T) { + channel := &Channel{ + Id: 2, + Key: "", + Status: common.ChannelStatusEnabled, + ChannelInfo: ChannelInfo{ + IsMultiKey: true, + MultiKeySize: 0, + }, + } + + channel.SyncMultiKeyChannelStatus() + + assert.Equal(t, common.ChannelStatusAutoDisabled, channel.Status) +} + +func TestSyncMultiKeyChannelStatusKeepsEnabledWhileKeysCooling(t *testing.T) { + channel := &Channel{ + Id: 3, + Key: "k1\nk2", + Status: common.ChannelStatusEnabled, + ChannelInfo: ChannelInfo{ + IsMultiKey: true, + MultiKeySize: 2, + MultiKeyStatusList: map[int]int{0: common.ChannelStatusManuallyDisabled, 1: common.ChannelStatusTempDisabled}, + MultiKeyCooldownUntil: map[int]int64{1: common.GetTimestamp() + 3600}, + }, + } + + channel.SyncMultiKeyChannelStatus() + + assert.Equal(t, common.ChannelStatusEnabled, channel.Status) +} + +func TestSyncMultiKeyChannelStatusReenablesWhenKeyBecomesUsable(t *testing.T) { + channel := &Channel{ + Id: 4, + Key: "k1\nk2", + Status: common.ChannelStatusAutoDisabled, + OtherInfo: `{"status_reason":"` + allKeysDisabledReason + `","status_time":123}`, + ChannelInfo: ChannelInfo{ + IsMultiKey: true, + MultiKeySize: 2, + MultiKeyStatusList: map[int]int{0: common.ChannelStatusManuallyDisabled}, + }, + } + + channel.SyncMultiKeyChannelStatus() + + assert.Equal(t, common.ChannelStatusEnabled, channel.Status) + assert.NotContains(t, channel.GetOtherInfo(), "status_reason") +} + +func TestSyncMultiKeyChannelStatusLeavesUnrelatedDisabledStates(t *testing.T) { + manuallyDisabled := &Channel{ + Id: 5, + Key: "k1", + Status: common.ChannelStatusManuallyDisabled, + ChannelInfo: ChannelInfo{ + IsMultiKey: true, + MultiKeySize: 1, + MultiKeyStatusList: map[int]int{0: common.ChannelStatusManuallyDisabled}, + }, + } + manuallyDisabled.SyncMultiKeyChannelStatus() + assert.Equal(t, common.ChannelStatusManuallyDisabled, manuallyDisabled.Status) + + autoDisabledOtherReason := &Channel{ + Id: 6, + Key: "k1\nk2", + Status: common.ChannelStatusAutoDisabled, + OtherInfo: `{"status_reason":"provider rejected key"}`, + ChannelInfo: ChannelInfo{ + IsMultiKey: true, + MultiKeySize: 2, + }, + } + autoDisabledOtherReason.SyncMultiKeyChannelStatus() + assert.Equal(t, common.ChannelStatusAutoDisabled, autoDisabledOtherReason.Status) +} + +func TestSyncMultiKeyChannelStatusFlowsToDbAndAbilities(t *testing.T) { + setupChannelStatusTest(t) + + channel := &Channel{ + Name: "multi-key-sync", + Key: "k1\nk2", + Status: common.ChannelStatusEnabled, + Models: "sync-model", + Group: "default", + ChannelInfo: ChannelInfo{ + IsMultiKey: true, + MultiKeySize: 2, + MultiKeyMode: constant.MultiKeyModeRandom, + }, + } + require.NoError(t, DB.Create(&channel).Error) + require.NoError(t, channel.AddAbilities(nil)) + + // 禁用全部密钥后,渠道状态回算为禁用,并退出 abilities 选择池 + channel.ChannelInfo.MultiKeyStatusList = map[int]int{0: common.ChannelStatusManuallyDisabled, 1: common.ChannelStatusManuallyDisabled} + channel.SyncMultiKeyChannelStatus() + require.NoError(t, channel.Update()) + + var stored Channel + require.NoError(t, DB.First(&stored, channel.Id).Error) + assert.Equal(t, common.ChannelStatusAutoDisabled, stored.Status) + assert.Equal(t, allKeysDisabledReason, stored.GetOtherInfo()["status_reason"]) + var ability Ability + require.NoError(t, DB.First(&ability, "channel_id = ?", channel.Id).Error) + assert.False(t, ability.Enabled) + + // 重新启用一个密钥后,渠道随密钥恢复可用 + delete(channel.ChannelInfo.MultiKeyStatusList, 0) + channel.SyncMultiKeyChannelStatus() + require.NoError(t, channel.Update()) + + require.NoError(t, DB.First(&stored, channel.Id).Error) + assert.Equal(t, common.ChannelStatusEnabled, stored.Status) + require.NoError(t, DB.First(&ability, "channel_id = ?", channel.Id).Error) + assert.True(t, ability.Enabled) +} + func TestSaveStatusStateFromSingleKeySnapshotPreservesUnownedColumns(t *testing.T) { setupChannelStatusTest(t) diff --git a/model/user.go b/model/user.go index 7bc060ad1bf8..1d38ea392ad5 100644 --- a/model/user.go +++ b/model/user.go @@ -1334,6 +1334,10 @@ func DeltaUpdateUserQuota(id int, delta int) (err error) { //} func GetRootUser() (user *User) { + // 异步通知等后台任务可能在 DB 尚未初始化或已被回收后调用,此时视为无 root 用户 + if DB == nil { + return &User{} + } DB.Where("role = ?", common.RoleRootUser).First(&user) return user } diff --git a/service/channel.go b/service/channel.go index 81de54c609bd..294ba0e22a2c 100644 --- a/service/channel.go +++ b/service/channel.go @@ -3,6 +3,7 @@ package service import ( "fmt" "regexp" + "strconv" "strings" "time" @@ -116,7 +117,60 @@ func DetectKeyLimit(channelInfo model.ChannelInfo, errMessage string) (matched b return true, parsed.Unix(), fmt.Sprintf("%s (reset at %s)", pattern.Name, resetCapture) } } + if cycleReset, ok := nextCycleResetTime(pattern.ResetCycle, time.Now()); ok { + return true, cycleReset.Unix(), fmt.Sprintf("%s (reset cycle %s)", pattern.Name, pattern.ResetCycle) + } return true, now + int64(fallbackMinutes)*60, fmt.Sprintf("%s (fallback %d min)", pattern.Name, fallbackMinutes) } return false, 0, "" } + +// nextCycleResetTime 计算周期重置规则在 now 之后的下一个边界(本地时区 0 点): +// "daily" 为次日 0 点;"weekly:N" 为下一个周 N(1=周一 .. 7=周日)的 0 点; +// "monthly:N" 为下一个月中的 N 号 0 点,当月不存在 N 号(如 2 月的 31 号)时顺延到 +// 下月 1 号。规则为空或无法解析时返回 false,由调用方回退到 DefaultMinutes。 +func nextCycleResetTime(resetCycle string, now time.Time) (time.Time, bool) { + if resetCycle == "" { + return time.Time{}, false + } + now = now.In(time.Local) + midnight := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.Local) + firstOfNextMonth := func(t time.Time) time.Time { + return time.Date(t.Year(), t.Month(), 1, 0, 0, 0, 0, time.Local).AddDate(0, 1, 0) + } + dayOfMonthIn := func(firstOfMonth time.Time, day int) time.Time { + candidate := time.Date(firstOfMonth.Year(), firstOfMonth.Month(), day, 0, 0, 0, 0, time.Local) + if candidate.Day() != day { + return firstOfNextMonth(firstOfMonth) + } + return candidate + } + switch { + case resetCycle == "daily": + return midnight.AddDate(0, 0, 1), true + case strings.HasPrefix(resetCycle, "weekly:"): + day, err := strconv.Atoi(strings.TrimPrefix(resetCycle, "weekly:")) + if err != nil || day < 1 || day > 7 { + return time.Time{}, false + } + // Go 的 Weekday 以周日为 0,将 1=周一..7=周日 转换为 Go 编号 + targetGoDay := day % 7 + next := midnight.AddDate(0, 0, 1) + for next.Weekday() != time.Weekday(targetGoDay) { + next = next.AddDate(0, 0, 1) + } + return next, true + case strings.HasPrefix(resetCycle, "monthly:"): + day, err := strconv.Atoi(strings.TrimPrefix(resetCycle, "monthly:")) + if err != nil || day < 1 || day > 31 { + return time.Time{}, false + } + thisMonth := time.Date(now.Year(), now.Month(), 1, 0, 0, 0, 0, time.Local) + candidate := dayOfMonthIn(thisMonth, day) + if !candidate.After(now) { + candidate = dayOfMonthIn(firstOfNextMonth(now), day) + } + return candidate, true + } + return time.Time{}, false +} diff --git a/service/channel_limit_test.go b/service/channel_limit_test.go index 4199d0ee18a4..dde61b220671 100644 --- a/service/channel_limit_test.go +++ b/service/channel_limit_test.go @@ -106,3 +106,97 @@ func TestDetectKeyLimitSkipsInvalidRegex(t *testing.T) { require.True(t, matched) assert.Contains(t, reason, "good") } + +func TestNextCycleResetTime(t *testing.T) { + // 2026-09-10 是周四 + now := time.Date(2026, 9, 10, 12, 30, 0, 0, time.Local) + local := func(year int, month time.Month, day int) time.Time { + return time.Date(year, month, day, 0, 0, 0, 0, time.Local) + } + + cases := []struct { + name string + cycle string + want time.Time + }{ + {"daily to next midnight", "daily", local(2026, 9, 11)}, + {"weekly to next friday", "weekly:5", local(2026, 9, 11)}, + {"weekly same weekday rolls to next week", "weekly:4", local(2026, 9, 17)}, + {"weekly to next sunday", "weekly:7", local(2026, 9, 13)}, + {"monthly later this month", "monthly:15", local(2026, 9, 15)}, + {"monthly same day rolls to next month", "monthly:10", local(2026, 10, 10)}, + {"monthly last day of month", "monthly:30", local(2026, 9, 30)}, + {"monthly day beyond month end rolls to next month 1st", "monthly:31", local(2026, 10, 1)}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + got, ok := nextCycleResetTime(c.cycle, now) + require.True(t, ok) + assert.True(t, got.After(now), "reset time must be in the future") + assert.Equal(t, c.want, got) + }) + } + + for _, cycle := range []string{"", "garbage", "weekly:0", "weekly:8", "monthly:0", "monthly:32", "weekly:x", "monthly:"} { + _, ok := nextCycleResetTime(cycle, now) + assert.False(t, ok, "cycle %q should be rejected", cycle) + } +} + +func TestDetectKeyLimitUsesResetCycle(t *testing.T) { + info := model.ChannelInfo{ + MultiKeyLimitPatterns: []model.LimitPattern{ + { + Name: "aliyun token-plan", + Regex: `Your token-plan quota has been exhausted\.`, + ResetCycle: "monthly:1", + }, + }, + } + matched, cooldownUntil, reason := DetectKeyLimit(info, "Your token-plan quota has been exhausted.") + require.True(t, matched) + assert.Contains(t, reason, "reset cycle monthly:1") + // 无论何时触发,重置时间都必须落在未来某月 1 号 0 点(本地时区) + reset := time.Unix(cooldownUntil, 0).In(time.Local) + assert.True(t, reset.After(time.Now())) + assert.Equal(t, 1, reset.Day()) + assert.Equal(t, 0, reset.Hour()) + assert.Equal(t, 0, reset.Minute()) +} + +func TestDetectKeyLimitInvalidCycleFallsBackToMinutes(t *testing.T) { + info := model.ChannelInfo{ + MultiKeyLimitPatterns: []model.LimitPattern{ + { + Name: "bad cycle", + Regex: `quota exhausted`, + ResetCycle: "garbage", + DefaultMinutes: 15, + }, + }, + } + matched, cooldownUntil, reason := DetectKeyLimit(info, "quota exhausted again") + require.True(t, matched) + assert.Contains(t, reason, "fallback") + assert.InDelta(t, common.GetTimestamp()+15*60, cooldownUntil, 3) +} + +func TestDetectKeyLimitCapturedDateWinsOverCycle(t *testing.T) { + info := model.ChannelInfo{ + MultiKeyLimitPatterns: []model.LimitPattern{ + { + Name: "captured", + Regex: `limit (?P\d{4}-\d{2}-\d{2} \d{2}:\d{2}:\d{2})`, + DateLayout: "2006-01-02 15:04:05", + ResetCycle: "monthly:1", + DefaultMinutes: 10, + }, + }, + } + resetTime := time.Now().In(time.Local).Add(2 * time.Hour).Truncate(time.Second) + msg := "limit " + resetTime.Format("2006-01-02 15:04:05") + matched, cooldownUntil, reason := DetectKeyLimit(info, msg) + require.True(t, matched) + assert.Contains(t, reason, "reset at") + assert.InDelta(t, resetTime.Unix(), cooldownUntil, 2) +} diff --git a/service/notify-limit.go b/service/notify-limit.go index cad5d7bc1824..23cb52ada34b 100644 --- a/service/notify-limit.go +++ b/service/notify-limit.go @@ -48,7 +48,8 @@ func startCleanupTask() { // CheckNotificationLimit checks if the user has exceeded their notification limit // Returns true if the user can send notification, false if limit exceeded func CheckNotificationLimit(userId int, notifyType string) (bool, error) { - if common.RedisEnabled { + // 后台通知协程可能在 Redis 初始化前/回收后执行,此时回落到内存限流 + if common.RedisEnabled && common.RDB != nil { return checkRedisLimit(userId, notifyType) } return checkMemoryLimit(userId, notifyType) diff --git a/web/src/features/channels/components/limit-patterns-editor.tsx b/web/src/features/channels/components/limit-patterns-editor.tsx index cd4e42a5054b..27df7340894f 100644 --- a/web/src/features/channels/components/limit-patterns-editor.tsx +++ b/web/src/features/channels/components/limit-patterns-editor.tsx @@ -30,8 +30,12 @@ import { import { LIMIT_PATTERN_PRESETS, PREDEFINED_DATE_LAYOUTS, + formatResetCycle, + parseResetCycle, + resetCycleMaxDay, validateLimitPatternRegex, } from '../lib/limit-pattern-utils' +import type { ResetCycleType } from '../lib/limit-pattern-utils' import type { LimitPattern } from '../types' type LimitPatternsEditorProps = { @@ -91,8 +95,9 @@ export function LimitPatternsEditor(props: LimitPatternsEditorProps) { { name: '', regex: '', - date_layout: '2006-01-02 15:04:05', + date_layout: '', default_minutes: 10, + reset_cycle: '', }, ], [...keysRef.current, nextKey()] @@ -125,7 +130,11 @@ export function LimitPatternsEditor(props: LimitPatternsEditorProps) { {props.value.map((pattern, index) => { const validation = validateLimitPatternRegex(pattern.regex) - const usingCustomLayout = !isPredefinedLayout(pattern.date_layout) + // 日期布局仅对带 (?P...) 捕获组的正则有意义;无捕获组的 + // 模式(如阿里云额度耗尽)只用周期或默认分钟数冷却。 + const usesResetCapture = pattern.regex.includes('(?P') + const usingCustomLayout = usesResetCapture && !isPredefinedLayout(pattern.date_layout) + const cycle = parseResetCycle(pattern.reset_cycle) return (
...) group')} + placeholder={t('Regex with optional (?P...) group')} className='font-mono text-sm' onChange={(e) => updatePattern(index, { regex: e.target.value })} aria-invalid={!validation.valid} @@ -145,47 +154,106 @@ export function LimitPatternsEditor(props: LimitPatternsEditorProps) { {t(validation.error ?? '')}

)} + {usesResetCapture && ( +
+ { + const next = e.target.value + updatePattern(index, { + date_layout: next === 'custom' ? '' : next, + }) + }} + > + {PREDEFINED_DATE_LAYOUTS.map((layout) => ( + + {layout.label} + + ))} + + {usingCustomLayout && ( + + updatePattern(index, { date_layout: e.target.value }) + } + /> + )} +
+ )} + {usesResetCapture && !pattern.date_layout.trim() && ( +

+ {t( + 'A date layout is required to parse the captured reset time, otherwise the reset cycle or fallback minutes apply' + )} +

+ )}
+ + updatePattern(index, { + default_minutes: Number(e.target.value), + }) + } + /> { - const next = e.target.value + const type = e.target.value as ResetCycleType updatePattern(index, { - date_layout: next === 'custom' ? '' : next, + reset_cycle: formatResetCycle(type, cycle.day), }) }} > - {PREDEFINED_DATE_LAYOUTS.map((layout) => ( - - {layout.label} - - ))} + + {t('No cycle')} + + + {t('Daily')} + + + {t('Weekly')} + + + {t('Monthly')} + - {usingCustomLayout && ( + {(cycle.type === 'weekly' || cycle.type === 'monthly') && ( - updatePattern(index, { date_layout: e.target.value }) + updatePattern(index, { + reset_cycle: formatResetCycle( + cycle.type, + Number(e.target.value) + ), + }) } /> )} - - updatePattern(index, { - default_minutes: Number(e.target.value), - }) - } - />