From 7c62382d044d1f9f0abefa494d52d6b947cf302f Mon Sep 17 00:00:00 2001 From: IanShaw027 Date: Fri, 7 Aug 2026 14:53:18 +0800 Subject: [PATCH] =?UTF-8?q?feat(grok):=20=E5=90=B8=E6=94=B6=E8=B0=83?= =?UTF-8?q?=E5=BA=A6=E9=98=88=E5=80=BC=E3=80=81=E9=85=8D=E9=A2=9D=E8=A7=A3?= =?UTF-8?q?=E6=9E=90=E3=80=81=E6=89=B9=E9=87=8F=E7=94=A8=E9=87=8F=E4=B8=8E?= =?UTF-8?q?=20CLI=20=E8=BA=AB=E4=BB=BD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 从 personal-dev 取优并接线到 feat/grok-complete-integration: - 调度阈值:grok_sched_* 写入、ApplyAccountSchedulingThreshold、 account_scheduling_thresholds 设置与 Settings UI、网关过滤 - 配额头:x-rate-limit 别名、相对秒 reset、tier/entitlement 扩展 - 批量 usage:GetUsageBatch + POST /usage/batch + AccountsView 合并加载 - CLI 身份:xai.cli_identity 统一 pin/UA;service 层 applyGrokCLIHeaders 与 transport 最终改写对齐;media eligibility 模块落地 - UsageCell:月度 cents→USD 与 30d 进度条展示 保留既有 failure taxonomy / sticky / free recovery 等 HEAD cools。 --- .../internal/handler/admin/account_handler.go | 35 ++ .../internal/handler/admin/setting_handler.go | 3 +- .../handler/admin/setting_handler_audit.go | 24 + .../handler/admin/setting_handler_update.go | 15 +- backend/internal/handler/dto/settings.go | 3 + backend/internal/pkg/xai/billing.go | 6 +- backend/internal/pkg/xai/cli_identity.go | 79 +++ backend/internal/pkg/xai/cli_identity_test.go | 67 +++ backend/internal/pkg/xai/quota.go | 72 ++- backend/internal/pkg/xai/quota_test.go | 99 ++++ backend/internal/repository/http_upstream.go | 24 +- backend/internal/server/routes/admin.go | 1 + .../service/account_grok_media_eligibility.go | 65 +++ .../account_grok_media_eligibility_test.go | 42 +- .../account_scheduling_threshold_eval.go | 464 ++++++++++++++++++ .../account_scheduling_threshold_eval_test.go | 308 ++++++++++++ ...t_scheduling_threshold_integration_test.go | 136 +++++ .../account_scheduling_threshold_reason.go | 188 +++++++ ...ccount_scheduling_threshold_reason_test.go | 76 +++ ...t_scheduling_threshold_snapshot_cleanup.go | 18 + ...eduling_threshold_snapshot_cleanup_test.go | 21 + .../internal/service/account_usage_service.go | 121 ++++- .../account_usage_service_batch_test.go | 191 +++++++ backend/internal/service/admin_account.go | 58 +-- backend/internal/service/domain_constants.go | 15 + .../internal/service/gateway_scheduling.go | 43 +- .../service/grok_model_quota_block.go | 3 +- backend/internal/service/grok_p2_test.go | 15 + .../internal/service/grok_team_rate_limit.go | 3 +- .../service/grok_team_rate_limit_test.go | 16 + .../internal/service/grok_upstream_headers.go | 78 +++ .../service/grok_upstream_headers_test.go | 86 ++++ .../service/openai_account_scheduler.go | 11 +- .../internal/service/openai_gateway_grok.go | 118 ++++- .../service/openai_gateway_grok_test.go | 33 +- .../service/openai_gateway_scheduling.go | 41 +- .../service/openai_gateway_service_test.go | 23 + .../service/openai_ws_http_bridge_test.go | 2 +- backend/internal/service/ratelimit_service.go | 91 ++++ ...limit_service_scheduling_threshold_test.go | 166 +++++++ backend/internal/service/setting_features.go | 57 +++ .../service/setting_gateway_runtime.go | 13 + backend/internal/service/setting_parse.go | 8 + ...setting_service_platform_threshold_test.go | 151 ++++++ backend/internal/service/setting_update.go | 90 ++++ backend/internal/service/settings_view.go | 3 + frontend/src/api/admin/accounts.ts | 14 + frontend/src/api/admin/settings.ts | 36 ++ .../components/account/AccountUsageCell.vue | 166 ++++++- .../src/i18n/locales/en/admin/accounts.ts | 2 + .../src/i18n/locales/en/admin/settings.ts | 7 +- .../src/i18n/locales/zh/admin/accounts.ts | 2 + .../src/i18n/locales/zh/admin/settings.ts | 7 +- frontend/src/views/admin/AccountsView.vue | 200 +++++++- frontend/src/views/admin/SettingsView.vue | 88 ++++ 55 files changed, 3542 insertions(+), 162 deletions(-) create mode 100644 backend/internal/pkg/xai/cli_identity.go create mode 100644 backend/internal/pkg/xai/cli_identity_test.go create mode 100644 backend/internal/service/account_grok_media_eligibility.go create mode 100644 backend/internal/service/account_scheduling_threshold_eval.go create mode 100644 backend/internal/service/account_scheduling_threshold_eval_test.go create mode 100644 backend/internal/service/account_scheduling_threshold_integration_test.go create mode 100644 backend/internal/service/account_scheduling_threshold_reason.go create mode 100644 backend/internal/service/account_scheduling_threshold_reason_test.go create mode 100644 backend/internal/service/account_scheduling_threshold_snapshot_cleanup.go create mode 100644 backend/internal/service/account_scheduling_threshold_snapshot_cleanup_test.go create mode 100644 backend/internal/service/account_usage_service_batch_test.go create mode 100644 backend/internal/service/grok_upstream_headers.go create mode 100644 backend/internal/service/grok_upstream_headers_test.go create mode 100644 backend/internal/service/ratelimit_service_scheduling_threshold_test.go create mode 100644 backend/internal/service/setting_service_platform_threshold_test.go diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index 9b80dd05d1..3b3283e1aa 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -2422,6 +2422,11 @@ type BatchTodayStatsRequest struct { AccountIDs []int64 `json:"account_ids" binding:"required"` } +type BatchUsageRequest struct { + AccountIDs []int64 `json:"account_ids" binding:"required"` + Force bool `json:"force"` +} + // GetBatchTodayStats 批量获取多个账号的今日统计。 // POST /api/v1/admin/accounts/today-stats/batch func (h *AccountHandler) GetBatchTodayStats(c *gin.Context) { @@ -2468,6 +2473,36 @@ func (h *AccountHandler) GetBatchTodayStats(c *gin.Context) { response.Success(c, payload) } +// GetBatchUsage 批量获取多个账号的 current usage。 +// POST /api/v1/admin/accounts/usage/batch +func (h *AccountHandler) GetBatchUsage(c *gin.Context) { + var req BatchUsageRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + + accountIDs := normalizeInt64IDList(req.AccountIDs) + if len(accountIDs) == 0 { + response.Success(c, gin.H{ + "usage": map[string]any{}, + "errors": map[string]string{}, + }) + return + } + + usageByAccount, errorsByAccount, err := h.accountUsageService.GetUsageBatch(c.Request.Context(), accountIDs, req.Force) + if err != nil { + response.ErrorFrom(c, err) + return + } + + response.Success(c, gin.H{ + "usage": usageByAccount, + "errors": errorsByAccount, + }) +} + // SetSchedulableRequest represents the request body for setting schedulable status type SetSchedulableRequest struct { Schedulable bool `json:"schedulable"` diff --git a/backend/internal/handler/admin/setting_handler.go b/backend/internal/handler/admin/setting_handler.go index 7ff2f87cc2..b7bbf223a8 100644 --- a/backend/internal/handler/admin/setting_handler.go +++ b/backend/internal/handler/admin/setting_handler.go @@ -383,7 +383,8 @@ func (h *SettingHandler) GetSettings(c *gin.Context) { AffiliateEnabled: settings.AffiliateEnabled, - AllowUserViewErrorRequests: settings.AllowUserViewErrorRequests, + AccountSchedulingThresholds: settings.AccountSchedulingThresholds, + AllowUserViewErrorRequests: settings.AllowUserViewErrorRequests, } // OpenAI fast policy (stored under a dedicated setting key) diff --git a/backend/internal/handler/admin/setting_handler_audit.go b/backend/internal/handler/admin/setting_handler_audit.go index 5594b6c9c1..e5370bf299 100644 --- a/backend/internal/handler/admin/setting_handler_audit.go +++ b/backend/internal/handler/admin/setting_handler_audit.go @@ -598,6 +598,9 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings, if !equalPlatformQuotaSettings(before.DefaultPlatformQuotas, after.DefaultPlatformQuotas) { changed = append(changed, service.SettingKeyDefaultPlatformQuotas) } + if !equalAccountSchedulingThresholds(before.AccountSchedulingThresholds, after.AccountSchedulingThresholds) { + changed = append(changed, service.SettingKeyAccountSchedulingThresholds) + } changed = appendAuthSourceDefaultChanges(changed, beforeAuthSourceDefaults, afterAuthSourceDefaults) return changed } @@ -811,6 +814,27 @@ func slotOf(s *service.DefaultPlatformQuotaSetting, win string) *float64 { } // equalPlatformQuotaSettings reports whether two platform-quota maps are identical across all allowed slots. +func equalAccountSchedulingThresholds(before, after map[string]int) bool { + for _, platform := range service.AllowedSchedulingThresholdPlatforms { + beforeValue := 100 + if before != nil { + if value, ok := before[platform]; ok { + beforeValue = value + } + } + afterValue := 100 + if after != nil { + if value, ok := after[platform]; ok { + afterValue = value + } + } + if beforeValue != afterValue { + return false + } + } + return true +} + func equalPlatformQuotaSettings(before, after map[string]*service.DefaultPlatformQuotaSetting) bool { for _, platform := range service.AllowedQuotaPlatforms { b := before[platform] diff --git a/backend/internal/handler/admin/setting_handler_update.go b/backend/internal/handler/admin/setting_handler_update.go index 34c084d82b..581cda6290 100644 --- a/backend/internal/handler/admin/setting_handler_update.go +++ b/backend/internal/handler/admin/setting_handler_update.go @@ -358,6 +358,9 @@ type UpdateSettingsRequest struct { // 系统全局 platform quota 默认值(整体替换语义:nil = 不修改,non-nil = 整体覆盖)。 DefaultPlatformQuotas map[string]*service.DefaultPlatformQuotaSetting `json:"default_platform_quotas"` + // 各平台账号自动停调阈值(整体替换语义:nil = 不修改,non-nil = 整体覆盖)。 + AccountSchedulingThresholds map[string]int `json:"account_scheduling_thresholds"` + // auth-source 层 platform quota 覆盖(override 语义:nil = 不修改,non-nil = 整体覆盖该 source 的 quota 配置)。 AuthSourceEmailPlatformQuotas map[string]*service.DefaultPlatformQuotaSetting `json:"auth_source_default_email_platform_quotas"` AuthSourceLinuxDoPlatformQuotas map[string]*service.DefaultPlatformQuotaSetting `json:"auth_source_default_linuxdo_platform_quotas"` @@ -1480,7 +1483,8 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { settings := &service.SystemSettings{ // 系统全局 platform quota 默认值(整体替换语义) - DefaultPlatformQuotas: req.DefaultPlatformQuotas, + DefaultPlatformQuotas: req.DefaultPlatformQuotas, + AccountSchedulingThresholds: req.AccountSchedulingThresholds, RegistrationEnabled: req.RegistrationEnabled, EmailVerifyEnabled: req.EmailVerifyEnabled, @@ -2320,10 +2324,11 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { AffiliateEnabled: updatedSettings.AffiliateEnabled, - RiskControlEnabled: updatedSettings.RiskControlEnabled, - CyberSessionBlockEnabled: updatedSettings.CyberSessionBlockEnabled, - CyberSessionBlockTTLSeconds: updatedSettings.CyberSessionBlockTTLSeconds, - AllowUserViewErrorRequests: updatedSettings.AllowUserViewErrorRequests, + RiskControlEnabled: updatedSettings.RiskControlEnabled, + CyberSessionBlockEnabled: updatedSettings.CyberSessionBlockEnabled, + CyberSessionBlockTTLSeconds: updatedSettings.CyberSessionBlockTTLSeconds, + AccountSchedulingThresholds: updatedSettings.AccountSchedulingThresholds, + AllowUserViewErrorRequests: updatedSettings.AllowUserViewErrorRequests, } if fastPolicy, err := h.settingService.GetOpenAIFastPolicySettings(c.Request.Context()); err != nil { slog.Error("openai_fast_policy_settings_get_failed", "error", err) diff --git a/backend/internal/handler/dto/settings.go b/backend/internal/handler/dto/settings.go index 75e7972a4d..ed1d81a193 100644 --- a/backend/internal/handler/dto/settings.go +++ b/backend/internal/handler/dto/settings.go @@ -331,6 +331,9 @@ type SystemSettings struct { // 系统全局默认平台配额(key = platform,nil/缺省 = 不限制) DefaultPlatformQuotas map[string]*service.DefaultPlatformQuotaSetting `json:"default_platform_quotas,omitempty"` + // 系统全局账号自动停调阈值(key = platform,100 = disabled) + AccountSchedulingThresholds map[string]int `json:"account_scheduling_thresholds,omitempty"` + // 允许终端用户在用量页查看自己的失败请求 AllowUserViewErrorRequests bool `json:"allow_user_view_error_requests"` } diff --git a/backend/internal/pkg/xai/billing.go b/backend/internal/pkg/xai/billing.go index 3c620a9884..6eb74aa293 100644 --- a/backend/internal/pkg/xai/billing.go +++ b/backend/internal/pkg/xai/billing.go @@ -20,7 +20,9 @@ const ( // one bump here covers OAuth traffic and billing probes together. // Keep in sync with https://x.ai/cli/stable. CLIClientVersion = "0.2.114" - CLIUserAgent = "grok-pager/" + CLIClientVersion + " grok-shell/" + CLIClientVersion + " (macos; aarch64)" + // billingCLIUserAgent is the legacy pager/shell UA used by billing probes. + // Distinct from CLIUserAgent() in cli_identity.go (workspace-style UA). + billingCLIUserAgent = "grok-pager/" + CLIClientVersion + " grok-shell/" + CLIClientVersion + " (macos; aarch64)" BillingWeeklyPath = "/billing?format=credits" BillingMonthlyPath = "/billing" @@ -127,7 +129,7 @@ func ApplyCLIBillingHeaders(req *http.Request, accessToken string) { req.Header.Set("Content-Type", "application/json") req.Header.Set(CLITokenAuthHeader, CLITokenAuthValue) req.Header.Set(CLIClientVersionHeader, CLIClientVersion) - req.Header.Set("User-Agent", CLIUserAgent) + req.Header.Set("User-Agent", billingCLIUserAgent) } // ParseBillingPayload unmarshals a billing API response body. diff --git a/backend/internal/pkg/xai/cli_identity.go b/backend/internal/pkg/xai/cli_identity.go new file mode 100644 index 0000000000..480e2e40af --- /dev/null +++ b/backend/internal/pkg/xai/cli_identity.go @@ -0,0 +1,79 @@ +package xai + +import ( + "net/http" + "os" + "strings" + + "golang.org/x/mod/semver" +) + +// Fixed Grok Build / CLI-chat-proxy client identity. +// These values are intentionally pinned in-binary (not scraped from live CLI). +// Operators may bump the version via XAI_GROK_CLI_VERSION without a release. +const ( + // CLIProxyHost is the hostname that requires the official CLI identity headers. + CLIProxyHost = "cli-chat-proxy.grok.com" + + // CLIStableVersion is the known-good minimum client version accepted by cli-chat-proxy. + CLIStableVersion = "0.2.93" + + // CLIVersionEnv is the optional operator override for CLIStableVersion. + CLIVersionEnv = "XAI_GROK_CLI_VERSION" + + // CLITokenAuth is required by cli-chat-proxy for Grok Build OAuth tokens. + CLITokenAuth = "xai-grok-cli" + + // CLIClientIdentifier is the x-grok-client-identifier value used by Grok shell/CLI. + CLIClientIdentifier = "grok-shell" + + // CLIClientMode is used by billing / quota probes on the CLI surface. + CLIClientMode = "cli" +) + +// ResolveCLIVersion returns a supported CLI client version. +// Empty or invalid overrides fall back to CLIClientVersion (the pinned +// preferred client pin in billing.go). CLIStableVersion is only the minimum +// accepted by IsSupportedCLIVersion, not the default identity we advertise. +func ResolveCLIVersion() string { + version := strings.TrimSpace(os.Getenv(CLIVersionEnv)) + if !IsSupportedCLIVersion(version) { + return CLIClientVersion + } + return version +} + +// IsSupportedCLIVersion reports whether version is a valid semver string at or +// above CLIStableVersion (prereleases below a higher release are rejected when +// they compare less than the stable pin). +func IsSupportedCLIVersion(version string) bool { + canonical := "v" + version + minimum := "v" + CLIStableVersion + return semver.IsValid(canonical) && + semver.Canonical(canonical) == canonical && + semver.Compare(canonical, minimum) >= 0 +} + +// CLIUserAgent builds the workspace-style User-Agent for a CLI client version. +func CLIUserAgent(version string) string { + if strings.TrimSpace(version) == "" { + version = CLIClientVersion + } + return "xai-grok-workspace/" + version +} + +// ApplyCLIProxyHeaders stamps the fixed Grok CLI identity when the request +// targets cli-chat-proxy. Direct api.x.ai traffic is left unchanged. +func ApplyCLIProxyHeaders(req *http.Request) { + if req == nil || req.URL == nil || !strings.EqualFold(strings.TrimSpace(req.URL.Hostname()), CLIProxyHost) { + return + } + if req.Header == nil { + req.Header = make(http.Header) + } + version := ResolveCLIVersion() + req.Header.Set("X-XAI-Token-Auth", CLITokenAuth) + req.Header.Set("x-grok-client-version", version) + req.Header.Set("x-grok-client-identifier", CLIClientIdentifier) + req.Header.Set("User-Agent", CLIUserAgent(version)) +} diff --git a/backend/internal/pkg/xai/cli_identity_test.go b/backend/internal/pkg/xai/cli_identity_test.go new file mode 100644 index 0000000000..521df7483d --- /dev/null +++ b/backend/internal/pkg/xai/cli_identity_test.go @@ -0,0 +1,67 @@ +package xai + +import ( + "net/http" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestResolveCLIVersionDefaultsToPinnedClientVersion(t *testing.T) { + t.Setenv(CLIVersionEnv, "") + // Default advertise pin is CLIClientVersion; CLIStableVersion is only the floor. + require.Equal(t, CLIClientVersion, ResolveCLIVersion()) + require.True(t, IsSupportedCLIVersion(CLIClientVersion)) + require.True(t, IsSupportedCLIVersion(CLIStableVersion)) +} + +func TestResolveCLIVersionAcceptsValidOverride(t *testing.T) { + t.Setenv(CLIVersionEnv, "0.2.95-alpha.1") + require.Equal(t, "0.2.95-alpha.1", ResolveCLIVersion()) +} + +func TestResolveCLIVersionRejectsUnsafeOrTooOld(t *testing.T) { + for _, version := range []string{ + "0.2.92", + "0.2.93-beta.1", + "0.2.95\r\nX-Injected: true", + "0.2.093", + "0.3", + "1", + } { + t.Run(version, func(t *testing.T) { + t.Setenv(CLIVersionEnv, version) + require.Equal(t, CLIClientVersion, ResolveCLIVersion()) + }) + } +} + +func TestApplyCLIProxyHeaders(t *testing.T) { + t.Setenv(CLIVersionEnv, "") + + req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) + require.NoError(t, err) + req.Header.Set("User-Agent", "sub2api-grok/1.0") + + ApplyCLIProxyHeaders(req) + + require.Equal(t, CLIClientVersion, req.Header.Get("x-grok-client-version")) + require.Equal(t, CLIClientIdentifier, req.Header.Get("x-grok-client-identifier")) + require.Equal(t, CLITokenAuth, req.Header.Get("X-XAI-Token-Auth")) + require.Equal(t, CLIUserAgent(CLIClientVersion), req.Header.Get("User-Agent")) +} + +func TestApplyCLIProxyHeadersLeavesAPIHostUnchanged(t *testing.T) { + t.Setenv(CLIVersionEnv, "0.2.95") + + req, err := http.NewRequest(http.MethodPost, "https://api.x.ai/v1/responses", nil) + require.NoError(t, err) + req.Header.Set("User-Agent", "sub2api-grok/1.0") + + ApplyCLIProxyHeaders(req) + + require.Empty(t, req.Header.Get("x-grok-client-version")) + require.Empty(t, req.Header.Get("x-grok-client-identifier")) + require.Empty(t, req.Header.Get("X-XAI-Token-Auth")) + require.Equal(t, "sub2api-grok/1.0", req.Header.Get("User-Agent")) +} diff --git a/backend/internal/pkg/xai/quota.go b/backend/internal/pkg/xai/quota.go index fe269d1d84..3802b4fa4f 100644 --- a/backend/internal/pkg/xai/quota.go +++ b/backend/internal/pkg/xai/quota.go @@ -61,11 +61,27 @@ var quotaHeaderAllowlist = []string{ "x-ratelimit-limit-tokens", "x-ratelimit-remaining-tokens", "x-ratelimit-reset-tokens", + "x-rate-limit-limit-requests", + "x-rate-limit-remaining-requests", + "x-rate-limit-reset-requests", + "x-rate-limit-limit-tokens", + "x-rate-limit-remaining-tokens", + "x-rate-limit-reset-tokens", "retry-after", "x-subscription-tier", "xai-subscription-tier", + "x-xai-subscription-tier", + "x-xai-user-tier", + "xai-user-tier", + "xai-tier", + "x-user-tier", + "x-plan-tier", + "x-subscription-plan", "x-entitlement-status", "xai-entitlement-status", + "x-xai-entitlement-status", + "x-xai-user-entitlement-status", + "x-user-entitlement-status", } func ParseQuotaHeaders(headers http.Header, statusCode int) *QuotaSnapshot { @@ -95,8 +111,24 @@ func parseQuotaHeaders(headers http.Header, statusCode int, source string, keepE if retryAfter := parseRetryAfter(headers.Get("retry-after")); retryAfter != nil { snapshot.RetryAfterSeconds = retryAfter } - snapshot.SubscriptionTier = firstHeader(headers, "xai-subscription-tier", "x-subscription-tier") - snapshot.EntitlementStatus = firstHeader(headers, "xai-entitlement-status", "x-entitlement-status") + snapshot.SubscriptionTier = firstHeader(headers, + "xai-subscription-tier", + "x-subscription-tier", + "x-xai-subscription-tier", + "x-xai-user-tier", + "xai-user-tier", + "xai-tier", + "x-user-tier", + "x-plan-tier", + "x-subscription-plan", + ) + snapshot.EntitlementStatus = firstHeader(headers, + "xai-entitlement-status", + "x-entitlement-status", + "x-xai-entitlement-status", + "x-xai-user-entitlement-status", + "x-user-entitlement-status", + ) for _, name := range quotaHeaderAllowlist { if value := strings.TrimSpace(headers.Get(name)); value != "" { @@ -121,11 +153,23 @@ func parseQuotaHeaders(headers http.Header, statusCode int, source string, keepE } func parseQuotaWindow(headers http.Header, dimension string) *QuotaWindow { + limitHeader := firstHeader(headers, + "x-ratelimit-limit-"+dimension, + "x-rate-limit-limit-"+dimension, + ) + remainingHeader := firstHeader(headers, + "x-ratelimit-remaining-"+dimension, + "x-rate-limit-remaining-"+dimension, + ) + resetHeader := firstHeader(headers, + "x-ratelimit-reset-"+dimension, + "x-rate-limit-reset-"+dimension, + ) window := &QuotaWindow{ - Limit: parseInt64Ptr(headers.Get("x-ratelimit-limit-" + dimension)), - Remaining: parseInt64Ptr(headers.Get("x-ratelimit-remaining-" + dimension)), + Limit: parseInt64Ptr(limitHeader), + Remaining: parseInt64Ptr(remainingHeader), } - if reset := parseResetHeader(headers.Get("x-ratelimit-reset-" + dimension)); reset != nil { + if reset := parseResetHeader(resetHeader); reset != nil { window.ResetUnix = reset window.ResetAt = time.Unix(*reset, 0).UTC().Format(time.RFC3339) } @@ -141,11 +185,27 @@ func parseResetHeader(raw string) *int64 { return nil } if value, err := strconv.ParseInt(raw, 10, 64); err == nil { - if value > 1_000_000_000_000 { + // xAI (and OpenAI-compatible upstreams) may express the reset as a + // millisecond epoch, a second epoch, or a *relative* number of seconds + // until reset (e.g. "60"). Disambiguate by magnitude, mirroring the + // Kiro reset parser, so a relative "60" is not misread as 1970-01-01. + switch { + case value >= 1_000_000_000_000: // milliseconds epoch → seconds value = value / 1000 + case value >= 1_000_000_000: // already a plausible unix-seconds epoch (>= 2001-09) + // keep as-is + default: // relative seconds from now + value = time.Now().Unix() + value } return &value } + if duration, err := time.ParseDuration(raw); err == nil && duration > 0 { + if duration < time.Second { + duration = time.Second + } + value := time.Now().Add(duration).Unix() + return &value + } if t, err := time.Parse(time.RFC3339, raw); err == nil { value := t.Unix() return &value diff --git a/backend/internal/pkg/xai/quota_test.go b/backend/internal/pkg/xai/quota_test.go index 983a87e84f..aafa7486b9 100644 --- a/backend/internal/pkg/xai/quota_test.go +++ b/backend/internal/pkg/xai/quota_test.go @@ -5,6 +5,7 @@ package xai import ( "net/http" "testing" + "time" "github.com/stretchr/testify/require" ) @@ -41,6 +42,104 @@ func TestParseQuotaHeaders(t *testing.T) { require.NotContains(t, snapshot.Headers, "authorization") } +func TestParseQuotaHeadersAcceptsXAITierAliases(t *testing.T) { + t.Parallel() + + headers := http.Header{} + headers.Set("x-xai-user-tier", "supergrok-heavy") + headers.Set("x-xai-user-entitlement-status", "enabled") + + snapshot := ParseQuotaHeaders(headers, http.StatusOK) + require.NotNil(t, snapshot) + require.True(t, snapshot.HeadersObserved) + require.Equal(t, "supergrok-heavy", snapshot.SubscriptionTier) + require.Equal(t, "enabled", snapshot.EntitlementStatus) + require.Equal(t, "supergrok-heavy", snapshot.Headers["x-xai-user-tier"]) + require.Equal(t, "enabled", snapshot.Headers["x-xai-user-entitlement-status"]) +} + +func TestParseQuotaHeadersAcceptsRateLimitAliases(t *testing.T) { + t.Parallel() + + headers := http.Header{} + headers.Set("x-rate-limit-limit-tokens", "500000") + headers.Set("x-rate-limit-remaining-tokens", "100") + headers.Set("x-rate-limit-reset-tokens", "1893456000") + + snapshot := ParseQuotaHeaders(headers, http.StatusOK) + require.NotNil(t, snapshot) + require.NotNil(t, snapshot.Tokens) + require.Equal(t, int64(500000), *snapshot.Tokens.Limit) + require.Equal(t, int64(100), *snapshot.Tokens.Remaining) + require.Equal(t, int64(1893456000), *snapshot.Tokens.ResetUnix) + require.Contains(t, snapshot.Headers, "x-rate-limit-limit-tokens") +} + +func TestParseResetHeaderRelativeSecondsNotMisreadAsEpoch(t *testing.T) { + t.Parallel() + + headers := http.Header{} + // xAI may return the reset window as a relative number of seconds ("60"). + // It must resolve to ~now+60s, NOT 1970-01-01 (epoch 60). + headers.Set("x-ratelimit-reset-requests", "60") + headers.Set("x-ratelimit-remaining-requests", "0") + + before := time.Now().Unix() + snapshot := ParseQuotaHeaders(headers, http.StatusTooManyRequests) + require.NotNil(t, snapshot) + require.NotNil(t, snapshot.Requests) + require.NotNil(t, snapshot.Requests.ResetUnix) + got := *snapshot.Requests.ResetUnix + require.GreaterOrEqual(t, got, before+59) + require.LessOrEqual(t, got, time.Now().Unix()+61) +} + +func TestParseResetHeaderDurationWindow(t *testing.T) { + t.Parallel() + + headers := http.Header{} + headers.Set("x-ratelimit-reset-requests", "6m0s") + headers.Set("x-ratelimit-remaining-requests", "0") + + before := time.Now().Unix() + snapshot := ParseQuotaHeaders(headers, http.StatusTooManyRequests) + require.NotNil(t, snapshot) + require.NotNil(t, snapshot.Requests) + require.NotNil(t, snapshot.Requests.ResetUnix) + got := *snapshot.Requests.ResetUnix + require.GreaterOrEqual(t, got, before+359) + require.LessOrEqual(t, got, time.Now().Unix()+361) +} + +func TestParseResetHeaderSubsecondDurationCeilsToFutureSecond(t *testing.T) { + t.Parallel() + + headers := http.Header{} + headers.Set("x-rate-limit-reset-tokens", "250ms") + headers.Set("x-rate-limit-remaining-tokens", "0") + + before := time.Now().Unix() + snapshot := ParseQuotaHeaders(headers, http.StatusTooManyRequests) + require.NotNil(t, snapshot) + require.NotNil(t, snapshot.Tokens) + require.NotNil(t, snapshot.Tokens.ResetUnix) + require.GreaterOrEqual(t, *snapshot.Tokens.ResetUnix, before) + require.LessOrEqual(t, *snapshot.Tokens.ResetUnix, time.Now().Unix()+2) +} + +func TestParseResetHeaderMillisecondsEpochNormalized(t *testing.T) { + t.Parallel() + + headers := http.Header{} + headers.Set("x-ratelimit-reset-tokens", "1893456000000") // ms epoch + headers.Set("x-ratelimit-remaining-tokens", "0") + + snapshot := ParseQuotaHeaders(headers, http.StatusTooManyRequests) + require.NotNil(t, snapshot) + require.NotNil(t, snapshot.Tokens) + require.Equal(t, int64(1893456000), *snapshot.Tokens.ResetUnix) +} + func TestParseQuotaHeadersReturnsNilForMissingHeaders(t *testing.T) { t.Parallel() diff --git a/backend/internal/repository/http_upstream.go b/backend/internal/repository/http_upstream.go index 68a986648a..9e996a43b8 100644 --- a/backend/internal/repository/http_upstream.go +++ b/backend/internal/repository/http_upstream.go @@ -23,6 +23,7 @@ import ( "github.com/andybalholm/brotli" "github.com/klauspost/compress/zstd" + "golang.org/x/mod/semver" "golang.org/x/net/http2" "github.com/Wei-Shaw/sub2api/internal/config" @@ -33,7 +34,6 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" - "golang.org/x/mod/semver" ) // 默认配置常量 @@ -73,12 +73,12 @@ const ( openAIHTTP2PingTimeout = 15 * time.Second // The Grok CLI proxy rejects requests that do not identify a supported - // client version. Keep a known-good stable version in the binary while - // allowing operators to bump it without waiting for a Sub2API release. - grokCLIProxyHost = "cli-chat-proxy.grok.com" + // client version. Host/env/version pins live in package xai so service, + // billing, and transport layers advertise the same identity. + grokCLIProxyHost = xai.CLIProxyHost grokOfficialAPIHost = "api.x.ai" - grokCLIStableVersion = xai.CLIClientVersion - grokCLIVersionOverride = "XAI_GROK_CLI_VERSION" + grokCLIStableVersion = xai.CLIClientVersion // preferred pin (not the minimum floor) + grokCLIVersionOverride = xai.CLIVersionEnv grokFallbackBodyLimit = 64 << 10 ) @@ -438,6 +438,11 @@ type prefixedReadCloser struct { // the final shared transport boundary. Keying this behavior to the exact CLI // proxy host keeps direct api.x.ai traffic unchanged and automatically covers // Responses, Chat Completions, media, quota probes, and account tests. +// +// Operator overrides must be >= CLIClientVersion (the preferred pin). Package +// xai.IsSupportedCLIVersion uses a lower floor (CLIStableVersion) for general +// validation; transport is stricter so we never silently advertise an older pin +// than the binary default. func applyGrokCLIProxyHeaders(req *http.Request) { if req == nil || req.URL == nil || !strings.EqualFold(strings.TrimSpace(req.URL.Hostname()), grokCLIProxyHost) { return @@ -449,14 +454,15 @@ func applyGrokCLIProxyHeaders(req *http.Request) { if !isSupportedGrokCLIVersion(version) { version = grokCLIStableVersion } - req.Header.Set("X-XAI-Token-Auth", "xai-grok-cli") + req.Header.Set("X-XAI-Token-Auth", xai.CLITokenAuth) req.Header.Set("x-grok-client-version", version) - req.Header.Set("User-Agent", "xai-grok-workspace/"+version) + req.Header.Set("x-grok-client-identifier", xai.CLIClientIdentifier) + req.Header.Set("User-Agent", xai.CLIUserAgent(version)) } func isSupportedGrokCLIVersion(version string) bool { canonical := "v" + version - minimum := "v" + grokCLIStableVersion + minimum := "v" + xai.CLIClientVersion return semver.IsValid(canonical) && semver.Canonical(canonical) == canonical && semver.Compare(canonical, minimum) >= 0 diff --git a/backend/internal/server/routes/admin.go b/backend/internal/server/routes/admin.go index 4af5f8f4c9..dcce82c413 100644 --- a/backend/internal/server/routes/admin.go +++ b/backend/internal/server/routes/admin.go @@ -379,6 +379,7 @@ func registerAccountRoutes(admin *gin.RouterGroup, h *handler.Handlers, stepUpAu accounts.POST("/:id/revert-proxy-fallback", h.Admin.Account.RevertProxyFallback) accounts.GET("/:id/usage", h.Admin.Account.GetUsage) accounts.GET("/:id/today-stats", h.Admin.Account.GetTodayStats) + accounts.POST("/usage/batch", h.Admin.Account.GetBatchUsage) accounts.POST("/today-stats/batch", h.Admin.Account.GetBatchTodayStats) accounts.POST("/:id/clear-rate-limit", h.Admin.Account.ClearRateLimit) accounts.POST("/:id/reset-quota", h.Admin.Account.ResetQuota) diff --git a/backend/internal/service/account_grok_media_eligibility.go b/backend/internal/service/account_grok_media_eligibility.go new file mode 100644 index 0000000000..5f05a92772 --- /dev/null +++ b/backend/internal/service/account_grok_media_eligibility.go @@ -0,0 +1,65 @@ +package service + +import infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + +// ValidateGrokMediaEligibilityExtra validates the optional per-account media +// routing override. A nil value removes the override and restores automatic +// provider-observation routing. +func ValidateGrokMediaEligibilityExtra(platform string, extra map[string]any) error { + if platform != PlatformGrok || extra == nil { + return nil + } + raw, exists := extra[GrokMediaEligibleExtraKey] + if !exists || raw == nil { + return nil + } + if _, ok := raw.(bool); !ok { + return infraerrors.BadRequest("GROK_MEDIA_ELIGIBILITY_INVALID", "grok_media_eligible must be a boolean or null") + } + return nil +} + +func normalizeGrokMediaEligibilityExtra(platform string, extra map[string]any) (map[string]any, error) { + if platform != PlatformGrok { + return extra, nil + } + if err := ValidateGrokMediaEligibilityExtra(platform, extra); err != nil { + return nil, err + } + if extra == nil { + return nil, nil + } + normalized := shallowCopyMap(extra) + if normalized[GrokMediaEligibleExtraKey] == nil { + delete(normalized, GrokMediaEligibleExtraKey) + } + return normalized, nil +} + +func normalizeGrokMediaEligibilityUpdateExtra(account *Account, input *UpdateAccountInput, normalized map[string]any) (map[string]any, error) { + if account == nil || account.Platform != PlatformGrok { + return normalized, nil + } + if input == nil { + return nil, infraerrors.BadRequest("INVALID_ACCOUNT_INPUT", "account update input is required") + } + if err := ValidateGrokMediaEligibilityExtra(account.Platform, input.Extra); err != nil { + return nil, err + } + if normalized == nil { + normalized = make(map[string]any) + } else { + normalized = shallowCopyMap(normalized) + } + raw, provided := input.Extra[GrokMediaEligibleExtraKey] + if provided { + if raw == nil { + delete(normalized, GrokMediaEligibleExtraKey) + } + return normalized, nil + } + if current, ok := account.Extra[GrokMediaEligibleExtraKey].(bool); ok { + normalized[GrokMediaEligibleExtraKey] = current + } + return normalized, nil +} diff --git a/backend/internal/service/account_grok_media_eligibility_test.go b/backend/internal/service/account_grok_media_eligibility_test.go index a79f9c02ec..be7b4759af 100644 --- a/backend/internal/service/account_grok_media_eligibility_test.go +++ b/backend/internal/service/account_grok_media_eligibility_test.go @@ -13,7 +13,6 @@ import ( ) func TestGrokMediaGenerationEligibility(t *testing.T) { - weeklyUsagePercent := 12.5 forbiddenBilling := &xai.BillingSummary{ StatusCode: http.StatusForbidden, WeeklyStatusCode: http.StatusForbidden, @@ -21,24 +20,9 @@ func TestGrokMediaGenerationEligibility(t *testing.T) { } weeklyAllowance := &xai.BillingSummary{ PeriodType: "weekly", - UsagePercent: &weeklyUsagePercent, StatusCode: http.StatusOK, WeeklyStatusCode: http.StatusOK, } - freeBilling := &xai.BillingSummary{ - PeriodType: "monthly", - StatusCode: http.StatusOK, - WeeklyStatusCode: http.StatusOK, - MonthlyStatusCode: http.StatusOK, - MonthlyUpdatedAt: "2026-07-17T00:00:00Z", - } - inconclusiveBilling := &xai.BillingSummary{ - StatusCode: http.StatusOK, - WeeklyStatusCode: http.StatusOK, - MonthlyStatusCode: http.StatusBadGateway, - Partial: true, - FailedWindows: []string{"monthly"}, - } weeklyForbidden := &xai.BillingSummary{ StatusCode: http.StatusOK, WeeklyStatusCode: http.StatusForbidden, @@ -59,14 +43,12 @@ func TestGrokMediaGenerationEligibility(t *testing.T) { {name: "nil account", account: nil, want: false, wantReason: "not_grok"}, {name: "non grok account", account: &Account{Platform: PlatformOpenAI}, want: false, wantReason: "not_grok"}, {name: "non oauth grok account stays eligible", account: &Account{Platform: PlatformGrok, Type: AccountTypeAPIKey}, want: true, wantReason: "non_oauth"}, - {name: "unobserved oauth fails closed", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth}, want: false, wantReason: "billing_unobserved"}, - {name: "weekly paid usage is eligible without inferring from period type", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: weeklyAllowance}}, want: true, wantReason: "eligible"}, - {name: "observed free account is rejected", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: freeBilling}}, want: false, wantReason: "billing_free_tier"}, - {name: "inconclusive billing fails closed", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: inconclusiveBilling}}, want: false, wantReason: "billing_inconclusive"}, + {name: "unobserved oauth preserves legacy routing", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth}, want: true, wantReason: "billing_unobserved"}, + {name: "weekly allowance is not treated as weekly subscription", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: weeklyAllowance}}, want: true, wantReason: "eligible"}, {name: "billing forbidden is rejected", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: forbiddenBilling}}, want: false, wantReason: "billing_forbidden"}, {name: "weekly billing forbidden is rejected after partial success", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: weeklyForbidden}}, want: false, wantReason: "billing_forbidden"}, {name: "monthly billing forbidden is rejected after partial success", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: monthlyForbidden}}, want: false, wantReason: "billing_forbidden"}, - {name: "malformed billing observation fails closed", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: make(chan int)}}, want: false, wantReason: "billing_unobserved"}, + {name: "malformed billing observation preserves legacy routing", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: make(chan int)}}, want: true, wantReason: "billing_unobserved"}, {name: "malformed override falls back to observations", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{GrokMediaEligibleExtraKey: "false", grokBillingExtraKey: weeklyAllowance}}, want: true, wantReason: "eligible"}, {name: "explicit disable wins", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{GrokMediaEligibleExtraKey: false}}, want: false, wantReason: "override_disabled"}, {name: "explicit enable wins over forbidden probe", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{GrokMediaEligibleExtraKey: true, grokBillingExtraKey: forbiddenBilling}}, want: true, wantReason: "override_enabled"}, @@ -81,24 +63,6 @@ func TestGrokMediaGenerationEligibility(t *testing.T) { } } -func TestGrokMediaCapabilityKeepsOnlyUnobservedOAuthAsProbeCandidate(t *testing.T) { - unobserved := &Account{Platform: PlatformGrok, Type: AccountTypeOAuth} - eligible, reason := unobserved.GrokMediaGenerationEligibility() - require.False(t, eligible) - require.Equal(t, "billing_unobserved", reason) - require.True(t, unobserved.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityGrokMediaGeneration)) - - inconclusive := &Account{ - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Extra: map[string]any{grokBillingExtraKey: &xai.BillingSummary{ - StatusCode: http.StatusOK, - Partial: true, - }}, - } - require.False(t, inconclusive.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityGrokMediaGeneration)) -} - func TestGrokMediaCapabilityFiltersOnlyGeneration(t *testing.T) { account := &Account{ ID: 1, diff --git a/backend/internal/service/account_scheduling_threshold_eval.go b/backend/internal/service/account_scheduling_threshold_eval.go new file mode 100644 index 0000000000..49ca80db35 --- /dev/null +++ b/backend/internal/service/account_scheduling_threshold_eval.go @@ -0,0 +1,464 @@ +package service + +import ( + "encoding/json" + "math" + "strconv" + "strings" + "time" +) + +// AccountSchedulingThresholdDecision captures the pure pause decision for one account. +type AccountSchedulingThresholdDecision struct { + ShouldPause bool + Platform string + Window string + Scope string + ThresholdPercent int + UsedPercent float64 + Until *time.Time +} + +type accountSchedulingThresholdCandidate struct { + window string + scope string + usedPercent float64 + until *time.Time +} + +const accountSchedulingThresholdCredentialKey = "account_scheduling_threshold" + +// EvaluateAccountSchedulingThreshold evaluates whether an account should be paused +// based on the current per-platform scheduling threshold snapshot. +func EvaluateAccountSchedulingThreshold(account *Account, thresholds map[string]int, now time.Time) AccountSchedulingThresholdDecision { + decision := AccountSchedulingThresholdDecision{} + if account == nil { + return decision + } + + decision.Platform = strings.ToLower(strings.TrimSpace(account.Platform)) + if decision.Platform == "" { + return decision + } + if !isAllowedSchedulingThresholdPlatform(decision.Platform) { + return decision + } + + threshold, ok := resolveEffectiveAccountSchedulingThreshold(account, thresholds, decision.Platform) + decision.ThresholdPercent = threshold + if !ok || threshold >= 100 { + return decision + } + + var winner *accountSchedulingThresholdCandidate + switch decision.Platform { + case PlatformOpenAI: + winner = pickLatestResetSchedulingCandidate(openAIThresholdCandidates(account), threshold, now) + case PlatformAnthropic: + winner = pickLatestResetSchedulingCandidate(anthropicThresholdCandidates(account), threshold, now) + case PlatformGrok: + winner = pickLatestResetSchedulingCandidate(grokThresholdCandidates(account), threshold, now) + default: + return decision + } + + if winner == nil { + return decision + } + + decision.ShouldPause = true + decision.Window = winner.window + decision.Scope = winner.scope + decision.UsedPercent = winner.usedPercent + decision.Until = winner.until + return decision +} + +func isAllowedSchedulingThresholdPlatform(platform string) bool { + for _, allowed := range AllowedSchedulingThresholdPlatforms { + if platform == allowed { + return true + } + } + return false +} + +func resolveEffectiveAccountSchedulingThreshold(account *Account, thresholds map[string]int, platform string) (int, bool) { + if account != nil { + if threshold, ok := accountSchedulingThresholdOverride(account); ok { + return threshold, true + } + } + return lookupAccountSchedulingThreshold(thresholds, platform) +} + +func accountSchedulingThresholdOverride(account *Account) (int, bool) { + if account == nil || len(account.Credentials) == 0 { + return 0, false + } + raw, ok := account.Credentials[accountSchedulingThresholdCredentialKey] + if !ok { + return 0, false + } + return parseAccountSchedulingThresholdValue(raw) +} + +func parseAccountSchedulingThresholdValue(raw any) (int, bool) { + var value int + switch v := raw.(type) { + case int: + value = v + case int64: + value = int(v) + case float64: + value = int(math.Round(v)) + case float32: + value = int(math.Round(float64(v))) + case json.Number: + parsed, err := v.Float64() + if err != nil { + return 0, false + } + value = int(math.Round(parsed)) + case string: + raw := strings.TrimSpace(v) + parsed, err := strconv.Atoi(raw) + if err == nil { + value = parsed + break + } + parsedFloat, floatErr := strconv.ParseFloat(raw, 64) + if floatErr != nil { + return 0, false + } + value = int(math.Round(parsedFloat)) + default: + return 0, false + } + if value < 1 || value > 100 { + return 0, false + } + return value, true +} + +func lookupAccountSchedulingThreshold(thresholds map[string]int, platform string) (int, bool) { + if len(thresholds) == 0 { + return 0, false + } + value, ok := thresholds[platform] + return value, ok +} + +func openAIThresholdCandidates(account *Account) []*accountSchedulingThresholdCandidate { + if account == nil { + return nil + } + if !openAICodexSnapshotIdentityTrusted(account) { + return nil + } + return []*accountSchedulingThresholdCandidate{ + openAIThresholdCandidate(account.Extra, "5h"), + openAIThresholdCandidate(account.Extra, "7d"), + } +} + +func openAICodexSnapshotIdentityTrusted(account *Account) bool { + if account == nil || !account.IsOpenAIOAuth() || len(account.Extra) == 0 { + return true + } + + if identityValuesConflict( + firstStringValue(account.Credentials, "email"), + firstStringValue(account.Extra, "email", "email_address"), + ) { + return false + } + if identityValuesConflict( + firstStringValue(account.Credentials, "chatgpt_account_id"), + firstStringValue(account.Extra, "chatgpt_account_id", "account_id"), + ) { + return false + } + if identityValuesConflict( + firstStringValue(account.Credentials, "workspace_id", "chatgpt_workspace_id", "organization_id", "org_id"), + firstStringValue(account.Extra, "workspace_id", "chatgpt_workspace_id", "organization_id", "org_id"), + ) { + return false + } + return true +} + +func identityValuesConflict(left, right string) bool { + left = strings.TrimSpace(left) + right = strings.TrimSpace(right) + return left != "" && right != "" && !strings.EqualFold(left, right) +} + +// firstStringValue returns the first non-empty string among the given map keys. +// Used by OpenAI codex snapshot identity matching for scheduling thresholds. +func firstStringValue(values map[string]any, keys ...string) string { + if len(values) == 0 { + return "" + } + for _, key := range keys { + raw, ok := values[key] + if !ok || raw == nil { + continue + } + switch typed := raw.(type) { + case string: + if v := strings.TrimSpace(typed); v != "" { + return v + } + default: + if v := strings.TrimSpace(stringValue(raw)); v != "" { + return v + } + } + } + return "" +} + +func openAIThresholdCandidate(extra map[string]any, window string) *accountSchedulingThresholdCandidate { + if len(extra) == 0 { + return nil + } + + var ( + usedPercentKey string + resetAtKey string + ) + switch window { + case "5h": + usedPercentKey = "codex_5h_used_percent" + resetAtKey = "codex_5h_reset_at" + case "7d": + usedPercentKey = "codex_7d_used_percent" + resetAtKey = "codex_7d_reset_at" + default: + return nil + } + + usedPercent, ok := extra[usedPercentKey] + if !ok { + return nil + } + return &accountSchedulingThresholdCandidate{ + window: window, + usedPercent: utilizationAsPercent(usedPercent), + until: parseSchedulingResetAt(extra[resetAtKey]), + } +} + +func anthropicThresholdCandidates(account *Account) []*accountSchedulingThresholdCandidate { + if account == nil { + return nil + } + + var candidates []*accountSchedulingThresholdCandidate + if usedPercent := utilizationAsPercent(account.Extra["session_window_utilization"]); usedPercent > 0 { + candidates = append(candidates, &accountSchedulingThresholdCandidate{ + window: "5h", + usedPercent: usedPercent, + until: cloneTimePtr(account.SessionWindowEnd), + }) + } + if usedPercent := utilizationAsPercent(account.Extra["passive_usage_7d_utilization"]); usedPercent > 0 { + candidates = append(candidates, &accountSchedulingThresholdCandidate{ + window: "7d", + usedPercent: usedPercent, + until: parseSchedulingResetAt(account.Extra["passive_usage_7d_reset"]), + }) + } + return candidates +} + +// NOTE: Gemini / Kiro / Antigravity are intentionally NOT threshold-pausing +// platforms (see AllowedSchedulingThresholdPlatforms and the evaluator switch, +// asserted by TestEvaluateAccountSchedulingThreshold_UnsupportedPlatformsDoNotPause). +// Their former per-platform candidate readers were dead code — never reachable +// from EvaluateAccountSchedulingThreshold — and have been removed to avoid the +// false impression that configuring a threshold for them has any effect. The +// kiro_sched_* / antigravity_sched_* extras are still written purely as +// observability snapshots. + +func grokThresholdCandidates(account *Account) []*accountSchedulingThresholdCandidate { + if account == nil { + return nil + } + return []*accountSchedulingThresholdCandidate{ + { + window: "quota", + scope: "grok", + usedPercent: schedulingPercentValue(account.Extra["grok_sched_utilization"]), + until: parseSchedulingResetAt(account.Extra["grok_sched_reset_at"]), + }, + } +} + +func pickLatestResetSchedulingCandidate(candidates []*accountSchedulingThresholdCandidate, threshold int, now time.Time) *accountSchedulingThresholdCandidate { + var winner *accountSchedulingThresholdCandidate + for _, candidate := range candidates { + if !candidateMatchesThreshold(candidate, threshold, now) { + continue + } + if winner == nil || candidate.until.After(*winner.until) { + winner = candidate + continue + } + if winner.until.Equal(*candidate.until) && candidate.usedPercent > winner.usedPercent { + winner = candidate + } + } + return winner +} + +func candidateMatchesThreshold(candidate *accountSchedulingThresholdCandidate, threshold int, now time.Time) bool { + if candidate == nil || candidate.until == nil || !candidate.until.After(now) { + return false + } + return candidate.usedPercent >= float64(threshold) +} + +func utilizationAsPercent(raw any) float64 { + switch v := raw.(type) { + case float64: + if v >= 0 && v <= 1 { + return v * 100 + } + return v + case float32: + value := float64(v) + if value >= 0 && value <= 1 { + return value * 100 + } + return value + case int: + return float64(v) + case int64: + return float64(v) + case json.Number: + value, err := v.Float64() + if err != nil { + return 0 + } + if strings.Contains(v.String(), ".") && value >= 0 && value <= 1 { + return value * 100 + } + return value + case string: + trimmed := strings.TrimSpace(v) + value, err := strconv.ParseFloat(trimmed, 64) + if err != nil { + return 0 + } + if strings.Contains(trimmed, ".") && value >= 0 && value <= 1 { + return value * 100 + } + return value + default: + return 0 + } +} + +func schedulingPercentValue(raw any) float64 { + switch v := raw.(type) { + case float64: + return v + case float32: + return float64(v) + case int: + return float64(v) + case int64: + return float64(v) + case json.Number: + value, err := v.Float64() + if err != nil { + return 0 + } + return value + case string: + value, err := strconv.ParseFloat(strings.TrimSpace(v), 64) + if err != nil { + return 0 + } + return value + default: + return 0 + } +} + +func parseSchedulingResetAt(raw any) *time.Time { + switch v := raw.(type) { + case nil: + return nil + case time.Time: + ts := v + return &ts + case *time.Time: + return cloneTimePtr(v) + case string: + trimmed := strings.TrimSpace(v) + if trimmed == "" { + return nil + } + ts, err := parseSchedulingTime(trimmed) + if err != nil { + return nil + } + return &ts + case json.Number: + if value, err := v.Int64(); err == nil && value > 0 { + ts := time.Unix(value, 0) + return &ts + } + if value, err := v.Float64(); err == nil && value > 0 { + ts := time.Unix(int64(value), 0) + return &ts + } + case float64: + if v > 0 { + ts := time.Unix(int64(v), 0) + return &ts + } + case float32: + if v > 0 { + ts := time.Unix(int64(v), 0) + return &ts + } + case int: + if v > 0 { + ts := time.Unix(int64(v), 0) + return &ts + } + case int64: + if v > 0 { + ts := time.Unix(v, 0) + return &ts + } + } + return nil +} + +func parseSchedulingTime(raw string) (time.Time, error) { + formats := []string{ + time.RFC3339, + time.RFC3339Nano, + "2006-01-02T15:04:05Z", + "2006-01-02T15:04:05.000Z", + } + for _, format := range formats { + if ts, err := time.Parse(format, raw); err == nil { + return ts, nil + } + } + return time.Time{}, strconv.ErrSyntax +} + +func cloneTimePtr(src *time.Time) *time.Time { + if src == nil { + return nil + } + value := *src + return &value +} diff --git a/backend/internal/service/account_scheduling_threshold_eval_test.go b/backend/internal/service/account_scheduling_threshold_eval_test.go new file mode 100644 index 0000000000..43546be923 --- /dev/null +++ b/backend/internal/service/account_scheduling_threshold_eval_test.go @@ -0,0 +1,308 @@ +//go:build unit + +package service + +import ( + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestEvaluateAccountSchedulingThreshold_OpenAIChoosesLatestResetWindow(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + wantUntil := now.Add(72 * time.Hour) + account := &Account{ + Platform: PlatformOpenAI, + Extra: map[string]any{ + "codex_5h_used_percent": 90.0, + "codex_5h_reset_at": now.Add(2 * time.Hour).Format(time.RFC3339), + "codex_7d_used_percent": 85.0, + "codex_7d_reset_at": wantUntil.Format(time.RFC3339), + }, + } + + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{ + PlatformOpenAI: 80, + }, now) + + require.True(t, decision.ShouldPause) + require.Equal(t, PlatformOpenAI, decision.Platform) + require.Equal(t, "7d", decision.Window) + require.Empty(t, decision.Scope) + require.Equal(t, 85.0, decision.UsedPercent) + require.NotNil(t, decision.Until) + require.True(t, wantUntil.Equal(*decision.Until)) +} + +func TestEvaluateAccountSchedulingThreshold_OpenAIIgnoresMismatchedCodexSnapshotIdentity(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 13, 8, 50, 0, 0, time.UTC) + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "email": "CageLeen9208@outlook.com", + "chatgpt_account_id": "1f945aa7-d9a9-4369-9542-0c702ff4adb0", + "workspace_id": "org-nU4goUxMmureroyswT5oYPv4", + "chatgpt_workspace_id": "org-nU4goUxMmureroyswT5oYPv4", + }, + Extra: map[string]any{ + "email": "MasonDobies01@outlook.com", + "name": "Paul Clark", + "workspace_id": "org-avRk1G4qdXg7qph3cRIraNKf", + "codex_7d_used_percent": 100.0, + "codex_7d_reset_at": now.Add(7 * 24 * time.Hour).Format(time.RFC3339), + }, + } + + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{ + PlatformOpenAI: 99, + }, now) + + require.False(t, decision.ShouldPause) +} + +func TestEvaluateAccountSchedulingThreshold_AnthropicIgnoresExpiredFiveHourWindow(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + expiredEnd := now.Add(-30 * time.Minute) + wantUntil := now.Add(5 * 24 * time.Hour) + account := &Account{ + Platform: PlatformAnthropic, + SessionWindowEnd: &expiredEnd, + Extra: map[string]any{ + "session_window_utilization": 0.99, + "passive_usage_7d_utilization": 0.82, + "passive_usage_7d_reset": float64(wantUntil.Unix()), + }, + } + + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{ + PlatformAnthropic: 80, + }, now) + + require.True(t, decision.ShouldPause) + require.Equal(t, PlatformAnthropic, decision.Platform) + require.Equal(t, "7d", decision.Window) + require.Empty(t, decision.Scope) + require.Equal(t, 82.0, decision.UsedPercent) + require.NotNil(t, decision.Until) + require.True(t, wantUntil.Equal(*decision.Until)) +} + +func TestEvaluateAccountSchedulingThreshold_FractionalPlatformsKeepFractionSemantics(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + openAIUntil := now.Add(24 * time.Hour) + openAIAccount := &Account{ + Platform: PlatformOpenAI, + Extra: map[string]any{ + "codex_5h_used_percent": 0.91, + "codex_5h_reset_at": openAIUntil.Format(time.RFC3339), + }, + } + + openAIDecision := EvaluateAccountSchedulingThreshold(openAIAccount, map[string]int{ + PlatformOpenAI: 90, + }, now) + + require.True(t, openAIDecision.ShouldPause) + require.Equal(t, 91.0, openAIDecision.UsedPercent) + + anthropicUntil := now.Add(5 * time.Hour) + anthropicAccount := &Account{ + Platform: PlatformAnthropic, + SessionWindowEnd: &anthropicUntil, + Extra: map[string]any{ + "session_window_utilization": 0.92, + }, + } + + anthropicDecision := EvaluateAccountSchedulingThreshold(anthropicAccount, map[string]int{ + PlatformAnthropic: 90, + }, now) + + require.True(t, anthropicDecision.ShouldPause) + require.Equal(t, 92.0, anthropicDecision.UsedPercent) +} + +func TestEvaluateAccountSchedulingThreshold_AccountOverrideCanLowerOpenAIThreshold(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + wantUntil := now.Add(12 * time.Hour) + account := &Account{ + Platform: PlatformOpenAI, + Credentials: map[string]any{ + "account_scheduling_threshold": 80, + }, + Extra: map[string]any{ + "codex_7d_used_percent": 85.0, + "codex_7d_reset_at": wantUntil.Format(time.RFC3339), + }, + } + + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{ + PlatformOpenAI: 90, + }, now) + + require.True(t, decision.ShouldPause) + require.Equal(t, PlatformOpenAI, decision.Platform) + require.Equal(t, 80, decision.ThresholdPercent) + require.Equal(t, "7d", decision.Window) + require.Empty(t, decision.Scope) + require.Equal(t, 85.0, decision.UsedPercent) + require.NotNil(t, decision.Until) + require.True(t, wantUntil.Equal(*decision.Until)) +} + +func TestEvaluateAccountSchedulingThreshold_AccountOverrideHundredDisablesOpenAI(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + account := &Account{ + Platform: PlatformOpenAI, + Credentials: map[string]any{ + "account_scheduling_threshold": 100, + }, + Extra: map[string]any{ + "codex_7d_used_percent": 99.0, + "codex_7d_reset_at": now.Add(24 * time.Hour).Format(time.RFC3339), + }, + } + + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{ + PlatformOpenAI: 80, + }, now) + + require.False(t, decision.ShouldPause) + require.Equal(t, 100, decision.ThresholdPercent) +} + +func TestEvaluateAccountSchedulingThreshold_AccountOverrideRoundsDecimalThreshold(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + wantUntil := now.Add(12 * time.Hour) + account := &Account{ + Platform: PlatformOpenAI, + Credentials: map[string]any{ + "account_scheduling_threshold": 75.5, + }, + Extra: map[string]any{ + "codex_7d_used_percent": 80.0, + "codex_7d_reset_at": wantUntil.Format(time.RFC3339), + }, + } + + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{ + PlatformOpenAI: 90, + }, now) + + require.True(t, decision.ShouldPause) + require.Equal(t, 76, decision.ThresholdPercent) + require.Equal(t, 80.0, decision.UsedPercent) + require.NotNil(t, decision.Until) + require.True(t, wantUntil.Equal(*decision.Until)) +} + +func TestEvaluateAccountSchedulingThreshold_UnsupportedPlatformsDoNotPause(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + cases := []struct { + name string + platform string + threshold int + extra map[string]any + }{ + { + name: "gemini", + platform: PlatformGemini, + threshold: 80, + extra: map[string]any{ + "gemini_usage_raw": map[string]any{ + "buckets": []any{ + map[string]any{ + "modelId": "gemini-2.5-pro", + "remainingFraction": 0.05, + "resetTime": now.Add(2 * time.Hour).Format(time.RFC3339), + }, + }, + }, + }, + }, + { + name: "kiro", + platform: PlatformKiro, + threshold: 90, + extra: map[string]any{ + "kiro_sched_utilization": 99.0, + "kiro_sched_reset_at": now.Add(24 * time.Hour).Format(time.RFC3339), + }, + }, + { + name: "antigravity", + platform: PlatformAntigravity, + threshold: 90, + extra: map[string]any{ + "antigravity_sched_utilization": 92.0, + "antigravity_sched_reset_at": now.Add(48 * time.Hour).Format(time.RFC3339), + "antigravity_sched_scope": "gemini", + }, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + account := &Account{ + Platform: tc.platform, + Credentials: map[string]any{ + "account_scheduling_threshold": 1, + }, + Extra: tc.extra, + } + + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{ + tc.platform: tc.threshold, + }, now) + + require.False(t, decision.ShouldPause) + require.Zero(t, decision.ThresholdPercent) + }) + } +} + +func TestEvaluateAccountSchedulingThreshold_GrokUsesConfiguredThresholds(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + wantUntil := now.Add(2 * time.Hour) + account := &Account{ + Platform: PlatformGrok, + Extra: map[string]any{ + "grok_sched_utilization": 92.0, + "grok_sched_reset_at": wantUntil.Format(time.RFC3339), + }, + } + + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{ + PlatformGrok: 90, + }, now) + + require.True(t, decision.ShouldPause) + require.Equal(t, PlatformGrok, decision.Platform) + require.Equal(t, 90, decision.ThresholdPercent) + require.Equal(t, "grok", decision.Scope) + require.Equal(t, 92.0, decision.UsedPercent) + require.NotNil(t, decision.Until) + require.True(t, wantUntil.Equal(*decision.Until)) +} diff --git a/backend/internal/service/account_scheduling_threshold_integration_test.go b/backend/internal/service/account_scheduling_threshold_integration_test.go new file mode 100644 index 0000000000..d2b80dbbcf --- /dev/null +++ b/backend/internal/service/account_scheduling_threshold_integration_test.go @@ -0,0 +1,136 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +type thresholdSelectionAccountRepoStub struct { + rateLimitAccountRepoStub + accounts []Account +} + +func (r *thresholdSelectionAccountRepoStub) ListSchedulableByPlatform(_ context.Context, platform string) ([]Account, error) { + filtered := make([]Account, 0, len(r.accounts)) + for _, account := range r.accounts { + if account.Platform == platform { + filtered = append(filtered, account) + } + } + return filtered, nil +} + +func (r *thresholdSelectionAccountRepoStub) ListSchedulableByGroupIDAndPlatform(ctx context.Context, _ int64, platform string) ([]Account, error) { + return r.ListSchedulableByPlatform(ctx, platform) +} + +func (r *thresholdSelectionAccountRepoStub) ListSchedulableUngroupedByPlatform(ctx context.Context, platform string) ([]Account, error) { + return r.ListSchedulableByPlatform(ctx, platform) +} + +func TestGatewayService_ListSchedulableAccounts_DoesNotFilterUnsupportedThresholdPlatforms(t *testing.T) { + accountSchedulingThresholdsSF.Forget(SettingKeyAccountSchedulingThresholds) + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{}) + + settingsRepo := newMockSettingRepo() + settingsRepo.data[SettingKeyAccountSchedulingThresholds] = `{"openai":90}` + + accountRepo := &thresholdSelectionAccountRepoStub{ + accounts: []Account{ + { + ID: 3101, + Platform: PlatformKiro, + Status: StatusActive, + Schedulable: true, + Credentials: map[string]any{ + "account_scheduling_threshold": 1, + }, + Extra: map[string]any{ + "kiro_sched_utilization": 95.0, + "kiro_sched_reset_at": time.Now().UTC().Add(2 * time.Hour).Format(time.RFC3339), + }, + }, + { + ID: 3102, + Platform: PlatformKiro, + Status: StatusActive, + Schedulable: true, + Extra: map[string]any{ + "kiro_sched_utilization": 42.0, + "kiro_sched_reset_at": time.Now().UTC().Add(2 * time.Hour).Format(time.RFC3339), + }, + }, + }, + } + + rateLimitService := NewRateLimitService(accountRepo, nil, &config.Config{}, nil, nil) + rateLimitService.SetSettingService(NewSettingService(settingsRepo, &config.Config{})) + svc := &GatewayService{ + accountRepo: accountRepo, + cfg: &config.Config{}, + rateLimitService: rateLimitService, + } + + accounts, useMixed, err := svc.listSchedulableAccounts(context.Background(), nil, PlatformKiro, false) + + require.NoError(t, err) + require.False(t, useMixed) + require.Len(t, accounts, 2) + require.Equal(t, int64(3101), accounts[0].ID) + require.Equal(t, int64(3102), accounts[1].ID) + require.Equal(t, 0, accountRepo.tempCalls) +} + +func TestOpenAIGatewayService_ListSchedulableAccounts_FiltersThresholdBlockedAccounts(t *testing.T) { + accountSchedulingThresholdsSF.Forget(SettingKeyAccountSchedulingThresholds) + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{}) + + settingsRepo := newMockSettingRepo() + settingsRepo.data[SettingKeyAccountSchedulingThresholds] = `{"openai":85}` + + accountRepo := &thresholdSelectionAccountRepoStub{ + accounts: []Account{ + { + ID: 4101, + Platform: PlatformOpenAI, + Status: StatusActive, + Schedulable: true, + Extra: map[string]any{ + "codex_7d_used_percent": 91.0, + "codex_7d_reset_at": time.Now().UTC().Add(12 * time.Hour).Format(time.RFC3339), + }, + }, + { + ID: 4102, + Platform: PlatformOpenAI, + Status: StatusActive, + Schedulable: true, + Extra: map[string]any{ + "codex_7d_used_percent": 40.0, + "codex_7d_reset_at": time.Now().UTC().Add(12 * time.Hour).Format(time.RFC3339), + }, + }, + }, + } + + rateLimitService := NewRateLimitService(accountRepo, nil, &config.Config{}, nil, nil) + rateLimitService.SetSettingService(NewSettingService(settingsRepo, &config.Config{})) + svc := &OpenAIGatewayService{ + accountRepo: accountRepo, + cfg: &config.Config{}, + rateLimitService: rateLimitService, + } + + accounts, err := svc.listSchedulableAccounts(context.Background(), nil, PlatformOpenAI) + + require.NoError(t, err) + require.Len(t, accounts, 1) + require.Equal(t, int64(4102), accounts[0].ID) + require.Equal(t, 1, accountRepo.tempCalls) +} diff --git a/backend/internal/service/account_scheduling_threshold_reason.go b/backend/internal/service/account_scheduling_threshold_reason.go new file mode 100644 index 0000000000..b43c7e22a1 --- /dev/null +++ b/backend/internal/service/account_scheduling_threshold_reason.go @@ -0,0 +1,188 @@ +package service + +import ( + "encoding/json" + "fmt" + "strings" + "time" +) + +const AccountSchedulingThresholdReasonSource = "account_scheduling_threshold" + +const ( + defaultTempUnschedReasonErrorMessage = "temporary scheduling block reason unavailable" + defaultAccountSchedulingThresholdErrorMessage = "account scheduling threshold reached" +) + +type tempUnschedReasonPayload struct { + Source string `json:"source,omitempty"` + Platform string `json:"platform,omitempty"` + Window string `json:"window,omitempty"` + Scope string `json:"scope,omitempty"` + ThresholdPercent int `json:"threshold_percent,omitempty"` + UsedPercent float64 `json:"used_percent,omitempty"` + UntilUnix int64 `json:"until_unix,omitempty"` + TriggeredAtUnix int64 `json:"triggered_at_unix,omitempty"` + ErrorMessage string `json:"error_message"` +} + +type AccountSchedulingThresholdReasonInput struct { + Platform string + Window string + Scope string + ThresholdPercent int + UsedPercent float64 + Until time.Time + Now time.Time +} + +func BuildTempUnschedReasonPayload(source string, errorMessage string) string { + payload := tempUnschedReasonPayload{ + Source: strings.TrimSpace(source), + ErrorMessage: normalizeTempUnschedReasonErrorMessage(errorMessage, defaultTempUnschedReasonErrorMessage), + } + + raw, err := json.Marshal(payload) + if err != nil { + return payload.ErrorMessage + } + return string(raw) +} + +func BuildAccountSchedulingThresholdReason(errorMessage string) string { + return BuildTempUnschedReasonPayload( + AccountSchedulingThresholdReasonSource, + normalizeTempUnschedReasonErrorMessage(errorMessage, defaultAccountSchedulingThresholdErrorMessage), + ) +} + +func BuildDetailedAccountSchedulingThresholdReason(input AccountSchedulingThresholdReasonInput) string { + triggeredAt := input.Now + if triggeredAt.IsZero() { + triggeredAt = time.Now().UTC() + } + payload := tempUnschedReasonPayload{ + Source: AccountSchedulingThresholdReasonSource, + Platform: strings.TrimSpace(input.Platform), + Window: strings.TrimSpace(input.Window), + Scope: strings.TrimSpace(input.Scope), + ThresholdPercent: input.ThresholdPercent, + UsedPercent: input.UsedPercent, + TriggeredAtUnix: triggeredAt.Unix(), + ErrorMessage: buildAccountSchedulingThresholdErrorMessage(input), + } + if !input.Until.IsZero() { + payload.UntilUnix = input.Until.UTC().Unix() + } + + raw, err := json.Marshal(payload) + if err != nil { + return payload.ErrorMessage + } + return string(raw) +} + +func IsAccountSchedulingThresholdReason(rawReason string) bool { + payload, ok := parseTempUnschedReasonPayload(rawReason) + if !ok { + return false + } + return payload.Source == AccountSchedulingThresholdReasonSource +} + +func parseTempUnschedReasonPayload(rawReason string) (tempUnschedReasonPayload, bool) { + rawReason = strings.TrimSpace(rawReason) + if rawReason == "" { + return tempUnschedReasonPayload{}, false + } + + var payload tempUnschedReasonPayload + if err := json.Unmarshal([]byte(rawReason), &payload); err != nil { + return tempUnschedReasonPayload{}, false + } + payload.Source = strings.TrimSpace(payload.Source) + payload.ErrorMessage = strings.TrimSpace(payload.ErrorMessage) + return payload, true +} + +func normalizeTempUnschedReasonErrorMessage(errorMessage string, fallback string) string { + errorMessage = strings.TrimSpace(errorMessage) + if errorMessage != "" { + return errorMessage + } + + fallback = strings.TrimSpace(fallback) + if fallback != "" { + return fallback + } + return defaultTempUnschedReasonErrorMessage +} + +func buildAccountSchedulingThresholdErrorMessage(input AccountSchedulingThresholdReasonInput) string { + platform := strings.TrimSpace(input.Platform) + if platform == "" { + platform = "account" + } + + target := strings.TrimSpace(input.Window) + if scope := strings.TrimSpace(input.Scope); scope != "" { + if target == "" { + target = scope + } else { + target = target + "/" + scope + } + } + if target == "" { + target = "usage window" + } + + threshold := input.ThresholdPercent + if threshold <= 0 { + threshold = 100 + } + + untilText := "the window reset" + if !input.Until.IsZero() { + untilText = input.Until.UTC().Format(time.RFC3339) + } + + return fmt.Sprintf( + "%s scheduling threshold reached for %s: %.1f%% used >= %d%%; paused until %s", + platform, + target, + input.UsedPercent, + threshold, + untilText, + ) +} + +func tempUnschedStateFromStoredReason(rawReason string, fallbackUntilUnix int64) *TempUnschedState { + state := &TempUnschedState{ + UntilUnix: fallbackUntilUnix, + RuleIndex: -1, + } + + rawReason = strings.TrimSpace(rawReason) + if rawReason == "" { + state.ErrorMessage = defaultTempUnschedReasonErrorMessage + return state + } + + parsed := TempUnschedState{RuleIndex: -1} + if err := json.Unmarshal([]byte(rawReason), &parsed); err == nil { + if fallbackUntilUnix > parsed.UntilUnix { + parsed.UntilUnix = fallbackUntilUnix + } + if strings.TrimSpace(parsed.ErrorMessage) == "" { + if IsAccountSchedulingThresholdReason(rawReason) { + parsed.ErrorMessage = defaultAccountSchedulingThresholdErrorMessage + } else { + parsed.ErrorMessage = rawReason + } + } + return &parsed + } + + state.ErrorMessage = rawReason + return state +} diff --git a/backend/internal/service/account_scheduling_threshold_reason_test.go b/backend/internal/service/account_scheduling_threshold_reason_test.go new file mode 100644 index 0000000000..34eea91cd4 --- /dev/null +++ b/backend/internal/service/account_scheduling_threshold_reason_test.go @@ -0,0 +1,76 @@ +//go:build unit + +package service + +import ( + "encoding/json" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestBuildAccountSchedulingThresholdReason_UsesSourceAndFallbackMessage(t *testing.T) { + raw := BuildAccountSchedulingThresholdReason(" \t ") + + var payload map[string]string + require.NoError(t, json.Unmarshal([]byte(raw), &payload)) + require.Equal(t, AccountSchedulingThresholdReasonSource, payload["source"]) + require.Equal(t, defaultAccountSchedulingThresholdErrorMessage, payload["error_message"]) + require.True(t, IsAccountSchedulingThresholdReason(raw)) +} + +func TestIsAccountSchedulingThresholdReason(t *testing.T) { + require.True(t, IsAccountSchedulingThresholdReason(BuildAccountSchedulingThresholdReason("threshold reached"))) + require.False(t, IsAccountSchedulingThresholdReason(BuildTempUnschedReasonPayload("", "temporary block"))) + require.False(t, IsAccountSchedulingThresholdReason("plain text reason")) +} + +func TestTempUnschedStateFromStoredReason_EmptyReasonUsesFallbackErrorMessage(t *testing.T) { + state := tempUnschedStateFromStoredReason(" \n ", 1735689600) + + require.NotNil(t, state) + require.Equal(t, int64(1735689600), state.UntilUnix) + require.Equal(t, defaultTempUnschedReasonErrorMessage, state.ErrorMessage) +} + +func TestTempUnschedStateFromStoredReason_MissingRuleIndexIsSystemRule(t *testing.T) { + state := tempUnschedStateFromStoredReason(`{"error_message":"system cooldown"}`, 123) + require.Equal(t, -1, state.RuleIndex) +} + +func TestTempUnschedStateFromStoredReason_SchedulingThresholdJSONWithoutMessageUsesThresholdFallback(t *testing.T) { + raw := `{"source":"` + AccountSchedulingThresholdReasonSource + `"}` + + state := tempUnschedStateFromStoredReason(raw, 1735689600) + + require.NotNil(t, state) + require.Equal(t, int64(1735689600), state.UntilUnix) + require.Equal(t, defaultAccountSchedulingThresholdErrorMessage, state.ErrorMessage) +} + +func TestBuildDetailedAccountSchedulingThresholdReason_IncludesReadableFields(t *testing.T) { + now := time.Unix(1735689600, 0).UTC() + until := now.Add(5 * time.Hour) + + raw := BuildDetailedAccountSchedulingThresholdReason(AccountSchedulingThresholdReasonInput{ + Platform: PlatformOpenAI, + Window: "7d", + ThresholdPercent: 90, + UsedPercent: 92.5, + Until: until, + Now: now, + }) + + var payload map[string]any + require.NoError(t, json.Unmarshal([]byte(raw), &payload)) + require.Equal(t, AccountSchedulingThresholdReasonSource, payload["source"]) + require.Equal(t, PlatformOpenAI, payload["platform"]) + require.Equal(t, "7d", payload["window"]) + require.Equal(t, float64(90), payload["threshold_percent"]) + require.Equal(t, float64(92.5), payload["used_percent"]) + require.Equal(t, float64(until.Unix()), payload["until_unix"]) + require.Equal(t, float64(now.Unix()), payload["triggered_at_unix"]) + require.Contains(t, payload["error_message"], "openai scheduling threshold reached") + require.Contains(t, payload["error_message"], "92.5% used >= 90%") +} diff --git a/backend/internal/service/account_scheduling_threshold_snapshot_cleanup.go b/backend/internal/service/account_scheduling_threshold_snapshot_cleanup.go new file mode 100644 index 0000000000..80c8fb2af1 --- /dev/null +++ b/backend/internal/service/account_scheduling_threshold_snapshot_cleanup.go @@ -0,0 +1,18 @@ +package service + +import ( + "context" + "fmt" +) + +type accountSchedulingThresholdSnapshotCleaner interface { + ClearAccountSchedulingThresholdSnapshots(ctx context.Context, id int64) error +} + +func clearAccountSchedulingThresholdSnapshots(ctx context.Context, repo AccountRepository, id int64) error { + cleaner, ok := repo.(accountSchedulingThresholdSnapshotCleaner) + if !ok { + return fmt.Errorf("account repository does not support account scheduling threshold snapshot cleanup") + } + return cleaner.ClearAccountSchedulingThresholdSnapshots(ctx, id) +} diff --git a/backend/internal/service/account_scheduling_threshold_snapshot_cleanup_test.go b/backend/internal/service/account_scheduling_threshold_snapshot_cleanup_test.go new file mode 100644 index 0000000000..59c7bdc4d2 --- /dev/null +++ b/backend/internal/service/account_scheduling_threshold_snapshot_cleanup_test.go @@ -0,0 +1,21 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" +) + +type missingThresholdSnapshotCleanerRepo struct { + AccountRepository +} + +func TestClearAccountSchedulingThresholdSnapshots_RequiresRepositorySupport(t *testing.T) { + err := clearAccountSchedulingThresholdSnapshots(context.Background(), missingThresholdSnapshotCleanerRepo{}, 1) + + require.Error(t, err) + require.Contains(t, err.Error(), "does not support account scheduling threshold snapshot cleanup") +} diff --git a/backend/internal/service/account_usage_service.go b/backend/internal/service/account_usage_service.go index 7033ad26c1..2bb77d5c9d 100644 --- a/backend/internal/service/account_usage_service.go +++ b/backend/internal/service/account_usage_service.go @@ -333,23 +333,28 @@ func NewAccountUsageService( } } -// GetUsage 获取账号使用量 -// OAuth账号: 调用Anthropic API获取真实数据(需要profile scope),API响应缓存10分钟,窗口统计缓存1分钟 -// Setup Token账号: 根据session_window推算5h窗口,7d数据不可用(没有profile scope) -// API Key账号: 不支持usage查询 -func (s *AccountUsageService) GetUsage(ctx context.Context, accountID int64, force ...bool) (*UsageInfo, error) { - forceProbe := len(force) > 0 && force[0] +func supportsAnthropicPassiveUsage(account *Account) bool { + return account != nil && account.IsAnthropicOAuthOrSetupToken() +} - account, err := s.accountRepo.GetByID(ctx, accountID) - if err != nil { - return nil, fmt.Errorf("get account failed: %w", err) +func batchUsageErrorMessage(err error) string { + if err == nil { + return "" } + return err.Error() +} + +func (s *AccountUsageService) getUsageForAccount(ctx context.Context, account *Account, forceProbe bool) (*UsageInfo, error) { + if account == nil { + return nil, fmt.Errorf("account is required") + } + accountID := account.ID // Dedicated UI load-test accounts must remain fully interactive without ever // contacting Anthropic with synthetic credentials. Reuse the same persisted // passive snapshot that the account table loads on mount. if account.IsSyntheticUITest() && account.IsAnthropicOAuthOrSetupToken() { - return s.GetPassiveUsage(ctx, accountID) + return s.getPassiveUsageForAccount(ctx, account) } if account.Platform == PlatformOpenAI && account.Type == AccountTypeOAuth { @@ -482,6 +487,96 @@ func (s *AccountUsageService) GetUsage(ctx context.Context, accountID int64, for return nil, fmt.Errorf("account type %s does not support usage query", account.Type) } +// GetUsage 获取账号使用量 +// OAuth账号: 调用Anthropic API获取真实数据(需要profile scope),API响应缓存10分钟,窗口统计缓存1分钟 +// Setup Token账号: 根据session_window推算5h窗口,7d数据不可用(没有profile scope) +// API Key账号: 不支持usage查询 +func (s *AccountUsageService) GetUsage(ctx context.Context, accountID int64, force ...bool) (*UsageInfo, error) { + forceProbe := len(force) > 0 && force[0] + + account, err := s.accountRepo.GetByID(ctx, accountID) + if err != nil { + return nil, fmt.Errorf("get account failed: %w", err) + } + + return s.getUsageForAccount(ctx, account, forceProbe) +} + +// GetUsageBatch 批量获取账号使用量。 +// Anthropic OAuth/SetupToken 统一走 passive 链路,其他账号复用现有主动查询逻辑。 +// 单个账号失败不会中断整批请求,错误会按账号返回。 +func (s *AccountUsageService) GetUsageBatch(ctx context.Context, accountIDs []int64, force bool) (map[int64]*UsageInfo, map[int64]string, error) { + uniqueIDs := make([]int64, 0, len(accountIDs)) + seen := make(map[int64]struct{}, len(accountIDs)) + for _, accountID := range accountIDs { + if accountID <= 0 { + continue + } + if _, ok := seen[accountID]; ok { + continue + } + seen[accountID] = struct{}{} + uniqueIDs = append(uniqueIDs, accountID) + } + + usageByAccount := make(map[int64]*UsageInfo, len(uniqueIDs)) + errorsByAccount := make(map[int64]string) + if len(uniqueIDs) == 0 { + return usageByAccount, errorsByAccount, nil + } + + accounts, err := s.accountRepo.GetByIDs(ctx, uniqueIDs) + if err != nil { + return nil, nil, fmt.Errorf("get accounts failed: %w", err) + } + + accountsByID := make(map[int64]*Account, len(accounts)) + for _, account := range accounts { + if account == nil { + continue + } + accountsByID[account.ID] = account + } + + var mu sync.Mutex + g, gctx := errgroup.WithContext(ctx) + g.SetLimit(6) + + for _, accountID := range uniqueIDs { + id := accountID + account := accountsByID[id] + if account == nil { + errorsByAccount[id] = ErrAccountNotFound.Error() + continue + } + + g.Go(func() error { + var usage *UsageInfo + var usageErr error + if supportsAnthropicPassiveUsage(account) { + usage, usageErr = s.getPassiveUsageForAccount(gctx, account) + } else { + usage, usageErr = s.getUsageForAccount(gctx, account, force) + } + + mu.Lock() + defer mu.Unlock() + if usageErr != nil { + errorsByAccount[id] = batchUsageErrorMessage(usageErr) + return nil + } + usageByAccount[id] = usage + return nil + }) + } + + if err := g.Wait(); err != nil { + return nil, nil, err + } + + return usageByAccount, errorsByAccount, nil +} + // GetPassiveUsage 从 Account.Extra 中的被动采样数据构建 UsageInfo,不调用外部 API。 // 仅适用于 Anthropic OAuth / SetupToken 账号。 func (s *AccountUsageService) GetPassiveUsage(ctx context.Context, accountID int64) (*UsageInfo, error) { @@ -490,7 +585,11 @@ func (s *AccountUsageService) GetPassiveUsage(ctx context.Context, accountID int return nil, fmt.Errorf("get account failed: %w", err) } - if !account.IsAnthropicOAuthOrSetupToken() { + return s.getPassiveUsageForAccount(ctx, account) +} + +func (s *AccountUsageService) getPassiveUsageForAccount(ctx context.Context, account *Account) (*UsageInfo, error) { + if !supportsAnthropicPassiveUsage(account) { return nil, fmt.Errorf("passive usage only supported for Anthropic OAuth/SetupToken accounts") } diff --git a/backend/internal/service/account_usage_service_batch_test.go b/backend/internal/service/account_usage_service_batch_test.go new file mode 100644 index 0000000000..067bf0ddd1 --- /dev/null +++ b/backend/internal/service/account_usage_service_batch_test.go @@ -0,0 +1,191 @@ +package service + +import ( + "context" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats" +) + +// Minimal UsageLogRepository stub for batch usage tests (HEAD lacks geminiUsageLogRepoStub). +type usageBatchLogRepoStub struct{} + +var _ UsageLogRepository = (*usageBatchLogRepoStub)(nil) + +func (r *usageBatchLogRepoStub) Create(context.Context, *UsageLog) (bool, error) { + return false, nil +} +func (r *usageBatchLogRepoStub) GetByID(context.Context, int64) (*UsageLog, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) Delete(context.Context, int64) error { return nil } +func (r *usageBatchLogRepoStub) ListByUser(context.Context, int64, pagination.PaginationParams) ([]UsageLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} +func (r *usageBatchLogRepoStub) ListByAPIKey(context.Context, int64, pagination.PaginationParams) ([]UsageLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} +func (r *usageBatchLogRepoStub) ListByAccount(context.Context, int64, pagination.PaginationParams) ([]UsageLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} +func (r *usageBatchLogRepoStub) ListByUserAndTimeRange(context.Context, int64, time.Time, time.Time) ([]UsageLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} +func (r *usageBatchLogRepoStub) ListByAPIKeyAndTimeRange(context.Context, int64, time.Time, time.Time) ([]UsageLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} +func (r *usageBatchLogRepoStub) ListByAccountAndTimeRange(context.Context, int64, time.Time, time.Time) ([]UsageLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} +func (r *usageBatchLogRepoStub) ListByModelAndTimeRange(context.Context, string, time.Time, time.Time) ([]UsageLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} +func (r *usageBatchLogRepoStub) GetAccountWindowStats(context.Context, int64, time.Time) (*usagestats.AccountStats, error) { + return &usagestats.AccountStats{}, nil +} +func (r *usageBatchLogRepoStub) GetAccountTodayStats(context.Context, int64) (*usagestats.AccountStats, error) { + return &usagestats.AccountStats{}, nil +} +func (r *usageBatchLogRepoStub) GetDashboardStats(context.Context) (*usagestats.DashboardStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUsageTrendWithFilters(context.Context, time.Time, time.Time, string, int64, int64, int64, int64, string, *int16, *bool, *int8) ([]usagestats.TrendDataPoint, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetModelStatsWithFilters(context.Context, time.Time, time.Time, int64, int64, int64, int64, *int16, *bool, *int8) ([]usagestats.ModelStat, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetEndpointStatsWithFilters(context.Context, time.Time, time.Time, int64, int64, int64, int64, string, *int16, *bool, *int8) ([]usagestats.EndpointStat, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUpstreamEndpointStatsWithFilters(context.Context, time.Time, time.Time, int64, int64, int64, int64, string, *int16, *bool, *int8) ([]usagestats.EndpointStat, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetGroupStatsWithFilters(context.Context, time.Time, time.Time, int64, int64, int64, int64, *int16, *bool, *int8) ([]usagestats.GroupStat, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUserBreakdownStats(context.Context, time.Time, time.Time, usagestats.UserBreakdownDimension, int) ([]usagestats.UserBreakdownItem, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetAllGroupUsageSummary(context.Context, time.Time) ([]usagestats.GroupUsageSummary, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetAPIKeyUsageTrend(context.Context, time.Time, time.Time, string, int) ([]usagestats.APIKeyUsageTrendPoint, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUserUsageTrend(context.Context, time.Time, time.Time, string, int) ([]usagestats.UserUsageTrendPoint, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUserSpendingRanking(context.Context, time.Time, time.Time, int) (*usagestats.UserSpendingRankingResponse, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetBatchUserUsageStats(context.Context, []int64, time.Time, time.Time) (map[int64]*usagestats.BatchUserUsageStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetBatchAPIKeyUsageStats(context.Context, []int64, time.Time, time.Time) (map[int64]*usagestats.BatchAPIKeyUsageStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUserDashboardStats(context.Context, int64) (*usagestats.UserDashboardStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetAPIKeyDashboardStats(context.Context, int64) (*usagestats.UserDashboardStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUserUsageTrendByUserID(context.Context, int64, time.Time, time.Time, string) ([]usagestats.TrendDataPoint, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUserModelStats(context.Context, int64, time.Time, time.Time) ([]usagestats.ModelStat, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) ListWithFilters(context.Context, pagination.PaginationParams, usagestats.UsageLogFilters) ([]UsageLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} +func (r *usageBatchLogRepoStub) GetGlobalStats(context.Context, time.Time, time.Time) (*usagestats.UsageStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetStatsWithFilters(context.Context, usagestats.UsageLogFilters) (*usagestats.UsageStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetAccountUsageStats(context.Context, int64, time.Time, time.Time) (*usagestats.AccountUsageStatsResponse, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUserStatsAggregated(context.Context, int64, time.Time, time.Time) (*usagestats.UsageStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetAPIKeyStatsAggregated(context.Context, int64, time.Time, time.Time) (*usagestats.UsageStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetAccountStatsAggregated(context.Context, int64, time.Time, time.Time) (*usagestats.UsageStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetModelStatsAggregated(context.Context, string, time.Time, time.Time) (*usagestats.UsageStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetDailyStatsAggregated(context.Context, int64, time.Time, time.Time) ([]map[string]any, error) { + return nil, nil +} + +func TestAccountUsageService_GetUsageBatch_BestEffortByAccount(t *testing.T) { + t.Parallel() + + resetAt := time.Now().Add(2 * time.Hour).UTC().Truncate(time.Second) + + repo := &stubOpenAIAccountRepo{ + accounts: []Account{ + { + ID: 7001, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Extra: map[string]any{ + "passive_usage_7d_utilization": 0.62, + }, + }, + { + ID: 7002, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{ + "codex_usage_updated_at": time.Now().UTC().Format(time.RFC3339), + "codex_5h_used_percent": 18.0, + "codex_5h_reset_at": resetAt.Format(time.RFC3339), + "codex_7d_used_percent": 34.0, + "codex_7d_reset_at": resetAt.Add(24 * time.Hour).Format(time.RFC3339), + "workspace_id": "org-test", + "chatgpt_account_id": "acct-test", + "openai_snapshot_version": "test", + }, + }, + { + ID: 7003, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + }, + }, + } + + svc := &AccountUsageService{ + accountRepo: repo, + usageLogRepo: &usageBatchLogRepoStub{}, + cache: NewUsageCache(), + } + + usageByAccount, errorsByAccount, err := svc.GetUsageBatch(context.Background(), []int64{7001, 7002, 7003, 7002}, false) + if err != nil { + t.Fatalf("GetUsageBatch() error = %v", err) + } + + if usageByAccount[7001] == nil || usageByAccount[7001].Source != "passive" { + t.Fatalf("expected anthropic passive usage, got %#v", usageByAccount[7001]) + } + + if usageByAccount[7002] == nil || usageByAccount[7002].FiveHour == nil || usageByAccount[7002].FiveHour.Utilization != 18.0 { + t.Fatalf("expected openai snapshot usage, got %#v", usageByAccount[7002]) + } + + if !strings.Contains(strings.ToLower(errorsByAccount[7003]), "does not support usage query") { + t.Fatalf("expected API key account error to be preserved, got %q", errorsByAccount[7003]) + } +} diff --git a/backend/internal/service/admin_account.go b/backend/internal/service/admin_account.go index 044a7ae87b..c32c36af8c 100644 --- a/backend/internal/service/admin_account.go +++ b/backend/internal/service/admin_account.go @@ -394,63 +394,7 @@ func normalizeOpenAILongContextBillingUpdateExtra(account *Account, input *Updat return normalized, nil } -// ValidateGrokMediaEligibilityExtra validates the optional media-routing -// override. null removes the override and returns the account to automatic -// provider-observation based routing. -func ValidateGrokMediaEligibilityExtra(platform string, extra map[string]any) error { - if platform != PlatformGrok || extra == nil { - return nil - } - raw, exists := extra[GrokMediaEligibleExtraKey] - if !exists || raw == nil { - return nil - } - if _, ok := raw.(bool); !ok { - return infraerrors.BadRequest( - "GROK_MEDIA_ELIGIBILITY_INVALID", - "grok_media_eligible must be a boolean or null", - ) - } - return nil -} - -func normalizeGrokMediaEligibilityExtra(platform string, extra map[string]any) (map[string]any, error) { - if platform != PlatformGrok { - return extra, nil - } - if err := ValidateGrokMediaEligibilityExtra(platform, extra); err != nil { - return nil, err - } - normalized := maps.Clone(extra) - if normalized != nil && normalized[GrokMediaEligibleExtraKey] == nil { - delete(normalized, GrokMediaEligibleExtraKey) - } - return normalized, nil -} - -func normalizeGrokMediaEligibilityUpdateExtra(account *Account, input *UpdateAccountInput, normalized map[string]any) (map[string]any, error) { - if account == nil || account.Platform != PlatformGrok { - return normalized, nil - } - if err := ValidateGrokMediaEligibilityExtra(account.Platform, input.Extra); err != nil { - return nil, err - } - normalized = maps.Clone(normalized) - if normalized == nil { - normalized = make(map[string]any) - } - raw, provided := input.Extra[GrokMediaEligibleExtraKey] - if provided { - if raw == nil { - delete(normalized, GrokMediaEligibleExtraKey) - } - return normalized, nil - } - if current, ok := account.Extra[GrokMediaEligibleExtraKey].(bool); ok { - normalized[GrokMediaEligibleExtraKey] = current - } - return normalized, nil -} +// Grok media eligibility helpers live in account_grok_media_eligibility.go. func buildAccountForCreate(input *CreateAccountInput, accountExtra map[string]any) (*Account, error) { // Probe/session state is system-managed. New accounts always start with automatic refresh disabled. diff --git a/backend/internal/service/domain_constants.go b/backend/internal/service/domain_constants.go index f24244ef7c..333f3ffe57 100644 --- a/backend/internal/service/domain_constants.go +++ b/backend/internal/service/domain_constants.go @@ -44,6 +44,9 @@ const ( PlatformAntigravity = domain.PlatformAntigravity PlatformGrok = domain.PlatformGrok PlatformComposite = domain.PlatformComposite + // PlatformKiro is retained for unsupported-platform threshold tests and legacy + // account rows. Scheduling-threshold evaluation never pauses kiro accounts. + PlatformKiro = "kiro" ) // AllowedQuotaPlatforms 是允许设置 user × platform quota 的平台列表(单一权威来源)。 @@ -57,6 +60,14 @@ var AllowedQuotaPlatforms = []string{ PlatformGrok, } +// AllowedSchedulingThresholdPlatforms 是允许设置账号自动停调阈值的平台列表。 +// 仅 openai / anthropic / grok 有原生用量窗口可供评估;其他平台写入阈值无效果。 +var AllowedSchedulingThresholdPlatforms = []string{ + PlatformOpenAI, + PlatformAnthropic, + PlatformGrok, +} + // IsAllowedQuotaPlatform 报告 s 是否为合法的 quota platform 标识。 func IsAllowedQuotaPlatform(s string) bool { for _, p := range AllowedQuotaPlatforms { @@ -585,6 +596,10 @@ const ( // 值为 map[platform]{daily,weekly,monthly},null/缺省 = 不限制;0 = 禁用;>0 = USD 上限。 const SettingKeyDefaultPlatformQuotas = "default_platform_quotas" +// SettingKeyAccountSchedulingThresholds —— 系统全局:按平台自动停调阈值(JSON map)。 +// 值为 map[platform]percent,1..100;100 = 禁用该平台自动停调。 +const SettingKeyAccountSchedulingThresholds = "account_scheduling_thresholds" + // SettingKeyAuthSourcePlatformQuotas 返回某 auth source 的 platform quota JSON key。 // 形如 auth_source_default_{source}_platform_quotas func SettingKeyAuthSourcePlatformQuotas(source string) string { diff --git a/backend/internal/service/gateway_scheduling.go b/backend/internal/service/gateway_scheduling.go index 9660abdd6f..3c5f1cb96e 100644 --- a/backend/internal/service/gateway_scheduling.go +++ b/backend/internal/service/gateway_scheduling.go @@ -961,6 +961,7 @@ func (s *GatewayService) listSchedulableAccounts(ctx context.Context, groupID *i if s.schedulerSnapshot != nil { accounts, useMixed, err := s.schedulerSnapshot.ListSchedulableAccounts(ctx, groupID, platform, hasForcePlatform) if err == nil { + accounts = s.filterAccountsBySchedulingThreshold(ctx, accounts) slog.Debug("account_scheduling_list_snapshot", "group_id", derefGroupID(groupID), "platform", platform, @@ -1022,7 +1023,7 @@ func (s *GatewayService) listSchedulableAccounts(ctx context.Context, groupID *i "tls_fingerprint", acc.IsTLSFingerprintEnabled()) } } - return filtered, useMixed, nil + return s.filterAccountsBySchedulingThreshold(ctx, filtered), useMixed, nil } var accounts []Account @@ -1057,7 +1058,7 @@ func (s *GatewayService) listSchedulableAccounts(ctx context.Context, groupID *i "tls_fingerprint", acc.IsTLSFingerprintEnabled()) } } - return accounts, useMixed, nil + return s.filterAccountsBySchedulingThreshold(ctx, accounts), useMixed, nil } // IsSingleAntigravityAccountGroup 检查指定分组是否只有一个 antigravity 平台的可调度账号。 @@ -1428,10 +1429,44 @@ func (s *GatewayService) checkAndRegisterSession(ctx context.Context, account *A } func (s *GatewayService) getSchedulableAccount(ctx context.Context, accountID int64) (*Account, error) { + var ( + account *Account + err error + ) if s.schedulerSnapshot != nil { - return s.schedulerSnapshot.GetAccount(ctx, accountID) + account, err = s.schedulerSnapshot.GetAccount(ctx, accountID) + } else { + account, err = s.accountRepo.GetByID(ctx, accountID) } - return s.accountRepo.GetByID(ctx, accountID) + if err != nil || account == nil { + return account, err + } + if s.isAccountBlockedBySchedulingThreshold(ctx, account) { + return nil, nil + } + return account, nil +} + +func (s *GatewayService) filterAccountsBySchedulingThreshold(ctx context.Context, accounts []Account) []Account { + if len(accounts) == 0 { + return accounts + } + + filtered := make([]Account, 0, len(accounts)) + for i := range accounts { + if s.isAccountBlockedBySchedulingThreshold(ctx, &accounts[i]) { + continue + } + filtered = append(filtered, accounts[i]) + } + return filtered +} + +func (s *GatewayService) isAccountBlockedBySchedulingThreshold(ctx context.Context, account *Account) bool { + if s == nil || s.rateLimitService == nil || account == nil { + return false + } + return s.rateLimitService.ApplyAccountSchedulingThreshold(ctx, account) } func (s *GatewayService) hydrateSelectedAccount(ctx context.Context, account *Account) (*Account, error) { diff --git a/backend/internal/service/grok_model_quota_block.go b/backend/internal/service/grok_model_quota_block.go index 05d894fc73..0bc9cd40b8 100644 --- a/backend/internal/service/grok_model_quota_block.go +++ b/backend/internal/service/grok_model_quota_block.go @@ -87,7 +87,8 @@ func filterGrokModelQuotaBlockedAccounts(accounts []Account, model string, now t } out := make([]Account, 0, len(accounts)) for i := range accounts { - if isGrokModelQuotaBlocked(accounts[i].ID, model, now) { + upstreamModel := canonicalOpenAIAccountSchedulingModel(&accounts[i], model) + if isGrokModelQuotaBlocked(accounts[i].ID, upstreamModel, now) { continue } out = append(out, accounts[i]) diff --git a/backend/internal/service/grok_p2_test.go b/backend/internal/service/grok_p2_test.go index 7cd0130d50..d3d6e67cf6 100644 --- a/backend/internal/service/grok_p2_test.go +++ b/backend/internal/service/grok_p2_test.go @@ -26,6 +26,21 @@ func TestGrokModelQuotaBlock_FiltersOnlyNamedModel(t *testing.T) { require.Equal(t, id+1, filtered[0].ID) } +func TestGrokModelQuotaBlockFiltersMappedUpstreamModel(t *testing.T) { + id := time.Now().UnixNano()%1_000_000 + 7000 + markGrokModelQuotaBlock(id, "grok-4.5", time.Now().Add(time.Hour)) + account := Account{ + ID: id, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "model_mapping": map[string]any{"gpt-*": "grok-4.5"}, + }, + } + + require.Empty(t, filterGrokModelQuotaBlockedAccounts([]Account{account}, "gpt-5", time.Now())) +} + func TestIsGrokModelSpecificFreeUsage(t *testing.T) { require.True(t, isGrokModelSpecificFreeUsage( "you've used all the included free usage for model grok-4.5", "grok-4.5")) diff --git a/backend/internal/service/grok_team_rate_limit.go b/backend/internal/service/grok_team_rate_limit.go index c4d3bae425..f10852615a 100644 --- a/backend/internal/service/grok_team_rate_limit.go +++ b/backend/internal/service/grok_team_rate_limit.go @@ -121,7 +121,8 @@ func filterGrokTeamModelRateLimitedAccounts(accounts []Account, model string, no out := accounts[:0] kept := false for i := range accounts { - if isGrokTeamModelRateLimited(&accounts[i], model, now) { + upstreamModel := canonicalOpenAIAccountSchedulingModel(&accounts[i], model) + if isGrokTeamModelRateLimited(&accounts[i], upstreamModel, now) { continue } out = append(out, accounts[i]) diff --git a/backend/internal/service/grok_team_rate_limit_test.go b/backend/internal/service/grok_team_rate_limit_test.go index 50c4ed5065..4926d4d7aa 100644 --- a/backend/internal/service/grok_team_rate_limit_test.go +++ b/backend/internal/service/grok_team_rate_limit_test.go @@ -57,3 +57,19 @@ func TestGrokTeamModelRateLimit_Expires(t *testing.T) { // After mark with past, resolveGrokTeamRateLimitUntil path isn't used; mark uses now+default when until not after now. require.True(t, isGrokTeamModelRateLimited(a, "grok-4.5", time.Now())) } + +func TestGrokTeamModelRateLimitFilterUsesMappedUpstreamModel(t *testing.T) { + now := time.Now() + account := &Account{ + ID: 301, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "team_id": "team-mapped-301", + "model_mapping": map[string]any{"gpt-*": "grok-4.5"}, + }, + } + markGrokTeamModelRateLimit(account, "grok-4.5", now.Add(time.Hour)) + + require.Empty(t, filterGrokTeamModelRateLimitedAccounts([]Account{*account}, "gpt-5", now)) +} diff --git a/backend/internal/service/grok_upstream_headers.go b/backend/internal/service/grok_upstream_headers.go new file mode 100644 index 0000000000..ad0ba21a2c --- /dev/null +++ b/backend/internal/service/grok_upstream_headers.go @@ -0,0 +1,78 @@ +package service + +import ( + "net/http" + "strings" + + "github.com/gin-gonic/gin" + + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" +) + +// grokUpstreamUserAgent is kept for compatibility with older Grok request +// tests. Current requests use the pinned default UA from this package. +const grokUpstreamUserAgent = "sub2api-grok/1.0" + +// defaultBrowserLikeUpstreamUserAgent remains available for non-Grok OpenAI-like +// fingerprint templates that historically shared this constant. +const defaultBrowserLikeUpstreamUserAgent = "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36" + +// Fixed CLI identity aliases — single source of truth is internal/pkg/xai. +const ( + grokClientVersionHeader = xai.CLIStableVersion + grokClientIdentifierHeader = xai.CLIClientIdentifier + grokClientModeHeader = xai.CLIClientMode +) + +// defaultGrokUpstreamUserAgent is the pinned Grok CLI / workspace UA. +// Grok upstream must not forward Claude Code / Codex / browser client UAs. +func defaultGrokUpstreamUserAgent() string { + return xai.CLIUserAgent(xai.ResolveCLIVersion()) +} + +func applyDefaultGrokUpstreamHeaders(req *http.Request) { + if req == nil { + return + } + // Always stamp CLI identity. Do not preserve inbound client UA (Claude Code, + // Codex, curl, etc.) — xAI chat/CLI surfaces fingerprint the client string. + req.Header.Set("User-Agent", defaultGrokUpstreamUserAgent()) + req.Header.Set("x-grok-client-version", xai.ResolveCLIVersion()) + req.Header.Set("x-grok-client-identifier", grokClientIdentifierHeader) +} + +func applyGrokTLSProfileHeaders(req *http.Request, profile *tlsfingerprint.Profile) { + // HEAD Profile is TLS-only (no HTTP UserAgent/Originator fields). Always stamp CLI identity. + applyDefaultGrokUpstreamHeaders(req) + _ = profile +} + +// openAITLSFingerprintRuntime is the resolved TLS fingerprint routing result +// used by OpenAI/Grok outbound header application. Defined here so Grok header +// helpers compile even when the full OpenAI TLS router is not present on HEAD. +type openAITLSFingerprintRuntime struct { + Profile *tlsfingerprint.Profile + UpstreamUserAgent string + UpstreamOriginator string + Matched bool +} + +func applyGrokRuntimeHeaders(req *http.Request, runtime openAITLSFingerprintRuntime) { + applyDefaultGrokUpstreamHeaders(req) + if req == nil { + return + } + // Apply Originator only; force CLI UA after so router overrides cannot + // leak Codex/Claude Code identity onto Grok upstream. + if originator := strings.TrimSpace(runtime.UpstreamOriginator); originator != "" { + req.Header.Set("Originator", originator) + } + req.Header.Set("User-Agent", defaultGrokUpstreamUserAgent()) +} + +// resolveGrokUpstreamUserAgent always returns the pinned Grok CLI User-Agent. +// Inbound client UAs (Claude Code, Codex, browsers, libraries) are never forwarded. +func resolveGrokUpstreamUserAgent(_ *gin.Context) string { + return defaultGrokUpstreamUserAgent() +} diff --git a/backend/internal/service/grok_upstream_headers_test.go b/backend/internal/service/grok_upstream_headers_test.go new file mode 100644 index 0000000000..18972f4801 --- /dev/null +++ b/backend/internal/service/grok_upstream_headers_test.go @@ -0,0 +1,86 @@ +package service + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" +) + +func TestApplyDefaultGrokUpstreamHeadersUsesCLIUserAgent(t *testing.T) { + t.Setenv(xai.CLIVersionEnv, "") + + req, err := http.NewRequest(http.MethodGet, "https://api.x.ai/v1/responses", nil) + require.NoError(t, err) + req.Header.Set("User-Agent", "claude-code/1.2.3") + req.Header.Set("x-grok-client-version", "none") + + applyDefaultGrokUpstreamHeaders(req) + + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), req.Header.Get("User-Agent")) + require.Equal(t, xai.CLIClientVersion, req.Header.Get("x-grok-client-version")) + require.Equal(t, xai.CLIClientIdentifier, req.Header.Get("x-grok-client-identifier")) +} + +func TestApplyDefaultGrokUpstreamHeadersHonorsCLIVersionOverride(t *testing.T) { + t.Setenv(xai.CLIVersionEnv, "0.2.95") + + req, err := http.NewRequest(http.MethodGet, "https://api.x.ai/v1/responses", nil) + require.NoError(t, err) + req.Header.Set("User-Agent", "codex_cli_rs/0.144.0") + + applyDefaultGrokUpstreamHeaders(req) + + require.Equal(t, "0.2.95", req.Header.Get("x-grok-client-version")) + require.Equal(t, xai.CLIUserAgent("0.2.95"), req.Header.Get("User-Agent")) + require.Equal(t, "grok-shell", req.Header.Get("x-grok-client-identifier")) +} + +func TestResolveGrokUpstreamUserAgentNeverPassthrough(t *testing.T) { + t.Setenv(xai.CLIVersionEnv, "") + gin.SetMode(gin.TestMode) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + c.Request.Header.Set("User-Agent", "claude-cli/2.0.0 (Mac OS; arm64)") + + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), resolveGrokUpstreamUserAgent(c)) + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), resolveGrokUpstreamUserAgent(nil)) +} + +func TestApplyGrokRuntimeHeadersKeepsCLIUserAgent(t *testing.T) { + t.Setenv(xai.CLIVersionEnv, "") + + req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) + require.NoError(t, err) + req.Header.Set("User-Agent", "claude-code/9.9.9") + + applyGrokRuntimeHeaders(req, openAITLSFingerprintRuntime{ + UpstreamUserAgent: "codex_cli_rs/0.144.0", + UpstreamOriginator: "codex_cli_rs", + }) + + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), req.Header.Get("User-Agent")) + require.Equal(t, "codex_cli_rs", req.Header.Get("Originator")) + require.Equal(t, xai.CLIClientVersion, req.Header.Get("x-grok-client-version")) +} + +func TestApplyGrokTLSProfileHeadersAlwaysUsesCLIUserAgent(t *testing.T) { + t.Setenv(xai.CLIVersionEnv, "") + + req, err := http.NewRequest(http.MethodPost, "https://api.x.ai/v1/responses", nil) + require.NoError(t, err) + req.Header.Set("User-Agent", "grok-native/1.0") + + // HEAD Profile is TLS-only; Originator/UserAgent HTTP fields are not present. + applyGrokTLSProfileHeaders(req, &tlsfingerprint.Profile{Name: "chrome"}) + + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), req.Header.Get("User-Agent")) + require.Equal(t, xai.CLIClientVersion, req.Header.Get("x-grok-client-version")) +} diff --git a/backend/internal/service/openai_account_scheduler.go b/backend/internal/service/openai_account_scheduler.go index d4627a459d..6457ef8fc2 100644 --- a/backend/internal/service/openai_account_scheduler.go +++ b/backend/internal/service/openai_account_scheduler.go @@ -508,11 +508,12 @@ func (s *defaultOpenAIAccountScheduler) selectBySessionHash( } // Team+model cool: sticky must not pin a sibling under the same team 429 window. now := time.Now() - if account != nil && isGrokTeamModelRateLimited(account, req.RequestedModel, now) { + upstreamModel := canonicalOpenAIAccountSchedulingModel(account, req.RequestedModel) + if account != nil && isGrokTeamModelRateLimited(account, upstreamModel, now) { _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) return nil, false, nil } - if account != nil && isGrokModelQuotaBlocked(account.ID, req.RequestedModel, now) { + if account != nil && isGrokModelQuotaBlocked(account.ID, upstreamModel, now) { _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) return nil, false, nil } @@ -1264,6 +1265,12 @@ func (s *defaultOpenAIAccountScheduler) tryFallbackToWeightedSticky( if len(s.filterGrokFreeQuotaAccounts(ctx, []Account{*account})) == 0 { continue } + upstreamModel := canonicalOpenAIAccountSchedulingModel(account, req.RequestedModel) + now := time.Now() + if isGrokTeamModelRateLimited(account, upstreamModel, now) || + isGrokModelQuotaBlocked(account.ID, upstreamModel, now) { + continue + } result, acquireErr := s.service.tryAcquireAccountSlot(ctx, account.ID, account.Concurrency) if acquireErr != nil { return nil, acquireErr diff --git a/backend/internal/service/openai_gateway_grok.go b/backend/internal/service/openai_gateway_grok.go index 15479ed87e..4160819e67 100644 --- a/backend/internal/service/openai_gateway_grok.go +++ b/backend/internal/service/openai_gateway_grok.go @@ -22,14 +22,14 @@ import ( const ( grokComposerImageBridgeVisionModel = "grok-build-0.1" grokComposerImageBridgeMaxOutputTokens = 512 - grokUpstreamUserAgent = "sub2api-grok/1.0" - grokCLIVersion = xai.CLIClientVersion - grokDefaultResponsesModel = "grok-4.5" - grokRateLimitFallbackCooldown = 2 * time.Minute - grokRateLimitRepeatCooldown = 10 * time.Minute - grokRateLimitSustainedCooldown = 30 * time.Minute - grokRateLimitMaxAdaptiveCooldown = time.Hour - grokRateLimitBackoffQuietPeriod = time.Hour + // grokUpstreamUserAgent lives in grok_upstream_headers.go (shared with TLS header helpers). + grokCLIVersion = xai.CLIClientVersion + grokDefaultResponsesModel = "grok-4.5" + grokRateLimitFallbackCooldown = 2 * time.Minute + grokRateLimitRepeatCooldown = 10 * time.Minute + grokRateLimitSustainedCooldown = 30 * time.Minute + grokRateLimitMaxAdaptiveCooldown = time.Hour + grokRateLimitBackoffQuietPeriod = time.Hour ) func (s *OpenAIGatewayService) forwardGrokResponses( @@ -1099,12 +1099,18 @@ func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Acc // applyGrokCLIHeaders identifies subscription traffic as a supported Grok CLI // version. The CLI gateway rejects otherwise valid OAuth requests without it. +// Identity pins come from package xai so service-layer headers match the final +// transport rewrite on cli-chat-proxy.grok.com. func applyGrokCLIHeaders(headers http.Header) { if headers == nil { return } - headers.Set("User-Agent", grokUpstreamUserAgent) - headers.Set("X-Grok-Client-Version", grokCLIVersion) + version := xai.ResolveCLIVersion() + headers.Set("User-Agent", xai.CLIUserAgent(version)) + headers.Set("X-Grok-Client-Version", version) + headers.Set("x-grok-client-version", version) + headers.Set("x-grok-client-identifier", xai.CLIClientIdentifier) + // Historical mode value expected by some unit tests / older CLI probes. headers.Set("X-Grok-Client-Mode", "interactive") } @@ -1127,6 +1133,15 @@ func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, acco } } + updates := map[string]any{ + grokQuotaSnapshotExtraKey: snapshot, + } + // Also derive the scheduling-threshold extras (grok_sched_*) the evaluator + // reads in grokThresholdCandidates. Without this writer the admin-configured + // Grok auto-pause threshold could never fire (the read side was dead config). + for k, v := range buildGrokSchedulerExtraUpdates(snapshot) { + updates[k] = v + } stateCtx := ctx if hasActiveLimit { var cancel context.CancelFunc @@ -1134,9 +1149,7 @@ func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, acco defer cancel() } if s.accountRepo != nil { - _ = s.accountRepo.UpdateExtra(stateCtx, accountID, map[string]any{ - grokQuotaSnapshotExtraKey: snapshot, - }) + _ = s.accountRepo.UpdateExtra(stateCtx, accountID, updates) } // Error responses are reconciled by handleGrokAccountUpstreamError. Pool-mode // API keys retain the snapshot for observability but leave account health to @@ -1364,6 +1377,85 @@ func (s *OpenAIGatewayService) rateLimitGrok(ctx context.Context, account *Accou } } +// buildGrokSchedulerExtraUpdates derives the grok_sched_* scheduling snapshot +// (utilization percent + reset time) consumed by EvaluateAccountSchedulingThreshold. +// Utilization is the most-constrained of the requests/tokens windows. +func buildGrokSchedulerExtraUpdates(snapshot *xai.QuotaSnapshot) map[string]any { + if snapshot == nil { + return nil + } + util, reset, ok := grokSnapshotUtilization(snapshot) + if !ok { + return nil + } + updates := map[string]any{ + "grok_sched_utilization": util, + "grok_sched_usage_updated_at": time.Now().UTC().Format(time.RFC3339), + } + if reset != nil { + // 防御:调度阈值暂停时长由 grok_sched_reset_at 决定。若上游返回脏的 + // reset 头(例如把相对毫秒 "6000" 误当相对秒解析出 ~33h 的未来时刻), + // 不设上限会把耗尽账号长时间锁死。xAI 配额窗口不会超过一天,因此对 + // 未来时刻做 grokMaxSchedulingResetHorizon 钳制;过去/无效值直接不写。 + now := time.Now() + if reset.After(now) { + capped := *reset + if horizon := now.Add(grokMaxSchedulingResetHorizon); capped.After(horizon) { + capped = horizon + } + updates["grok_sched_reset_at"] = capped.UTC().Format(time.RFC3339) + } + } + return updates +} + +// grokSnapshotUtilization returns the highest window utilization (0-100) across +// the requests/tokens quota windows and the reset time of that window. +func grokSnapshotUtilization(snapshot *xai.QuotaSnapshot) (float64, *time.Time, bool) { + if snapshot == nil { + return 0, nil, false + } + best := -1.0 + var bestReset *time.Time + consider := func(window *xai.QuotaWindow) { + if window == nil || window.Limit == nil || *window.Limit <= 0 || window.Remaining == nil { + return + } + remaining := *window.Remaining + if remaining < 0 { + remaining = 0 + } + util := (1 - float64(remaining)/float64(*window.Limit)) * 100 + if util < 0 { + util = 0 + } + if util > 100 { + util = 100 + } + if util > best { + best = util + if window.ResetUnix != nil { + t := time.Unix(*window.ResetUnix, 0).UTC() + bestReset = &t + } else { + bestReset = nil + } + } + } + consider(snapshot.Requests) + consider(snapshot.Tokens) + if best < 0 { + return 0, nil, false + } + return best, bestReset, true +} + +// grokMaxSchedulingResetHorizon bounds how far into the future a Grok +// scheduling-threshold pause (grok_sched_reset_at) may be set, so a malformed +// upstream reset header can't park an over-threshold account for days. xAI quota +// windows do not exceed ~a day. +const grokMaxSchedulingResetHorizon = 25 * time.Hour + // grokTeamRateLimitModelContextKey carries the upstream model for team cools. type grokTeamRateLimitModelContextKey struct{} diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index 208ff3b6fa..5302ec22b9 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -2082,7 +2082,7 @@ func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *te require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) require.Equal(t, "text/event-stream", upstream.lastReq.Header.Get("Accept")) - require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent")) + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), upstream.lastReq.Header.Get("User-Agent")) require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String()) require.True(t, gjson.GetBytes(upstream.lastBody, "stream_options.include_usage").Bool()) require.True(t, result.Stream) @@ -2255,7 +2255,7 @@ func TestForwardAsChatCompletionsForGrokStreamingStopFallsBackToRawXAIChatComple require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) require.Equal(t, "text/event-stream", upstream.lastReq.Header.Get("Accept")) - require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent")) + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), upstream.lastReq.Header.Get("User-Agent")) require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version")) require.NotEmpty(t, upstream.lastReq.Header.Get(grokConversationIDHeader)) require.NotEqual(t, "native-client-conversation", upstream.lastReq.Header.Get(grokConversationIDHeader)) @@ -2359,7 +2359,7 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) { require.NoError(t, err) require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) - require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent")) + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), upstream.lastReq.Header.Get("User-Agent")) require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version")) require.Equal(t, "grok-experimental", upstream.lastReq.Header.Get("OpenAI-Beta")) require.Empty(t, upstream.lastReq.Header.Get("originator")) @@ -3224,3 +3224,30 @@ func TestIsGrokImageGenerationModel(t *testing.T) { }) } } + +func TestBuildGrokSchedulerExtraUpdates_FeedsThresholdEvaluator(t *testing.T) { + int64p := func(v int64) *int64 { return &v } + resetUnix := time.Now().Add(90 * time.Minute).Unix() + snapshot := &xai.QuotaSnapshot{ + Requests: &xai.QuotaWindow{Limit: int64p(100), Remaining: int64p(30)}, // 70% used + Tokens: &xai.QuotaWindow{Limit: int64p(1000), Remaining: int64p(50), ResetUnix: &resetUnix}, // 95% used (most constrained) + } + + updates := buildGrokSchedulerExtraUpdates(snapshot) + require.NotNil(t, updates) + require.InDelta(t, 95.0, updates["grok_sched_utilization"], 0.001, "picks the most-constrained window") + require.Contains(t, updates, "grok_sched_reset_at") + + // The written extras must actually drive EvaluateAccountSchedulingThreshold + // (proves the previously-dead read side is now fed). + account := &Account{Platform: PlatformGrok, Extra: updates} + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{PlatformGrok: 90}, time.Now()) + require.True(t, decision.ShouldPause) + require.InDelta(t, 95.0, decision.UsedPercent, 0.001) + require.NotNil(t, decision.Until) +} + +func TestBuildGrokSchedulerExtraUpdates_NilWhenNoQuotaWindows(t *testing.T) { + require.Nil(t, buildGrokSchedulerExtraUpdates(&xai.QuotaSnapshot{})) + require.Nil(t, buildGrokSchedulerExtraUpdates(nil)) +} diff --git a/backend/internal/service/openai_gateway_scheduling.go b/backend/internal/service/openai_gateway_scheduling.go index 217f0a72b2..c232232715 100644 --- a/backend/internal/service/openai_gateway_scheduling.go +++ b/backend/internal/service/openai_gateway_scheduling.go @@ -1258,7 +1258,10 @@ func (s *OpenAIGatewayService) listSchedulableAccounts(ctx context.Context, grou platform = normalizeOpenAICompatiblePlatform(platform) if s.schedulerSnapshot != nil { accounts, _, err := s.schedulerSnapshot.ListSchedulableAccounts(ctx, groupID, platform, false) - return accounts, err + if err != nil { + return accounts, err + } + return s.filterOpenAIAccountsBySchedulingThreshold(ctx, accounts), nil } var accounts []Account var err error @@ -1272,7 +1275,7 @@ func (s *OpenAIGatewayService) listSchedulableAccounts(ctx context.Context, grou if err != nil { return nil, fmt.Errorf("query accounts failed: %w", err) } - return accounts, nil + return s.filterOpenAIAccountsBySchedulingThreshold(ctx, accounts), nil } func (s *OpenAIGatewayService) tryAcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int) (*AcquireResult, error) { @@ -1306,6 +1309,9 @@ func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context. if s.isOpenAIAccountRequestRuntimeBlocked(fresh, requestedModel) { return nil } + if s.isOpenAIAccountBlockedBySchedulingThreshold(ctx, fresh) { + return nil + } if s.isOpenAIProxyStreamQuarantined(ctx, fresh) { return nil } @@ -1335,6 +1341,9 @@ func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Co if !isOpenAICompatibleAccountEligibleForRequest(ctx, account, platform, requestedModel, requireCompact, requiredCapability) { return nil } + if s.isOpenAIAccountBlockedBySchedulingThreshold(ctx, account) { + return nil + } if !parentHealthyForShadow(account, s.parentAccountLookup(ctx)) { return nil } @@ -1360,6 +1369,9 @@ func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Co if s.isOpenAIAccountRequestRuntimeBlocked(latest, requestedModel) { return nil } + if s.isOpenAIAccountBlockedBySchedulingThreshold(ctx, latest) { + return nil + } if s.isOpenAIProxyStreamQuarantined(ctx, latest) { return nil } @@ -1386,9 +1398,34 @@ func (s *OpenAIGatewayService) getSchedulableAccount(ctx context.Context, accoun if err != nil || account == nil { return account, err } + if s.isOpenAIAccountBlockedBySchedulingThreshold(ctx, account) { + return nil, nil + } return account, nil } +func (s *OpenAIGatewayService) filterOpenAIAccountsBySchedulingThreshold(ctx context.Context, accounts []Account) []Account { + if len(accounts) == 0 { + return accounts + } + + filtered := make([]Account, 0, len(accounts)) + for i := range accounts { + if s.isOpenAIAccountBlockedBySchedulingThreshold(ctx, &accounts[i]) { + continue + } + filtered = append(filtered, accounts[i]) + } + return filtered +} + +func (s *OpenAIGatewayService) isOpenAIAccountBlockedBySchedulingThreshold(ctx context.Context, account *Account) bool { + if s == nil || s.rateLimitService == nil || account == nil { + return false + } + return s.rateLimitService.ApplyAccountSchedulingThreshold(ctx, account) +} + func (s *OpenAIGatewayService) hydrateSelectedAccount(ctx context.Context, account *Account) (*Account, error) { if account == nil || s.schedulerSnapshot == nil { return account, nil diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 4ec35804c7..eda073c527 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -68,6 +68,29 @@ func (r stubOpenAIAccountRepo) GetByID(ctx context.Context, id int64) (*Account, return nil, errors.New("account not found") } +func (r stubOpenAIAccountRepo) GetByIDs(ctx context.Context, ids []int64) ([]*Account, error) { + if len(ids) == 0 { + return []*Account{}, nil + } + index := make(map[int64]*Account, len(r.accounts)) + for i := range r.accounts { + account := &r.accounts[i] + index[account.ID] = account + } + out := make([]*Account, 0, len(ids)) + seen := make(map[int64]struct{}, len(ids)) + for _, id := range ids { + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + if account, ok := index[id]; ok { + out = append(out, account) + } + } + return out, nil +} + func (r stubOpenAIAccountRepo) ListSchedulableByGroupIDAndPlatform(ctx context.Context, groupID int64, platform string) ([]Account, error) { var result []Account for _, acc := range r.accounts { diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go index 802146cd3f..73fcb08586 100644 --- a/backend/internal/service/openai_ws_http_bridge_test.go +++ b/backend/internal/service/openai_ws_http_bridge_test.go @@ -671,7 +671,7 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridgeAndPreservesMa require.Len(t, upstream.bodies, 3) require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) - require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent")) + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), upstream.lastReq.Header.Get("User-Agent")) require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version")) require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.bodies[0], "model").String()) require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.bodies[1], "model").String()) diff --git a/backend/internal/service/ratelimit_service.go b/backend/internal/service/ratelimit_service.go index 266f806e83..5eda423fde 100644 --- a/backend/internal/service/ratelimit_service.go +++ b/backend/internal/service/ratelimit_service.go @@ -138,6 +138,97 @@ func (s *RateLimitService) notifyAccountSchedulingBlockCleared(accountID int64) s.runtimeBlocker.ClearAccountSchedulingBlock(accountID) } +// ApplyAccountSchedulingThreshold evaluates admin-configured per-platform +// utilization thresholds and, when breached, parks the account as temp- +// unschedulable until the winning window resets. Returns true when the account +// is blocked (either newly or already paused for the same threshold reason). +func (s *RateLimitService) ApplyAccountSchedulingThreshold(ctx context.Context, account *Account) bool { + if s == nil || s.settingService == nil || s.accountRepo == nil || account == nil || account.ID <= 0 { + return false + } + if !account.IsActive() || !account.Schedulable { + return false + } + + now := time.Now().UTC() + thresholds := s.settingService.GetAccountSchedulingThresholds(ctx) + decision := EvaluateAccountSchedulingThreshold(account, thresholds, now) + if !decision.ShouldPause || decision.Until == nil || !decision.Until.After(now) { + return false + } + + reason := BuildDetailedAccountSchedulingThresholdReason(AccountSchedulingThresholdReasonInput{ + Platform: decision.Platform, + Window: decision.Window, + Scope: decision.Scope, + ThresholdPercent: decision.ThresholdPercent, + UsedPercent: decision.UsedPercent, + Until: *decision.Until, + Now: now, + }) + + if accountHasSameSchedulingThresholdPause(account, *decision.Until, reason) { + return true + } + if !account.IsSchedulable() { + return false + } + + account.TempUnschedulableUntil = cloneTimePtr(decision.Until) + account.TempUnschedulableReason = reason + s.notifyAccountSchedulingBlocked(account, *decision.Until, "account_scheduling_threshold") + + if err := s.accountRepo.SetTempUnschedulable(ctx, account.ID, *decision.Until, reason); err != nil { + slog.Warn("account_scheduling_threshold_set_temp_unsched_failed", + "account_id", account.ID, + "platform", decision.Platform, + "window", decision.Window, + "scope", decision.Scope, + "threshold_percent", decision.ThresholdPercent, + "used_percent", decision.UsedPercent, + "until", decision.Until.UTC(), + "error", err) + } else if s.tempUnschedCache != nil { + if state := tempUnschedStateFromStoredReason(reason, decision.Until.Unix()); state != nil { + if err := s.tempUnschedCache.SetTempUnsched(ctx, account.ID, state); err != nil { + slog.Warn("account_scheduling_threshold_cache_set_failed", "account_id", account.ID, "error", err) + } + } + } + + slog.Info("account_scheduling_threshold_temp_unschedulable", + "account_id", account.ID, + "platform", decision.Platform, + "window", decision.Window, + "scope", decision.Scope, + "threshold_percent", decision.ThresholdPercent, + "used_percent", decision.UsedPercent, + "until", decision.Until.UTC()) + return true +} + +func accountHasSameSchedulingThresholdPause(account *Account, until time.Time, reason string) bool { + if account == nil || account.TempUnschedulableUntil == nil { + return false + } + if account.TempUnschedulableUntil.UTC().Unix() != until.UTC().Unix() { + return false + } + + existing, ok := parseTempUnschedReasonPayload(account.TempUnschedulableReason) + if !ok || existing.Source != AccountSchedulingThresholdReasonSource { + return false + } + next, ok := parseTempUnschedReasonPayload(reason) + if !ok || next.Source != AccountSchedulingThresholdReasonSource { + return false + } + + existing.TriggeredAtUnix = 0 + next.TriggeredAtUnix = 0 + return existing == next +} + // ErrorPolicyResult 表示错误策略检查的结果 type ErrorPolicyResult int diff --git a/backend/internal/service/ratelimit_service_scheduling_threshold_test.go b/backend/internal/service/ratelimit_service_scheduling_threshold_test.go new file mode 100644 index 0000000000..5069c4dec3 --- /dev/null +++ b/backend/internal/service/ratelimit_service_scheduling_threshold_test.go @@ -0,0 +1,166 @@ +//go:build unit + +package service + +import ( + "context" + "encoding/json" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +func TestRateLimitService_ApplyAccountSchedulingThreshold_SetsTempUnschedulable(t *testing.T) { + accountSchedulingThresholdsSF.Forget(SettingKeyAccountSchedulingThresholds) + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{}) + + settingsRepo := newMockSettingRepo() + settingsRepo.data[SettingKeyAccountSchedulingThresholds] = `{"openai":80}` + + accountRepo := &rateLimitAccountRepoStub{} + rl := NewRateLimitService(accountRepo, nil, &config.Config{}, nil, nil) + rl.SetSettingService(NewSettingService(settingsRepo, &config.Config{})) + + until := time.Now().UTC().Add(6 * time.Hour) + account := &Account{ + ID: 1001, + Platform: PlatformOpenAI, + Status: StatusActive, + Schedulable: true, + Extra: map[string]any{ + "codex_7d_used_percent": 91.5, + "codex_7d_reset_at": until.Format(time.RFC3339), + }, + } + + blocked := rl.ApplyAccountSchedulingThreshold(context.Background(), account) + + require.True(t, blocked) + require.Equal(t, 1, accountRepo.tempCalls) + require.NotNil(t, account.TempUnschedulableUntil) + require.WithinDuration(t, until, *account.TempUnschedulableUntil, time.Second) + require.True(t, IsAccountSchedulingThresholdReason(accountRepo.lastTempReason)) + + var payload map[string]any + require.NoError(t, json.Unmarshal([]byte(accountRepo.lastTempReason), &payload)) + require.Equal(t, PlatformOpenAI, payload["platform"]) + require.Equal(t, "7d", payload["window"]) + require.Equal(t, float64(80), payload["threshold_percent"]) + require.Equal(t, float64(91.5), payload["used_percent"]) + require.Contains(t, payload["error_message"], "91.5% used >= 80%") +} + +func TestRateLimitService_ApplyAccountSchedulingThreshold_UsesAccountOverrideInReason(t *testing.T) { + accountSchedulingThresholdsSF.Forget(SettingKeyAccountSchedulingThresholds) + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{}) + + settingsRepo := newMockSettingRepo() + settingsRepo.data[SettingKeyAccountSchedulingThresholds] = `{"openai":90}` + + accountRepo := &rateLimitAccountRepoStub{} + rl := NewRateLimitService(accountRepo, nil, &config.Config{}, nil, nil) + rl.SetSettingService(NewSettingService(settingsRepo, &config.Config{})) + + until := time.Now().UTC().Add(6 * time.Hour) + account := &Account{ + ID: 1003, + Platform: PlatformOpenAI, + Status: StatusActive, + Schedulable: true, + Credentials: map[string]any{ + "account_scheduling_threshold": 80, + }, + Extra: map[string]any{ + "codex_7d_used_percent": 85.5, + "codex_7d_reset_at": until.Format(time.RFC3339), + }, + } + + blocked := rl.ApplyAccountSchedulingThreshold(context.Background(), account) + + require.True(t, blocked) + require.Equal(t, 1, accountRepo.tempCalls) + + var payload map[string]any + require.NoError(t, json.Unmarshal([]byte(accountRepo.lastTempReason), &payload)) + require.Equal(t, float64(80), payload["threshold_percent"]) + require.Equal(t, float64(85.5), payload["used_percent"]) + require.Contains(t, payload["error_message"], "85.5% used >= 80%") +} + +func TestRateLimitService_ApplyAccountSchedulingThreshold_SkipsDuplicateTempUnschedulable(t *testing.T) { + accountSchedulingThresholdsSF.Forget(SettingKeyAccountSchedulingThresholds) + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{}) + + settingsRepo := newMockSettingRepo() + settingsRepo.data[SettingKeyAccountSchedulingThresholds] = `{"openai":80}` + + accountRepo := &rateLimitAccountRepoStub{} + rl := NewRateLimitService(accountRepo, nil, &config.Config{}, nil, nil) + rl.SetSettingService(NewSettingService(settingsRepo, &config.Config{})) + + until := time.Now().UTC().Add(6 * time.Hour).Truncate(time.Second) + existingReason := BuildDetailedAccountSchedulingThresholdReason(AccountSchedulingThresholdReasonInput{ + Platform: PlatformOpenAI, + Window: "7d", + ThresholdPercent: 80, + UsedPercent: 91.5, + Until: until, + Now: until.Add(-time.Hour), + }) + account := &Account{ + ID: 1002, + Platform: PlatformOpenAI, + Status: StatusActive, + Schedulable: true, + TempUnschedulableUntil: &until, + TempUnschedulableReason: existingReason, + Extra: map[string]any{ + "codex_7d_used_percent": 91.5, + "codex_7d_reset_at": until.Format(time.RFC3339), + }, + } + + blocked := rl.ApplyAccountSchedulingThreshold(context.Background(), account) + + require.True(t, blocked) + require.Equal(t, 0, accountRepo.tempCalls) + require.Equal(t, existingReason, account.TempUnschedulableReason) + require.NotNil(t, account.TempUnschedulableUntil) + require.True(t, until.Equal(*account.TempUnschedulableUntil)) +} + +func TestRateLimitService_ApplyAccountSchedulingThreshold_UnsupportedPlatformDoesNotBlock(t *testing.T) { + accountSchedulingThresholdsSF.Forget(SettingKeyAccountSchedulingThresholds) + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{}) + + settingsRepo := newMockSettingRepo() + settingsRepo.data[SettingKeyAccountSchedulingThresholds] = `{"openai":80}` + + accountRepo := &rateLimitAccountRepoStub{} + rl := NewRateLimitService(accountRepo, nil, &config.Config{}, nil, nil) + rl.SetSettingService(NewSettingService(settingsRepo, &config.Config{})) + + account := &Account{ + ID: 2002, + Platform: PlatformKiro, + Status: StatusActive, + Schedulable: true, + Credentials: map[string]any{ + "account_scheduling_threshold": 1, + }, + Extra: map[string]any{ + "kiro_sched_utilization": 99.0, + "kiro_sched_reset_at": time.Now().UTC().Add(24 * time.Hour).Format(time.RFC3339), + }, + } + + blocked := rl.ApplyAccountSchedulingThreshold(context.Background(), account) + + require.False(t, blocked) + require.Equal(t, 0, accountRepo.tempCalls) + require.Nil(t, account.TempUnschedulableUntil) + require.Empty(t, account.TempUnschedulableReason) +} diff --git a/backend/internal/service/setting_features.go b/backend/internal/service/setting_features.go index 6fba70948f..cfb83f0b5f 100644 --- a/backend/internal/service/setting_features.go +++ b/backend/internal/service/setting_features.go @@ -11,6 +11,7 @@ import ( "math" "strconv" "strings" + "time" ) // IsRegistrationEnabled 检查是否开放注册 @@ -1078,6 +1079,62 @@ func (s *SettingService) GetDefaultPlatformQuotas(ctx context.Context) (map[stri return out, nil // 补齐全部允许 platform key,保持与旧实现一致的下游契约 } +// GetAccountSchedulingThresholds returns per-platform auto-pause thresholds (1..100). +// 100 disables the threshold for that platform. Hot-path cached with singleflight. +func (s *SettingService) GetAccountSchedulingThresholds(ctx context.Context) map[string]int { + if s == nil || s.settingRepo == nil { + return defaultAccountSchedulingThresholds() + } + if cached, ok := accountSchedulingThresholdsCache.Load().(*cachedAccountSchedulingThresholds); ok { + if cached != nil && len(cached.thresholds) > 0 && time.Now().UnixNano() < cached.expiresAt { + return cloneAccountSchedulingThresholds(cached.thresholds) + } + } + + result, err, _ := accountSchedulingThresholdsSF.Do(SettingKeyAccountSchedulingThresholds, func() (any, error) { + if cached, ok := accountSchedulingThresholdsCache.Load().(*cachedAccountSchedulingThresholds); ok { + if cached != nil && len(cached.thresholds) > 0 && time.Now().UnixNano() < cached.expiresAt { + return cloneAccountSchedulingThresholds(cached.thresholds), nil + } + } + + thresholds := defaultAccountSchedulingThresholds() + dbCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), accountSchedulingThresholdsDBTimeout) + defer cancel() + + raw, err := s.settingRepo.GetValue(dbCtx, SettingKeyAccountSchedulingThresholds) + if err != nil { + slog.Warn("failed to get account scheduling thresholds, falling back to defaults", "error", err) + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{ + thresholds: cloneAccountSchedulingThresholds(thresholds), + expiresAt: time.Now().Add(accountSchedulingThresholdsErrorTTL).UnixNano(), + }) + return cloneAccountSchedulingThresholds(thresholds), nil + } + + if trimmed := strings.TrimSpace(raw); trimmed != "" { + if parsed, err := parseAccountSchedulingThresholdsSetting(trimmed); err != nil { + slog.Warn("failed to parse account scheduling thresholds, falling back to defaults", "error", err) + } else { + thresholds = parsed + } + } + + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{ + thresholds: cloneAccountSchedulingThresholds(thresholds), + expiresAt: time.Now().Add(accountSchedulingThresholdsCacheTTL).UnixNano(), + }) + return cloneAccountSchedulingThresholds(thresholds), nil + }) + if err != nil { + return defaultAccountSchedulingThresholds() + } + if thresholds, ok := result.(map[string]int); ok { + return cloneAccountSchedulingThresholds(thresholds) + } + return defaultAccountSchedulingThresholds() +} + // GetAuthSourcePlatformQuotas 读取指定 auth source 的 platform quota 覆盖(仅返回有配置的平台,override 语义)。 func (s *SettingService) GetAuthSourcePlatformQuotas(ctx context.Context, source string) map[string]*DefaultPlatformQuotaSetting { out := map[string]*DefaultPlatformQuotaSetting{} diff --git a/backend/internal/service/setting_gateway_runtime.go b/backend/internal/service/setting_gateway_runtime.go index d2d614b34e..2f13278e3b 100644 --- a/backend/internal/service/setting_gateway_runtime.go +++ b/backend/internal/service/setting_gateway_runtime.go @@ -72,6 +72,19 @@ const gatewayForwardingCacheTTL = 60 * time.Second const gatewayForwardingErrorTTL = 5 * time.Second const gatewayForwardingDBTimeout = 5 * time.Second +// cachedAccountSchedulingThresholds 缓存平台自动停调阈值(进程内缓存,60s TTL) +type cachedAccountSchedulingThresholds struct { + thresholds map[string]int + expiresAt int64 // unix nano +} + +var accountSchedulingThresholdsCache atomic.Value // *cachedAccountSchedulingThresholds +var accountSchedulingThresholdsSF singleflight.Group + +const accountSchedulingThresholdsCacheTTL = 60 * time.Second +const accountSchedulingThresholdsErrorTTL = 5 * time.Second +const accountSchedulingThresholdsDBTimeout = 5 * time.Second + // cachedAntigravityUserAgentVersion 缓存 Antigravity UA 版本号(进程内缓存,60s TTL) type cachedAntigravityUserAgentVersion struct { version string diff --git a/backend/internal/service/setting_parse.go b/backend/internal/service/setting_parse.go index a96dc44f70..d748e80061 100644 --- a/backend/internal/service/setting_parse.go +++ b/backend/internal/service/setting_parse.go @@ -941,6 +941,14 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin result.DefaultPlatformQuotas = parsed } } + result.AccountSchedulingThresholds = defaultAccountSchedulingThresholds() + if raw := strings.TrimSpace(settings[SettingKeyAccountSchedulingThresholds]); raw != "" { + if thresholds, err := parseAccountSchedulingThresholdsSetting(raw); err != nil { + slog.Warn("[Setting] parseSettings: unmarshal account_scheduling_thresholds failed", "error", err) + } else { + result.AccountSchedulingThresholds = thresholds + } + } result.AllowUserViewErrorRequests = settings[SettingKeyAllowUserViewErrorRequests] == "true" // default false diff --git a/backend/internal/service/setting_service_platform_threshold_test.go b/backend/internal/service/setting_service_platform_threshold_test.go new file mode 100644 index 0000000000..7b95cb8f78 --- /dev/null +++ b/backend/internal/service/setting_service_platform_threshold_test.go @@ -0,0 +1,151 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +func newSettingServiceForPlatformThresholdTest(seed map[string]string) *SettingService { + accountSchedulingThresholdsSF.Forget(SettingKeyAccountSchedulingThresholds) + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{}) + repo := newMockSettingRepo() + for k, v := range seed { + repo.data[k] = v + } + return NewSettingService(repo, &config.Config{}) +} + +func TestPlatformSchedulingThresholds_RoundTrip_DefaultsAndStoredValues(t *testing.T) { + svc := newSettingServiceForPlatformThresholdTest(nil) + + got := svc.parseSettings(map[string]string{}) + require.Equal(t, map[string]int{ + PlatformOpenAI: 100, + PlatformAnthropic: 100, + PlatformGrok: 100, + }, got.AccountSchedulingThresholds) + + got = svc.parseSettings(map[string]string{ + SettingKeyAccountSchedulingThresholds: `{"openai":91,"grok":77,"gemini":85,"kiro":99}`, + }) + require.Equal(t, 91, got.AccountSchedulingThresholds[PlatformOpenAI]) + require.Equal(t, 100, got.AccountSchedulingThresholds[PlatformAnthropic]) + require.Equal(t, 77, got.AccountSchedulingThresholds[PlatformGrok]) + require.NotContains(t, got.AccountSchedulingThresholds, PlatformGemini) + require.NotContains(t, got.AccountSchedulingThresholds, "kiro") +} + +func TestBuildSystemSettingsUpdates_PersistsAccountSchedulingThresholds(t *testing.T) { + svc := newSettingServiceForPlatformThresholdTest(nil) + + updates, err := svc.buildSystemSettingsUpdates(context.Background(), &SystemSettings{ + AccountSchedulingThresholds: map[string]int{ + PlatformOpenAI: 91, + PlatformAnthropic: 88, + PlatformGrok: 77, + }, + }) + require.NoError(t, err) + require.JSONEq(t, `{"openai":91,"anthropic":88,"grok":77}`, updates[SettingKeyAccountSchedulingThresholds]) +} + +func TestValidateAndNormalizeAccountSchedulingThresholds_FillsMissingPlatforms(t *testing.T) { + normalized, err := validateAndNormalizeAccountSchedulingThresholds(map[string]int{ + PlatformOpenAI: 91, + }) + require.NoError(t, err) + require.Equal(t, 91, normalized[PlatformOpenAI]) + require.Equal(t, 100, normalized[PlatformAnthropic]) + require.Equal(t, 100, normalized[PlatformGrok]) + require.NotContains(t, normalized, PlatformGemini) + require.NotContains(t, normalized, "kiro") + require.NotContains(t, normalized, PlatformAntigravity) +} + +func TestValidateAndNormalizeAccountSchedulingThresholds_RejectsUnsupportedPlatforms(t *testing.T) { + _, err := validateAndNormalizeAccountSchedulingThresholds(map[string]int{ + PlatformGemini: 85, + }) + require.Error(t, err) +} + +func TestUpdateSettings_StoresAccountSchedulingThresholds(t *testing.T) { + svc := newSettingServiceForPlatformThresholdTest(nil) + + err := svc.UpdateSettings(context.Background(), &SystemSettings{ + AccountSchedulingThresholds: map[string]int{ + PlatformOpenAI: 92, + PlatformAnthropic: 89, + PlatformGrok: 76, + }, + }) + require.NoError(t, err) + + got := svc.parseSettings(map[string]string{ + SettingKeyAccountSchedulingThresholds: svc.settingRepo.(*mockSettingRepo).data[SettingKeyAccountSchedulingThresholds], + }) + require.Equal(t, 92, got.AccountSchedulingThresholds[PlatformOpenAI]) + require.Equal(t, 89, got.AccountSchedulingThresholds[PlatformAnthropic]) + require.Equal(t, 76, got.AccountSchedulingThresholds[PlatformGrok]) + require.NotContains(t, got.AccountSchedulingThresholds, "kiro") +} + +func TestGetAccountSchedulingThresholds_ReadsStoredValue(t *testing.T) { + svc := newSettingServiceForPlatformThresholdTest(map[string]string{ + SettingKeyAccountSchedulingThresholds: `{"openai":93,"grok":88,"kiro":87}`, + }) + + got := svc.GetAccountSchedulingThresholds(context.Background()) + + require.Equal(t, 93, got[PlatformOpenAI]) + require.Equal(t, 100, got[PlatformAnthropic]) + require.Equal(t, 88, got[PlatformGrok]) + require.NotContains(t, got, "kiro") +} + +func TestUpdateSettings_OmittedAccountSchedulingThresholdsDoesNotCacheDefaults(t *testing.T) { + svc := newSettingServiceForPlatformThresholdTest(map[string]string{ + SettingKeyAccountSchedulingThresholds: `{"openai":85,"grok":88,"kiro":87}`, + }) + + err := svc.UpdateSettings(context.Background(), &SystemSettings{ + FrontendURL: "https://example.test", + }) + require.NoError(t, err) + + got := svc.GetAccountSchedulingThresholds(context.Background()) + require.Equal(t, 85, got[PlatformOpenAI]) + require.Equal(t, 88, got[PlatformGrok]) + require.NotContains(t, got, "kiro") +} + +func TestAccountSchedulingThresholds_InvalidStoredValueUsesSameDefaultsInSettingsAndCache(t *testing.T) { + svc := newSettingServiceForPlatformThresholdTest(map[string]string{ + SettingKeyAccountSchedulingThresholds: `{"openai":0,"grok":88,"kiro":87}`, + }) + + settings := svc.parseSettings(map[string]string{ + SettingKeyAccountSchedulingThresholds: `{"openai":0,"grok":88,"kiro":87}`, + }) + cached := svc.GetAccountSchedulingThresholds(context.Background()) + + require.Equal(t, settings.AccountSchedulingThresholds, cached) + require.Equal(t, 100, cached[PlatformOpenAI]) + require.Equal(t, 88, cached[PlatformGrok]) + require.NotContains(t, cached, "kiro") +} + +func TestGetAccountSchedulingThresholds_NilRepoReturnsDefaults(t *testing.T) { + svc := &SettingService{} + got := svc.GetAccountSchedulingThresholds(context.Background()) + require.Equal(t, map[string]int{ + PlatformOpenAI: 100, + PlatformAnthropic: 100, + PlatformGrok: 100, + }, got) +} diff --git a/backend/internal/service/setting_update.go b/backend/internal/service/setting_update.go index bdace5a429..9270dead24 100644 --- a/backend/internal/service/setting_update.go +++ b/backend/internal/service/setting_update.go @@ -519,12 +519,88 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting } updates[SettingKeyDefaultPlatformQuotas] = string(blob) } + if settings.AccountSchedulingThresholds != nil { + normalized, err := validateAndNormalizeAccountSchedulingThresholds(settings.AccountSchedulingThresholds) + if err != nil { + return nil, err + } + blob, err := json.Marshal(normalized) + if err != nil { + return nil, fmt.Errorf("marshal account scheduling thresholds: %w", err) + } + updates[SettingKeyAccountSchedulingThresholds] = string(blob) + } updates[SettingKeyAllowUserViewErrorRequests] = strconv.FormatBool(settings.AllowUserViewErrorRequests) return updates, nil } +func defaultAccountSchedulingThresholds() map[string]int { + return map[string]int{ + PlatformOpenAI: 100, + PlatformAnthropic: 100, + PlatformGrok: 100, + } +} + +func validateAndNormalizeAccountSchedulingThresholds(input map[string]int) (map[string]int, error) { + normalized := defaultAccountSchedulingThresholds() + for platform, value := range input { + allowed := false + for _, item := range AllowedSchedulingThresholdPlatforms { + if item == platform { + allowed = true + break + } + } + if !allowed { + return nil, infraerrors.BadRequest("INVALID_ACCOUNT_SCHEDULING_THRESHOLDS", fmt.Sprintf("unknown platform %q", platform)) + } + if value < 1 || value > 100 { + return nil, infraerrors.BadRequest("INVALID_ACCOUNT_SCHEDULING_THRESHOLDS", "platform scheduling threshold must be between 1 and 100") + } + normalized[platform] = value + } + return normalized, nil +} + +func parseAccountSchedulingThresholdsSetting(raw string) (map[string]int, error) { + thresholds := defaultAccountSchedulingThresholds() + raw = strings.TrimSpace(raw) + if raw == "" { + return thresholds, nil + } + parsed := map[string]int{} + if err := json.Unmarshal([]byte(raw), &parsed); err != nil { + return thresholds, err + } + for _, platform := range AllowedSchedulingThresholdPlatforms { + if value, ok := parsed[platform]; ok { + thresholds[platform] = boundedIntOrDefault(value, 1, 100, 100) + } + } + return thresholds, nil +} + +func boundedIntOrDefault(value, minValue, maxValue, defaultValue int) int { + if value < minValue || value > maxValue { + return defaultValue + } + return value +} + +func cloneAccountSchedulingThresholds(input map[string]int) map[string]int { + if len(input) == 0 { + return defaultAccountSchedulingThresholds() + } + cloned := make(map[string]int, len(input)) + for key, value := range input { + cloned[key] = value + } + return cloned +} + // validateDefaultPlatformQuotaMap 校验 platform quota map 的合法性: // 平台名须在 AllowedQuotaPlatforms 白名单内,每个非 nil 上限须 finite 且 >= 0。 // 系统层和 auth-source 层共用此 helper。 @@ -680,6 +756,20 @@ func (s *SettingService) refreshCachedSettings(settings *SystemSettings) { expiresAt: 0, }) } + accountSchedulingThresholdsSF.Forget(SettingKeyAccountSchedulingThresholds) + if settings.AccountSchedulingThresholds != nil { + normalizedThresholds, err := validateAndNormalizeAccountSchedulingThresholds(settings.AccountSchedulingThresholds) + if err != nil { + normalizedThresholds = defaultAccountSchedulingThresholds() + } + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{ + thresholds: cloneAccountSchedulingThresholds(normalizedThresholds), + expiresAt: time.Now().Add(accountSchedulingThresholdsCacheTTL).UnixNano(), + }) + } else { + // Partial/omitted payload: clear cache so the next hot-path read reloads from DB. + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{}) + } if s.cfg != nil { s.cfg.SetForwardedClientIPSettings(settings.APIKeyACLTrustForwardedIP, settings.ForwardedClientIPHeaders) } diff --git a/backend/internal/service/settings_view.go b/backend/internal/service/settings_view.go index bf262de0bd..d8563b937c 100644 --- a/backend/internal/service/settings_view.go +++ b/backend/internal/service/settings_view.go @@ -296,6 +296,9 @@ type SystemSettings struct { // 系统全局默认平台配额(key = platform,nil/缺省 = 不限制) DefaultPlatformQuotas map[string]*DefaultPlatformQuotaSetting `json:"default_platform_quotas"` + // 系统全局账号自动停调阈值(key = platform,100 = disabled) + AccountSchedulingThresholds map[string]int `json:"account_scheduling_thresholds"` + // 允许终端用户在用量页查看自己的失败请求 AllowUserViewErrorRequests bool } diff --git a/frontend/src/api/admin/accounts.ts b/frontend/src/api/admin/accounts.ts index bd9c54f8e1..21ce9f092f 100644 --- a/frontend/src/api/admin/accounts.ts +++ b/frontend/src/api/admin/accounts.ts @@ -317,6 +317,19 @@ export async function getUsage(id: number, source?: 'passive' | 'active', force? return data } +export interface BatchAccountUsageResponse { + usage: Record + errors: Record +} + +export async function getBatchUsage(accountIds: number[], force?: boolean): Promise { + const { data } = await apiClient.post('/admin/accounts/usage/batch', { + account_ids: accountIds, + force: force === true + }) + return data +} + /** * Clear account rate limit status * @param id - Account ID @@ -986,6 +999,7 @@ export const accountsAPI = { getStats, clearError, getUsage, + getBatchUsage, getTodayStats, getBatchTodayStats, clearRateLimit, diff --git a/frontend/src/api/admin/settings.ts b/frontend/src/api/admin/settings.ts index d825aac26a..434e0b9214 100644 --- a/frontend/src/api/admin/settings.ts +++ b/frontend/src/api/admin/settings.ts @@ -32,6 +32,38 @@ export type DefaultPlatformQuotasMap = Partial + +export const SCHEDULING_THRESHOLD_PLATFORMS: SchedulingThresholdPlatformType[] = [ + "openai", + "anthropic", + "grok", +] + +export function normalizeAccountSchedulingThresholdsMap( + input?: Partial> | null, +): AccountSchedulingThresholdsMap { + const result = {} as AccountSchedulingThresholdsMap + for (const platform of SCHEDULING_THRESHOLD_PLATFORMS) { + const value = input?.[platform] + result[platform] = typeof value === "number" && Number.isFinite(value) + ? Math.min(100, Math.max(1, Math.trunc(value))) + : 100 + } + return result +} + +export function sanitizeAccountSchedulingThresholdsMap( + input?: Partial> | null, +): AccountSchedulingThresholdsMap { + return normalizeAccountSchedulingThresholdsMap(input) +} + /** 归一化为全 4 平台 × 3 窗口(缺失填 null),供模板非空绑定 */ export function normalizePlatformQuotasMap(input?: DefaultPlatformQuotasMap | null): DefaultPlatformQuotasMap { const result: DefaultPlatformQuotasMap = {} @@ -559,6 +591,9 @@ export interface SystemSettings { grok_default_text_model: string; grok_cross_client_model_map_enabled: boolean; + // Per-platform account auto-pause thresholds (100 = disabled) + account_scheduling_thresholds: AccountSchedulingThresholdsMap; + // Identity patch configuration (Claude -> Gemini) enable_identity_patch: boolean; identity_patch_prompt: string; @@ -876,6 +911,7 @@ export interface UpdateSettingsRequest { fallback_model_antigravity?: string; grok_default_text_model?: string; grok_cross_client_model_map_enabled?: boolean; + account_scheduling_thresholds?: AccountSchedulingThresholdsMap; enable_identity_patch?: boolean; identity_patch_prompt?: string; ops_monitoring_enabled?: boolean; diff --git a/frontend/src/components/account/AccountUsageCell.vue b/frontend/src/components/account/AccountUsageCell.vue index bce09b37c8..d8d9dec56a 100644 --- a/frontend/src/components/account/AccountUsageCell.vue +++ b/frontend/src/components/account/AccountUsageCell.vue @@ -393,6 +393,29 @@ :show-now-when-idle="true" color="indigo" /> + +
+ + {{ t('admin.accounts.usageWindow.grokUsed') }} + {{ grokBillingMoneySummary.used }}/{{ grokBillingMoneySummary.limit }} + + + {{ grokBillingMoneySummary.usedPercent }}% + +
void) | null }>(), { todayStats: null, todayStatsLoading: false, - manualRefreshToken: 0 + manualRefreshToken: 0, + batchedUsage: null, + batchedUsageError: null, + batchedUsageLoading: false, + requestBatchedUsage: null } ) const emit = defineEmits<{ 'account-updated': [account: Account] + 'usage-loaded': [usage: AccountUsageInfo] }>() const { t } = useI18n() @@ -674,6 +706,9 @@ const loading = ref(false) const activeQueryLoading = ref(false) const error = ref(null) const usageInfo = ref(null) +watch(usageInfo, (usage) => { + if (usage) emit('usage-loaded', usage) +}) const suppressOpenAIUsageRefreshUntil = ref(0) const rootRef = ref(null) const isDesktopViewport = ref( @@ -713,6 +748,8 @@ const shouldFetchUsage = computed(() => { return false }) +const isBatchManaged = computed(() => typeof props.requestBatchedUsage === 'function') + const showGeminiTodayStats = computed(() => { return props.account.platform === 'gemini' && props.account.type === 'service_account' }) @@ -1094,6 +1131,58 @@ const grokWeeklyBillingBar = computed((): GrokQuotaBarInfo | null => { resetsAt: billing.period_end || null } }) +// Monthly used/limit % from billing probe (used_percent or derived from cents). +const grokMonthlyBillingBar = computed((): GrokQuotaBarInfo | null => { + const billing = grokBilling.value + if (!billing) return null + let utilization: number | null = null + if (billing.used_percent != null && Number.isFinite(billing.used_percent)) { + utilization = billing.used_percent + } else if ( + billing.monthly_limit_cents != null && + billing.monthly_limit_cents > 0 && + billing.used_cents != null + ) { + utilization = (billing.used_cents / billing.monthly_limit_cents) * 100 + } + if (utilization == null) return null + // Avoid duplicating the weekly bar when period_type is weekly-only without monthly. + if (billing.period_type?.toLowerCase() === 'weekly' && billing.monthly_limit_cents == null) { + return null + } + return { + utilization: Math.min(100, Math.max(0, utilization)), + resetsAt: billing.billing_period_end || billing.period_end || null + } +}) +const formatGrokMoneyFromCents = (cents?: number | null) => { + if (cents == null || Number.isNaN(cents)) return '0' + const dollars = cents / 100 + if (dollars >= 1000) return formatCompactNumber(dollars) + if (dollars >= 100) return dollars.toFixed(0) + if (dollars >= 10) return dollars.toFixed(1) + return dollars.toFixed(2) +} +// Absolute monthly used/limit derived from cents fields already on GrokBillingSummary. +// (personal-dev has separate prepaid/on_demand fields; HEAD only has cents.) +const grokBillingMoneySummary = computed(() => { + const billing = grokBilling.value + if (!billing) return null + const limit = billing.monthly_limit_cents + const used = billing.used_cents + if ((limit == null || limit <= 0) && (used == null || used <= 0)) return null + let usedPercent: number | null = null + if (billing.used_percent != null && Number.isFinite(billing.used_percent)) { + usedPercent = Math.round(Math.min(100, Math.max(0, billing.used_percent))) + } else if (limit != null && limit > 0 && used != null) { + usedPercent = Math.round(Math.min(100, Math.max(0, (used / limit) * 100))) + } + return { + used: formatGrokMoneyFromCents(used ?? 0), + limit: formatGrokMoneyFromCents(limit ?? 0), + usedPercent + } +}) const grokPlanLabelIsFree = (value: string) => value.includes('free') || value.includes('basic') const grokPlanLabelIsPaid = (value: string) => { return value !== '' && !grokPlanLabelIsFree(value) && !value.includes('unknown') @@ -1136,7 +1225,16 @@ const grokFreeTokenBar = computed(() => { }) const grokQuotaUnknown = computed(() => { if (props.account.platform !== 'grok') return false - if (grokBilling.value || grokFreeTokenBar.value || grokRequestQuotaBar.value || grokTokenQuotaBar.value) return false + if ( + grokBilling.value || + grokBillingMoneySummary.value || + grokMonthlyBillingBar.value || + grokFreeTokenBar.value || + grokRequestQuotaBar.value || + grokTokenQuotaBar.value + ) { + return false + } return usageInfo.value?.grok_quota_snapshot_state !== 'observed' }) const grokQuotaUnknownLabel = computed(() => { @@ -1273,8 +1371,24 @@ const isAnthropicOAuthOrSetupToken = computed(() => { return props.account.platform === 'anthropic' && (props.account.type === 'oauth' || props.account.type === 'setup-token') }) +const requestParentBatchUsage = (options?: { force?: boolean }) => { + if (!isBatchManaged.value || !shouldFetchUsage.value) return + props.requestBatchedUsage?.(props.account, options) +} + +const syncManagedUsageState = () => { + if (!isBatchManaged.value) return + usageInfo.value = props.batchedUsage ?? null + error.value = props.batchedUsageError ?? null + loading.value = props.batchedUsageLoading === true +} + const loadUsage = async (options?: { source?: 'passive' | 'active'; bypassCache?: boolean }) => { if (!shouldFetchUsage.value) return + if (isBatchManaged.value) { + requestParentBatchUsage({ force: options?.bypassCache === true }) + return + } // Check cache if (!options?.bypassCache) { @@ -1529,11 +1643,49 @@ onMounted(() => { } } + if (isBatchManaged.value) { + syncManagedUsageState() + requestParentBatchUsage() + return + } + if (!shouldAutoLoadUsageOnMount.value) return const source = isAnthropicOAuthOrSetupToken.value ? 'passive' : undefined requestAutoLoad(source) }) +watch( + () => [props.batchedUsage, props.batchedUsageError, props.batchedUsageLoading, isBatchManaged.value] as const, + () => { + syncManagedUsageState() + }, + { immediate: true, deep: true } +) + +watch(isBatchManaged, (managed, wasManaged) => { + if (managed && !wasManaged) { + syncManagedUsageState() + requestParentBatchUsage() + } +}) + +watch( + () => [props.account.id, props.account.platform, props.account.type, isBatchManaged.value] as const, + ([accountID, platform, accountType, managed], [previousAccountID, previousPlatform, previousAccountType]) => { + if ( + accountID === previousAccountID && + platform === previousPlatform && + accountType === previousAccountType + ) { + return + } + if (!managed || !shouldFetchUsage.value) return + syncManagedUsageState() + requestParentBatchUsage() + }, + { flush: 'post' } +) + watch(openAIUsageRefreshKey, (nextKey, prevKey) => { if (!prevKey || nextKey === prevKey) return if (props.account.platform !== 'openai' || props.account.type !== 'oauth') return @@ -1542,6 +1694,11 @@ watch(openAIUsageRefreshKey, (nextKey, prevKey) => { return } + if (isBatchManaged.value) { + requestParentBatchUsage({ force: true }) + return + } + _usageCache.delete(props.account.id) requestAutoLoad() }) @@ -1552,6 +1709,11 @@ watch( if (nextToken === prevToken) return if (!shouldFetchUsage.value) return + if (isBatchManaged.value) { + requestParentBatchUsage({ force: true }) + return + } + const source = isAnthropicOAuthOrSetupToken.value ? 'passive' : undefined _usageCache.delete(props.account.id) loadUsage({ source, bypassCache: true }).catch((e) => { diff --git a/frontend/src/i18n/locales/en/admin/accounts.ts b/frontend/src/i18n/locales/en/admin/accounts.ts index 9060639dc0..84e57a7b37 100644 --- a/frontend/src/i18n/locales/en/admin/accounts.ts +++ b/frontend/src/i18n/locales/en/admin/accounts.ts @@ -1357,6 +1357,8 @@ export default { grokTokens: 'Tok', grokFreeQuota24hHint: 'Estimated from local token usage over the rolling 24-hour window ({limit} limit)', grokWeeklyUsage: 'Weekly {percent}%', + grokUsed: 'Used $', + grokMonthlyLimit: 'Monthly used / limit (USD from billing cents)', grokUnknown: 'Grok quota is unknown until the first upstream response includes xAI rate-limit headers.', grokRetryAfter: 'Retry after {time}', grokProbe: 'Probe', diff --git a/frontend/src/i18n/locales/en/admin/settings.ts b/frontend/src/i18n/locales/en/admin/settings.ts index fb00d93196..c2fa0f92a3 100644 --- a/frontend/src/i18n/locales/en/admin/settings.ts +++ b/frontend/src/i18n/locales/en/admin/settings.ts @@ -409,7 +409,12 @@ export default { title: 'Gateway Scheduling Settings', description: 'Control API Key scheduling behavior', allowUngroupedKey: 'Allow Ungrouped Key Scheduling', - allowUngroupedKeyHint: 'When disabled, API Keys not assigned to any group cannot make requests (403 Forbidden). Keep disabled to ensure all Keys belong to a specific group.' + allowUngroupedKeyHint: 'When disabled, API Keys not assigned to any group cannot make requests (403 Forbidden). Keep disabled to ensure all Keys belong to a specific group.', + accountSchedulingThresholdsTitle: 'Platform Account Auto-Pause Thresholds', + accountSchedulingThresholdsDescription: 'Set per-platform thresholds that automatically pause scheduling for accounts on that platform when their current native usage window reaches the configured percentage.', + accountSchedulingThresholdsGlobalHint: 'This is a system-wide global setting and applies to all accounts on that platform.', + accountSchedulingThresholdsDisabledHint: 'A value of 100 disables the auto-pause threshold for that platform.', + accountSchedulingThresholdsRangeHint: 'Range 1-100, entered as a percentage.' }, upstreamBillingProbe: { title: 'Upstream Rate Auto Detection', diff --git a/frontend/src/i18n/locales/zh/admin/accounts.ts b/frontend/src/i18n/locales/zh/admin/accounts.ts index 664e84a146..762512f993 100644 --- a/frontend/src/i18n/locales/zh/admin/accounts.ts +++ b/frontend/src/i18n/locales/zh/admin/accounts.ts @@ -401,6 +401,8 @@ export default { grokTokens: 'Token', grokFreeQuota24hHint: '按 sub2api 近 24 小时本地 Token 用量估算(上限 {limit})', grokWeeklyUsage: '周额度已用 {percent}%', + grokUsed: '已用 $', + grokMonthlyLimit: '月度已用/上限(由 billing cents 换算 USD)', grokUnknown: 'Grok 配额需等待首次上游响应返回 xAI rate-limit 头后显示。', grokRetryAfter: '{time} 后重试', grokProbe: '探测', diff --git a/frontend/src/i18n/locales/zh/admin/settings.ts b/frontend/src/i18n/locales/zh/admin/settings.ts index a23f744aaf..8e0531ee8d 100644 --- a/frontend/src/i18n/locales/zh/admin/settings.ts +++ b/frontend/src/i18n/locales/zh/admin/settings.ts @@ -402,7 +402,12 @@ export default { title: '网关调度设置', description: '控制 API Key 的调度行为', allowUngroupedKey: '允许未分组 Key 调度', - allowUngroupedKeyHint: '关闭后,未分配到任何分组的 API Key 将无法发起请求(返回 403)。建议保持关闭以确保所有 Key 都归属明确的分组。' + allowUngroupedKeyHint: '关闭后,未分配到任何分组的 API Key 将无法发起请求(返回 403)。建议保持关闭以确保所有 Key 都归属明确的分组。', + accountSchedulingThresholdsTitle: '平台账号自动停调阈值', + accountSchedulingThresholdsDescription: '按平台设置账号自动停调阈值,当账号当前原生使用窗口达到配置百分比时,自动停调该平台账号。', + accountSchedulingThresholdsGlobalHint: '这是系统级全局设置,对该平台全部账号生效。', + accountSchedulingThresholdsDisabledHint: '100 表示禁用该平台的自动停调阈值。', + accountSchedulingThresholdsRangeHint: '范围 1-100,按百分比填写。' }, upstreamBillingProbe: { title: '上游倍率自动探测', diff --git a/frontend/src/views/admin/AccountsView.vue b/frontend/src/views/admin/AccountsView.vue index 784f7c37f4..3dfe2c2797 100644 --- a/frontend/src/views/admin/AccountsView.vue +++ b/frontend/src/views/admin/AccountsView.vue @@ -318,7 +318,12 @@ :today-stats="todayStatsByAccountId[String(row.id)] ?? null" :today-stats-loading="todayStatsLoading" :manual-refresh-token="usageManualRefreshToken" + :batched-usage="usageBatchByAccountId[String(row.id)] ?? null" + :batched-usage-error="usageBatchErrorByAccountId[String(row.id)] ?? null" + :batched-usage-loading="usageBatchLoadingByAccountId[String(row.id)] === true" + :request-batched-usage="isDesktopViewport ? queueBatchedUsage : null" @account-updated="handleAccountUpdated" + @usage-loaded="handleAccountUsageLoaded(row.id, $event)" />