mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:48:43 +08:00
feat(grok): 吸收调度阈值、配额解析、批量用量与 CLI 身份
从 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。
This commit is contained in:
@@ -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"`
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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"))
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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%")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
|
||||
@@ -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])
|
||||
}
|
||||
}
|
||||
@@ -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.
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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"))
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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{}
|
||||
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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{}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -317,6 +317,19 @@ export async function getUsage(id: number, source?: 'passive' | 'active', force?
|
||||
return data
|
||||
}
|
||||
|
||||
export interface BatchAccountUsageResponse {
|
||||
usage: Record<string, AccountUsageInfo>
|
||||
errors: Record<string, string>
|
||||
}
|
||||
|
||||
export async function getBatchUsage(accountIds: number[], force?: boolean): Promise<BatchAccountUsageResponse> {
|
||||
const { data } = await apiClient.post<BatchAccountUsageResponse>('/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,
|
||||
|
||||
@@ -32,6 +32,38 @@ export type DefaultPlatformQuotasMap = Partial<Record<PlatformType, PlatformQuot
|
||||
|
||||
const PLATFORMS: PlatformType[] = ["anthropic", "openai", "gemini", "antigravity", "grok"]
|
||||
|
||||
export type SchedulingThresholdPlatformType =
|
||||
| "openai"
|
||||
| "anthropic"
|
||||
| "grok"
|
||||
|
||||
export type AccountSchedulingThresholdsMap = Record<SchedulingThresholdPlatformType, number>
|
||||
|
||||
export const SCHEDULING_THRESHOLD_PLATFORMS: SchedulingThresholdPlatformType[] = [
|
||||
"openai",
|
||||
"anthropic",
|
||||
"grok",
|
||||
]
|
||||
|
||||
export function normalizeAccountSchedulingThresholdsMap(
|
||||
input?: Partial<Record<SchedulingThresholdPlatformType, number>> | 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<Record<SchedulingThresholdPlatformType, number>> | 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;
|
||||
|
||||
@@ -393,6 +393,29 @@
|
||||
:show-now-when-idle="true"
|
||||
color="indigo"
|
||||
/>
|
||||
<UsageProgressBar
|
||||
v-if="grokMonthlyBillingBar"
|
||||
label="30d"
|
||||
:utilization="grokMonthlyBillingBar.utilization"
|
||||
:resets-at="grokMonthlyBillingBar.resetsAt"
|
||||
:show-now-when-idle="true"
|
||||
color="indigo"
|
||||
/>
|
||||
<div
|
||||
v-if="grokBillingMoneySummary"
|
||||
class="flex flex-wrap items-center gap-1 text-[10px] text-gray-500 dark:text-gray-400"
|
||||
>
|
||||
<span :title="t('admin.accounts.usageWindow.grokMonthlyLimit')">
|
||||
{{ t('admin.accounts.usageWindow.grokUsed') }}
|
||||
{{ grokBillingMoneySummary.used }}/{{ grokBillingMoneySummary.limit }}
|
||||
</span>
|
||||
<span
|
||||
v-if="grokBillingMoneySummary.usedPercent != null"
|
||||
class="rounded bg-gray-100 px-1 py-0.5 dark:bg-gray-800"
|
||||
>
|
||||
{{ grokBillingMoneySummary.usedPercent }}%
|
||||
</span>
|
||||
</div>
|
||||
<UsageProgressBar
|
||||
v-if="!grokWeeklyBillingBar && !grokIsFree && grokRequestQuotaBar"
|
||||
:label="t('admin.accounts.usageWindow.grokRequests')"
|
||||
@@ -652,16 +675,25 @@ const props = withDefaults(
|
||||
todayStats?: WindowStats | null
|
||||
todayStatsLoading?: boolean
|
||||
manualRefreshToken?: number
|
||||
batchedUsage?: AccountUsageInfo | null
|
||||
batchedUsageError?: string | null
|
||||
batchedUsageLoading?: boolean
|
||||
requestBatchedUsage?: ((account: Account, options?: { force?: boolean }) => 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<string | null>(null)
|
||||
const usageInfo = ref<AccountUsageInfo | null>(null)
|
||||
watch(usageInfo, (usage) => {
|
||||
if (usage) emit('usage-loaded', usage)
|
||||
})
|
||||
const suppressOpenAIUsageRefreshUntil = ref(0)
|
||||
const rootRef = ref<HTMLElement | null>(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) => {
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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: '探测',
|
||||
|
||||
@@ -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: '上游倍率自动探测',
|
||||
|
||||
@@ -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)"
|
||||
/>
|
||||
</template>
|
||||
<template #cell-proxy="{ row }">
|
||||
@@ -527,7 +532,7 @@ import { extractApiErrorMessage } from '@/utils/apiError'
|
||||
import { sanitizeUrl } from '@/utils/url'
|
||||
import { getFloatingPanelPosition } from '@/utils/floatingPanel'
|
||||
import { formatMultiplier } from '@/utils/formatters'
|
||||
import type { Account, AccountPlatform, AccountSchedulerGroupScore, AccountType, Proxy as AccountProxy, AdminGroup, WindowStats, ClaudeModel, UpstreamBillingProbeSnapshot } from '@/types'
|
||||
import type { Account, AccountPlatform, AccountSchedulerGroupScore, AccountType, AccountUsageInfo, Proxy as AccountProxy, AdminGroup, WindowStats, ClaudeModel, UpstreamBillingProbeSnapshot } from '@/types'
|
||||
|
||||
const { t } = useI18n()
|
||||
const appStore = useAppStore()
|
||||
@@ -692,6 +697,24 @@ const todayStatsReqSeq = ref(0)
|
||||
const pendingTodayStatsRefresh = ref(false)
|
||||
const usageManualRefreshToken = ref(0)
|
||||
|
||||
const desktopViewportQuery = '(min-width: 768px)'
|
||||
const isDesktopViewport = ref(
|
||||
typeof window === 'undefined' ? true : window.matchMedia(desktopViewportQuery).matches
|
||||
)
|
||||
let desktopViewportMediaQuery: MediaQueryList | null = null
|
||||
let desktopViewportListener: ((event: MediaQueryListEvent) => void) | null = null
|
||||
|
||||
const usageBatchByAccountId = ref<Record<string, AccountUsageInfo | null>>({})
|
||||
const usageBatchErrorByAccountId = ref<Record<string, string | null>>({})
|
||||
const usageBatchLoadingByAccountId = ref<Record<string, boolean>>({})
|
||||
const usageBatchRequestTokenByAccountId = ref<Record<string, number>>({})
|
||||
const usageBatchCache = new Map<number, { data: AccountUsageInfo; ts: number }>()
|
||||
const USAGE_BATCH_CACHE_TTL = 5 * 60 * 1000
|
||||
const pendingUsageBatchIds = new Set<number>()
|
||||
let usageBatchFlushTimer: ReturnType<typeof setTimeout> | null = null
|
||||
let queuedUsageBatchForce = false
|
||||
let usageBatchRequestToken = 0
|
||||
|
||||
const buildDefaultTodayStats = (): WindowStats => ({
|
||||
requests: 0,
|
||||
tokens: 0,
|
||||
@@ -700,6 +723,138 @@ const buildDefaultTodayStats = (): WindowStats => ({
|
||||
user_cost: 0
|
||||
})
|
||||
|
||||
const accountSupportsBatchUsage = (account: Account) => {
|
||||
if (account.platform === 'anthropic') {
|
||||
return account.type === 'oauth' || account.type === 'setup-token'
|
||||
}
|
||||
if (account.platform === 'gemini') return true
|
||||
if (account.platform === 'antigravity') return account.type === 'oauth'
|
||||
if (account.platform === 'openai') return account.type === 'oauth'
|
||||
if (account.platform === 'grok') return account.type === 'oauth'
|
||||
return false
|
||||
}
|
||||
|
||||
const setUsageBatchLoading = (accountID: number, loadingState: boolean) => {
|
||||
usageBatchLoadingByAccountId.value = {
|
||||
...usageBatchLoadingByAccountId.value,
|
||||
[String(accountID)]: loadingState
|
||||
}
|
||||
}
|
||||
|
||||
const setUsageBatchState = (accountID: number, usage: AccountUsageInfo | null, error: string | null) => {
|
||||
const key = String(accountID)
|
||||
usageBatchByAccountId.value = {
|
||||
...usageBatchByAccountId.value,
|
||||
[key]: usage
|
||||
}
|
||||
usageBatchErrorByAccountId.value = {
|
||||
...usageBatchErrorByAccountId.value,
|
||||
[key]: error
|
||||
}
|
||||
}
|
||||
|
||||
const handleAccountUsageLoaded = (accountID: number, usage: AccountUsageInfo) => {
|
||||
if (usageBatchByAccountId.value[String(accountID)] === usage) return
|
||||
setUsageBatchState(accountID, usage, null)
|
||||
}
|
||||
|
||||
const flushQueuedUsageBatch = async () => {
|
||||
usageBatchFlushTimer = null
|
||||
const accountIDs = Array.from(pendingUsageBatchIds)
|
||||
const force = queuedUsageBatchForce
|
||||
pendingUsageBatchIds.clear()
|
||||
queuedUsageBatchForce = false
|
||||
|
||||
if (accountIDs.length === 0) return
|
||||
|
||||
const requestTokensByAccount = accountIDs.reduce<Record<string, number>>((acc, accountID) => {
|
||||
acc[String(accountID)] = usageBatchRequestTokenByAccountId.value[String(accountID)] ?? 0
|
||||
return acc
|
||||
}, {})
|
||||
|
||||
try {
|
||||
const result = await adminAPI.accounts.getBatchUsage(accountIDs, force)
|
||||
|
||||
const usageMap = result.usage ?? {}
|
||||
const errorMap = result.errors ?? {}
|
||||
const now = Date.now()
|
||||
const nextUsage = { ...usageBatchByAccountId.value }
|
||||
const nextErrors = { ...usageBatchErrorByAccountId.value }
|
||||
const nextLoading = { ...usageBatchLoadingByAccountId.value }
|
||||
|
||||
for (const accountID of accountIDs) {
|
||||
const key = String(accountID)
|
||||
if ((usageBatchRequestTokenByAccountId.value[key] ?? 0) !== requestTokensByAccount[key]) {
|
||||
continue
|
||||
}
|
||||
const usage = usageMap[key] ?? null
|
||||
nextUsage[key] = usage
|
||||
nextErrors[key] = errorMap[key] ?? null
|
||||
nextLoading[key] = false
|
||||
if (usage) {
|
||||
usageBatchCache.set(accountID, { data: usage, ts: now })
|
||||
} else {
|
||||
usageBatchCache.delete(accountID)
|
||||
}
|
||||
}
|
||||
|
||||
usageBatchByAccountId.value = nextUsage
|
||||
usageBatchErrorByAccountId.value = nextErrors
|
||||
usageBatchLoadingByAccountId.value = nextLoading
|
||||
} catch (error) {
|
||||
const nextErrors = { ...usageBatchErrorByAccountId.value }
|
||||
const nextLoading = { ...usageBatchLoadingByAccountId.value }
|
||||
for (const accountID of accountIDs) {
|
||||
const key = String(accountID)
|
||||
if ((usageBatchRequestTokenByAccountId.value[key] ?? 0) !== requestTokensByAccount[key]) {
|
||||
continue
|
||||
}
|
||||
nextErrors[key] = 'Failed'
|
||||
nextLoading[key] = false
|
||||
}
|
||||
usageBatchErrorByAccountId.value = nextErrors
|
||||
usageBatchLoadingByAccountId.value = nextLoading
|
||||
console.error('Failed to load account usage batch:', error)
|
||||
}
|
||||
}
|
||||
|
||||
const queueBatchedUsage = (account: Account, options?: { force?: boolean }) => {
|
||||
if (!isDesktopViewport.value) return
|
||||
if (!accountSupportsBatchUsage(account)) return
|
||||
|
||||
const force = options?.force === true
|
||||
const cacheKey = account.id
|
||||
const key = String(cacheKey)
|
||||
|
||||
if (force) {
|
||||
usageBatchCache.delete(cacheKey)
|
||||
} else {
|
||||
const cached = usageBatchCache.get(cacheKey)
|
||||
if (cached && Date.now() - cached.ts < USAGE_BATCH_CACHE_TTL) {
|
||||
setUsageBatchState(cacheKey, cached.data, null)
|
||||
setUsageBatchLoading(cacheKey, false)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
usageBatchErrorByAccountId.value = {
|
||||
...usageBatchErrorByAccountId.value,
|
||||
[key]: null
|
||||
}
|
||||
usageBatchRequestTokenByAccountId.value = {
|
||||
...usageBatchRequestTokenByAccountId.value,
|
||||
[key]: ++usageBatchRequestToken
|
||||
}
|
||||
setUsageBatchLoading(cacheKey, true)
|
||||
pendingUsageBatchIds.add(cacheKey)
|
||||
queuedUsageBatchForce = queuedUsageBatchForce || force
|
||||
|
||||
if (usageBatchFlushTimer !== null) return
|
||||
usageBatchFlushTimer = setTimeout(() => {
|
||||
void flushQueuedUsageBatch()
|
||||
}, 0)
|
||||
}
|
||||
|
||||
const refreshTodayStatsBatch = async () => {
|
||||
// Why this checks both columns:
|
||||
// - today_stats column shows dedicated today's metrics.
|
||||
@@ -1084,6 +1239,22 @@ watch(loading, (isLoading, wasLoading) => {
|
||||
}
|
||||
})
|
||||
|
||||
watch(accounts, (rows) => {
|
||||
const visibleIDs = new Set(rows.map((row) => String(row.id)))
|
||||
usageBatchByAccountId.value = Object.fromEntries(
|
||||
Object.entries(usageBatchByAccountId.value).filter(([key]) => visibleIDs.has(key))
|
||||
)
|
||||
usageBatchErrorByAccountId.value = Object.fromEntries(
|
||||
Object.entries(usageBatchErrorByAccountId.value).filter(([key]) => visibleIDs.has(key))
|
||||
)
|
||||
usageBatchLoadingByAccountId.value = Object.fromEntries(
|
||||
Object.entries(usageBatchLoadingByAccountId.value).filter(([key]) => visibleIDs.has(key))
|
||||
)
|
||||
usageBatchRequestTokenByAccountId.value = Object.fromEntries(
|
||||
Object.entries(usageBatchRequestTokenByAccountId.value).filter(([key]) => visibleIDs.has(key))
|
||||
)
|
||||
})
|
||||
|
||||
watch(upstreamBillingNow, () => {
|
||||
if (sortState.sort_by !== 'upstream_billing_rate' || loading.value) return
|
||||
if (typeof document !== 'undefined' && document.hidden) return
|
||||
@@ -2183,6 +2354,19 @@ const handleClickOutside = (event: MouseEvent) => {
|
||||
}
|
||||
|
||||
onMounted(async () => {
|
||||
if (typeof window !== 'undefined') {
|
||||
desktopViewportMediaQuery = window.matchMedia(desktopViewportQuery)
|
||||
isDesktopViewport.value = desktopViewportMediaQuery.matches
|
||||
desktopViewportListener = (event: MediaQueryListEvent) => {
|
||||
isDesktopViewport.value = event.matches
|
||||
}
|
||||
if (typeof desktopViewportMediaQuery.addEventListener === 'function') {
|
||||
desktopViewportMediaQuery.addEventListener('change', desktopViewportListener)
|
||||
} else {
|
||||
desktopViewportMediaQuery.addListener(desktopViewportListener)
|
||||
}
|
||||
}
|
||||
|
||||
load()
|
||||
loadUpstreamBillingProbeGlobalState()
|
||||
try {
|
||||
@@ -2205,9 +2389,23 @@ onMounted(async () => {
|
||||
})
|
||||
|
||||
onUnmounted(() => {
|
||||
if (usageBatchFlushTimer !== null) {
|
||||
clearTimeout(usageBatchFlushTimer)
|
||||
usageBatchFlushTimer = null
|
||||
}
|
||||
pendingUsageBatchIds.clear()
|
||||
window.removeEventListener('scroll', handleScroll, true)
|
||||
window.removeEventListener('resize', handleViewportResize)
|
||||
document.removeEventListener('click', handleClickOutside)
|
||||
if (desktopViewportMediaQuery && desktopViewportListener) {
|
||||
if (typeof desktopViewportMediaQuery.removeEventListener === 'function') {
|
||||
desktopViewportMediaQuery.removeEventListener('change', desktopViewportListener)
|
||||
} else {
|
||||
desktopViewportMediaQuery.removeListener(desktopViewportListener)
|
||||
}
|
||||
}
|
||||
desktopViewportListener = null
|
||||
desktopViewportMediaQuery = null
|
||||
})
|
||||
</script>
|
||||
|
||||
|
||||
@@ -4871,6 +4871,78 @@
|
||||
<Toggle v-model="form.allow_ungrouped_key_scheduling" />
|
||||
</div>
|
||||
|
||||
<div class="border-t border-gray-100 pt-4 dark:border-dark-700">
|
||||
<div class="mb-3">
|
||||
<label class="font-medium text-gray-900 dark:text-white">
|
||||
{{
|
||||
t(
|
||||
"admin.settings.scheduling.accountSchedulingThresholdsTitle",
|
||||
)
|
||||
}}
|
||||
</label>
|
||||
<p class="mt-1 text-sm text-gray-500 dark:text-gray-400">
|
||||
{{
|
||||
t(
|
||||
"admin.settings.scheduling.accountSchedulingThresholdsDescription",
|
||||
)
|
||||
}}
|
||||
</p>
|
||||
<p class="mt-0.5 text-xs text-gray-500 dark:text-gray-400">
|
||||
{{
|
||||
t(
|
||||
"admin.settings.scheduling.accountSchedulingThresholdsGlobalHint",
|
||||
)
|
||||
}}
|
||||
</p>
|
||||
<p class="mt-0.5 text-xs text-amber-600 dark:text-amber-400">
|
||||
{{
|
||||
t(
|
||||
"admin.settings.scheduling.accountSchedulingThresholdsDisabledHint",
|
||||
)
|
||||
}}
|
||||
</p>
|
||||
</div>
|
||||
<div class="grid grid-cols-1 gap-4 sm:grid-cols-2 xl:grid-cols-3">
|
||||
<div
|
||||
v-for="platform in schedulingThresholdPlatforms"
|
||||
:key="platform"
|
||||
class="rounded-lg border border-gray-200 p-4 dark:border-dark-700"
|
||||
>
|
||||
<div class="flex items-start justify-between gap-3">
|
||||
<div>
|
||||
<label
|
||||
class="font-mono text-sm font-medium text-gray-900 dark:text-white"
|
||||
>
|
||||
{{ platform }}
|
||||
</label>
|
||||
<p class="mt-0.5 text-xs text-gray-500 dark:text-gray-400">
|
||||
{{
|
||||
t(
|
||||
"admin.settings.scheduling.accountSchedulingThresholdsRangeHint",
|
||||
)
|
||||
}}
|
||||
</p>
|
||||
</div>
|
||||
<span
|
||||
class="rounded bg-gray-100 px-2 py-0.5 text-[11px] font-medium text-gray-600 dark:bg-dark-700 dark:text-gray-300"
|
||||
>
|
||||
%
|
||||
</span>
|
||||
</div>
|
||||
<input
|
||||
v-model.number="form.account_scheduling_thresholds[platform]"
|
||||
type="number"
|
||||
min="1"
|
||||
max="100"
|
||||
step="1"
|
||||
class="input mt-3"
|
||||
:data-testid="`account-scheduling-threshold-${platform}`"
|
||||
placeholder="100"
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div
|
||||
v-if="!form.openai_advanced_scheduler_enabled"
|
||||
class="flex items-center justify-between border-t border-gray-100 pt-5 dark:border-dark-700"
|
||||
@@ -8491,8 +8563,11 @@ import { adminAPI } from "@/api";
|
||||
import {
|
||||
appendAuthSourceDefaultsToUpdateRequest,
|
||||
buildAuthSourceDefaultsState,
|
||||
normalizeAccountSchedulingThresholdsMap,
|
||||
normalizePlatformQuotasMap,
|
||||
sanitizeAccountSchedulingThresholdsMap,
|
||||
sanitizePlatformQuotasMap,
|
||||
SCHEDULING_THRESHOLD_PLATFORMS,
|
||||
defaultWeChatConnectScopesForMode,
|
||||
deriveWeChatConnectStoredMode,
|
||||
normalizeDefaultSubscriptionSettings,
|
||||
@@ -9236,8 +9311,11 @@ type SettingsForm = Omit<
|
||||
openai_advanced_scheduler_weight_session_sticky: string;
|
||||
// 系统全局平台限额 map;form 内始终归一化为全 4 平台对象(模板非空绑定依赖此不变量)
|
||||
default_platform_quotas: DefaultPlatformQuotasMap;
|
||||
account_scheduling_thresholds: ReturnType<typeof normalizeAccountSchedulingThresholdsMap>;
|
||||
};
|
||||
|
||||
const schedulingThresholdPlatforms = SCHEDULING_THRESHOLD_PLATFORMS;
|
||||
|
||||
const form = reactive<SettingsForm>({
|
||||
registration_enabled: true,
|
||||
email_verify_enabled: false,
|
||||
@@ -9260,6 +9338,7 @@ const form = reactive<SettingsForm>({
|
||||
login_agreement_documents: defaultLoginAgreementDocuments(),
|
||||
default_balance: 0,
|
||||
default_platform_quotas: normalizePlatformQuotasMap() as DefaultPlatformQuotasMap,
|
||||
account_scheduling_thresholds: normalizeAccountSchedulingThresholdsMap(),
|
||||
affiliate_rebate_rate: 20,
|
||||
affiliate_rebate_freeze_hours: 0,
|
||||
affiliate_rebate_duration_days: 0,
|
||||
@@ -10512,6 +10591,9 @@ async function loadSettings() {
|
||||
: defaultLoginAgreementDocuments();
|
||||
Object.assign(authSourceDefaults, buildAuthSourceDefaultsState(settings));
|
||||
form.default_platform_quotas = normalizePlatformQuotasMap(settings.default_platform_quotas);
|
||||
form.account_scheduling_thresholds = normalizeAccountSchedulingThresholdsMap(
|
||||
settings.account_scheduling_thresholds,
|
||||
);
|
||||
form.backend_mode_enabled = settings.backend_mode_enabled;
|
||||
form.default_subscriptions = normalizeDefaultSubscriptionSettings(
|
||||
settings.default_subscriptions,
|
||||
@@ -11186,6 +11268,9 @@ async function saveSettings() {
|
||||
}
|
||||
|
||||
payload.default_platform_quotas = sanitizePlatformQuotasMap(form.default_platform_quotas);
|
||||
payload.account_scheduling_thresholds = sanitizeAccountSchedulingThresholdsMap(
|
||||
form.account_scheduling_thresholds,
|
||||
);
|
||||
appendAuthSourceDefaultsToUpdateRequest(payload, authSourceDefaults);
|
||||
|
||||
const updated = await settingsStepUp.run(() =>
|
||||
@@ -11199,6 +11284,9 @@ async function saveSettings() {
|
||||
}
|
||||
Object.assign(authSourceDefaults, buildAuthSourceDefaultsState(updated));
|
||||
form.default_platform_quotas = normalizePlatformQuotasMap(updated.default_platform_quotas);
|
||||
form.account_scheduling_thresholds = normalizeAccountSchedulingThresholdsMap(
|
||||
updated.account_scheduling_thresholds,
|
||||
);
|
||||
registrationEmailSuffixWhitelistTags.value =
|
||||
normalizeRegistrationEmailSuffixDomains(
|
||||
updated.registration_email_suffix_whitelist,
|
||||
|
||||
Reference in New Issue
Block a user