diff --git a/README.md b/README.md index c2715eae0c..74ab9af258 100644 --- a/README.md +++ b/README.md @@ -91,6 +91,11 @@ Sub2API is an AI API gateway platform designed to distribute and manage API quot Thanks to AIGoCode for sponsoring this project! AIGoCode is an all-in-one platform that integrates Claude Code, Codex, and the latest Gemini models, providing you with stable, efficient, and highly cost-effective AI coding services. The platform offers flexible subscription plans, zero risk of account suspension, direct access with no VPN required, and lightning-fast responses. AIGoCode has prepared a special benefit for sub2api users: if you register via this link, you'll receive an extra 10% bonus credit on your first top-up! + +bmoplus +Huge thanks to BmoPlus for sponsoring this project! BmoPlus is a highly reliable AI account provider built strictly for heavy AI users and developers. They offer rock-solid, ready-to-use accounts and official top-up services for ChatGPT Plus / ChatGPT Pro (Full Warranty) / Claude Pro / Super Grok / Gemini Pro. By registering and ordering through BmoPlus - Premium AI Accounts & Top-ups, users can unlock the mind-blowing rate of 10% of the official GPT subscription price (90% OFF) + + ## Ecosystem diff --git a/README_CN.md b/README_CN.md index 0ace1f77a6..c701372c43 100644 --- a/README_CN.md +++ b/README_CN.md @@ -90,6 +90,11 @@ Sub2API 是一个 AI API 网关平台,用于分发和管理 AI 产品订阅的 感谢 AIGoCode 赞助了本项目!AIGoCode 是一站式集成 Claude Code、Codex 以及最新 Gemini 模型的综合平台,为您提供稳定、高效、高性价比的 AI 编程服务。平台提供灵活的订阅方案,零封号风险,免 VPN 直连,响应极速。AIGoCode 为 sub2api 用户准备了专属福利:通过此链接注册,首次充值可额外获得 10% 赠送额度! + +bmoplus +感谢 BmoPlus 赞助了本项目!BmoPlus 是一家专为AI订阅重度用户打造的可靠 AI 账号代充服务商,提供稳定的 ChatGPT Plus / ChatGPT Pro(全程质保) / Claude Pro / Super Grok / Gemini Pro 的官方代充&成品账号。 通过BmoPlus AI成品号专卖/代充注册下单的用户,可享GPT 官网订阅一折 的震撼价格! + + ## 生态项目 diff --git a/README_JA.md b/README_JA.md index d74ca9cebf..0d4db616f0 100644 --- a/README_JA.md +++ b/README_JA.md @@ -90,6 +90,11 @@ Sub2API は、AI 製品のサブスクリプションから API クォータを AIGoCode のご支援に感謝します!AIGoCode は Claude Code、Codex、最新の Gemini モデルを統合したオールインワンプラットフォームで、安定的かつ効率的でコストパフォーマンスに優れた AI コーディングサービスを提供します。柔軟なサブスクリプションプラン、アカウント停止リスクゼロ、VPN 不要の直接アクセス、超高速レスポンスが特長です。AIGoCode は sub2api ユーザー向けに特別特典を用意しています:こちらのリンクから登録すると、初回チャージ時に 10% のボーナスクレジットを追加プレゼント! + +bmoplus +本プロジェクトにご支援いただいた BmoPlus に感謝いたします!BmoPlusは、AIサブスクリプションのヘビーユーザー向けに特化した信頼性の高いAIアカウントサービスプロバイダーであり、安定した ChatGPT Plus / ChatGPT Pro (完全保証) / Claude Pro / Super Grok / Gemini Pro の公式代行チャージおよび即納アカウントを提供しています。こちらのBmoPlus AIアカウント専門店/代行チャージ経由でご登録・ご注文いただいたユーザー様は、GPTを 公式サイト価格の約1割(90% OFF) という驚異的な価格でご利用いただけます! + + ## エコシステム diff --git a/assets/partners/logos/bmoplus.jpg b/assets/partners/logos/bmoplus.jpg new file mode 100644 index 0000000000..1a9b4d8b7b Binary files /dev/null and b/assets/partners/logos/bmoplus.jpg differ diff --git a/backend/internal/domain/constants.go b/backend/internal/domain/constants.go index 429486c3bf..a57f7067dc 100644 --- a/backend/internal/domain/constants.go +++ b/backend/internal/domain/constants.go @@ -71,6 +71,7 @@ const ( // 与前端 useModelWhitelist.ts 中的 antigravityDefaultMappings 保持一致 var DefaultAntigravityModelMapping = map[string]string{ // Claude 白名单 + "claude-opus-4-7": "claude-opus-4-7", // 官方模型 "claude-opus-4-6-thinking": "claude-opus-4-6-thinking", // 官方模型 "claude-opus-4-6": "claude-opus-4-6-thinking", // 简称映射 "claude-opus-4-5-thinking": "claude-opus-4-6-thinking", // 迁移旧模型 @@ -120,6 +121,7 @@ var DefaultAntigravityModelMapping = map[string]string{ // aws_region 自动调整为匹配的区域前缀(如 eu.、apac.、jp. 等) var DefaultBedrockModelMapping = map[string]string{ // Claude Opus + "claude-opus-4-7": "us.anthropic.claude-opus-4-7-v1", "claude-opus-4-6-thinking": "us.anthropic.claude-opus-4-6-v1", "claude-opus-4-6": "us.anthropic.claude-opus-4-6-v1", "claude-opus-4-5-thinking": "us.anthropic.claude-opus-4-5-20251101-v1:0", diff --git a/backend/internal/pkg/antigravity/claude_types.go b/backend/internal/pkg/antigravity/claude_types.go index ce144bb911..0b8ae5f2bc 100644 --- a/backend/internal/pkg/antigravity/claude_types.go +++ b/backend/internal/pkg/antigravity/claude_types.go @@ -154,6 +154,7 @@ var claudeModels = []modelDef{ {ID: "claude-sonnet-4-5-thinking", DisplayName: "Claude Sonnet 4.5 Thinking", CreatedAt: "2025-09-29T00:00:00Z"}, {ID: "claude-opus-4-6", DisplayName: "Claude Opus 4.6", CreatedAt: "2026-02-05T00:00:00Z"}, {ID: "claude-opus-4-6-thinking", DisplayName: "Claude Opus 4.6 Thinking", CreatedAt: "2026-02-05T00:00:00Z"}, + {ID: "claude-opus-4-7", DisplayName: "Claude Opus 4.7", CreatedAt: "2026-04-17T00:00:00Z"}, {ID: "claude-sonnet-4-6", DisplayName: "Claude Sonnet 4.6", CreatedAt: "2026-02-17T00:00:00Z"}, } diff --git a/backend/internal/pkg/antigravity/request_transformer.go b/backend/internal/pkg/antigravity/request_transformer.go index d13a84983f..b5de8166ce 100644 --- a/backend/internal/pkg/antigravity/request_transformer.go +++ b/backend/internal/pkg/antigravity/request_transformer.go @@ -582,8 +582,12 @@ func maxOutputTokensLimit(model string) int { return maxOutputTokensUpperBound } -func isAntigravityOpus46Model(model string) bool { - return strings.HasPrefix(strings.ToLower(model), "claude-opus-4-6") +// isAntigravityOpusHighTierModel 判断是否为高阶 Opus 模型(4.6+), +// 用于 adaptive thinking 时覆写为高预算。 +func isAntigravityOpusHighTierModel(model string) bool { + lower := strings.ToLower(model) + return strings.HasPrefix(lower, "claude-opus-4-6") || + strings.HasPrefix(lower, "claude-opus-4-7") } func buildGenerationConfig(req *ClaudeRequest) *GeminiGenerationConfig { @@ -605,12 +609,12 @@ func buildGenerationConfig(req *ClaudeRequest) *GeminiGenerationConfig { } // - thinking.type=enabled:budget_tokens>0 用显式预算 - // - thinking.type=adaptive:仅在 Antigravity 的 Opus 4.6 上覆写为 (24576) + // - thinking.type=adaptive:在 Antigravity 的高阶 Opus(4.6+)上覆写为 (24576) budget := -1 if req.Thinking.BudgetTokens > 0 { budget = req.Thinking.BudgetTokens } - if req.Thinking.Type == "adaptive" && isAntigravityOpus46Model(req.Model) { + if req.Thinking.Type == "adaptive" && isAntigravityOpusHighTierModel(req.Model) { budget = ClaudeAdaptiveHighThinkingBudgetTokens } diff --git a/backend/internal/pkg/claude/constants.go b/backend/internal/pkg/claude/constants.go index dfca252f48..21c723d208 100644 --- a/backend/internal/pkg/claude/constants.go +++ b/backend/internal/pkg/claude/constants.go @@ -83,6 +83,12 @@ var DefaultModels = []Model{ DisplayName: "Claude Opus 4.6", CreatedAt: "2026-02-06T00:00:00Z", }, + { + ID: "claude-opus-4-7", + Type: "model", + DisplayName: "Claude Opus 4.7", + CreatedAt: "2026-04-17T00:00:00Z", + }, { ID: "claude-sonnet-4-6", Type: "model", diff --git a/backend/internal/repository/scheduler_cache.go b/backend/internal/repository/scheduler_cache.go index e9be8c7a6e..add0e50175 100644 --- a/backend/internal/repository/scheduler_cache.go +++ b/backend/internal/repository/scheduler_cache.go @@ -426,6 +426,13 @@ func filterSchedulerExtra(extra map[string]any) map[string]any { "window_cost_sticky_reserve", "max_sessions", "session_idle_timeout_minutes", + "openai_oauth_responses_websockets_v2_enabled", + "openai_oauth_responses_websockets_v2_mode", + "openai_apikey_responses_websockets_v2_enabled", + "openai_apikey_responses_websockets_v2_mode", + "responses_websockets_v2_enabled", + "openai_ws_enabled", + "openai_ws_force_http", } filtered := make(map[string]any) for _, key := range keys { diff --git a/backend/internal/repository/scheduler_cache_unit_test.go b/backend/internal/repository/scheduler_cache_unit_test.go new file mode 100644 index 0000000000..bcfd0e7aba --- /dev/null +++ b/backend/internal/repository/scheduler_cache_unit_test.go @@ -0,0 +1,33 @@ +//go:build unit + +package repository + +import ( + "testing" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestBuildSchedulerMetadataAccount_KeepsOpenAIWSFlags(t *testing.T) { + account := service.Account{ + ID: 42, + Platform: service.PlatformOpenAI, + Type: service.AccountTypeOAuth, + Extra: map[string]any{ + "openai_oauth_responses_websockets_v2_enabled": true, + "openai_oauth_responses_websockets_v2_mode": service.OpenAIWSIngressModePassthrough, + "openai_ws_force_http": true, + "mixed_scheduling": true, + "unused_large_field": "drop-me", + }, + } + + got := buildSchedulerMetadataAccount(account) + + require.Equal(t, true, got.Extra["openai_oauth_responses_websockets_v2_enabled"]) + require.Equal(t, service.OpenAIWSIngressModePassthrough, got.Extra["openai_oauth_responses_websockets_v2_mode"]) + require.Equal(t, true, got.Extra["openai_ws_force_http"]) + require.Equal(t, true, got.Extra["mixed_scheduling"]) + require.Nil(t, got.Extra["unused_large_field"]) +} diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index 68a2fd0b9a..4d95eca483 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -63,6 +63,11 @@ func TestAPIContracts(t *testing.T) { "allowed_groups": null, "created_at": "2025-01-02T03:04:05Z", "updated_at": "2025-01-02T03:04:05Z", + "balance_notify_enabled": false, + "balance_notify_threshold_type": "", + "balance_notify_threshold": null, + "balance_notify_extra_emails": null, + "total_recharged": 0, "run_mode": "standard" } }`, diff --git a/backend/internal/service/account_test_service.go b/backend/internal/service/account_test_service.go index 55865945c6..a5559b7de2 100644 --- a/backend/internal/service/account_test_service.go +++ b/backend/internal/service/account_test_service.go @@ -515,22 +515,10 @@ func (s *AccountTestService) testOpenAIAccountConnection(c *gin.Context, account _ = s.accountRepo.UpdateExtra(ctx, account.ID, updates) mergeAccountExtra(account, updates) } - if snapshot := ParseCodexRateLimitHeaders(resp.Header); snapshot != nil { - if resetAt := codexRateLimitResetAtFromSnapshot(snapshot, time.Now()); resetAt != nil { - _ = s.accountRepo.SetRateLimited(ctx, account.ID, *resetAt) - account.RateLimitResetAt = resetAt - } - } } if resp.StatusCode != http.StatusOK { body, _ := io.ReadAll(resp.Body) - if isOAuth && s.accountRepo != nil { - if resetAt := (&RateLimitService{}).calculateOpenAI429ResetTime(resp.Header); resetAt != nil { - _ = s.accountRepo.SetRateLimited(ctx, account.ID, *resetAt) - account.RateLimitResetAt = resetAt - } - } // 401 Unauthorized: 标记账号为永久错误 if resp.StatusCode == http.StatusUnauthorized && s.accountRepo != nil { errMsg := fmt.Sprintf("Authentication failed (401): %s", string(body)) diff --git a/backend/internal/service/account_test_service_openai_test.go b/backend/internal/service/account_test_service_openai_test.go index 5125db5ba5..8260697991 100644 --- a/backend/internal/service/account_test_service_openai_test.go +++ b/backend/internal/service/account_test_service_openai_test.go @@ -111,7 +111,7 @@ func TestAccountTestService_OpenAISuccessPersistsSnapshotFromHeaders(t *testing. require.Contains(t, recorder.Body.String(), "test_complete") } -func TestAccountTestService_OpenAI429PersistsSnapshotAndRateLimit(t *testing.T) { +func TestAccountTestService_OpenAI429PersistsSnapshotWithoutRateLimit(t *testing.T) { gin.SetMode(gin.TestMode) ctx, _ := newTestContext() @@ -138,10 +138,7 @@ func TestAccountTestService_OpenAI429PersistsSnapshotAndRateLimit(t *testing.T) require.Error(t, err) require.NotEmpty(t, repo.updatedExtra) require.Equal(t, 100.0, repo.updatedExtra["codex_5h_used_percent"]) - require.Equal(t, int64(88), repo.rateLimitedID) - require.NotNil(t, repo.rateLimitedAt) - require.NotNil(t, account.RateLimitResetAt) - if account.RateLimitResetAt != nil && repo.rateLimitedAt != nil { - require.WithinDuration(t, *repo.rateLimitedAt, *account.RateLimitResetAt, time.Second) - } + require.Zero(t, repo.rateLimitedID) + require.Nil(t, repo.rateLimitedAt) + require.Nil(t, account.RateLimitResetAt) } diff --git a/backend/internal/service/account_usage_service.go b/backend/internal/service/account_usage_service.go index 0e5741d818..8d5bcec8d7 100644 --- a/backend/internal/service/account_usage_service.go +++ b/backend/internal/service/account_usage_service.go @@ -499,7 +499,6 @@ func (s *AccountUsageService) getOpenAIUsage(ctx context.Context, account *Accou if account == nil { return usage, nil } - syncOpenAICodexRateLimitFromExtra(ctx, s.accountRepo, account, now) if progress := buildCodexUsageProgressFromExtra(account.Extra, "5h", now); progress != nil { usage.FiveHour = progress @@ -509,11 +508,8 @@ func (s *AccountUsageService) getOpenAIUsage(ctx context.Context, account *Accou } if shouldRefreshOpenAICodexSnapshot(account, usage, now) && s.shouldProbeOpenAICodexSnapshot(account.ID, now) { - if updates, resetAt, err := s.probeOpenAICodexSnapshot(ctx, account); err == nil && (len(updates) > 0 || resetAt != nil) { + if updates, err := s.probeOpenAICodexSnapshot(ctx, account); err == nil && len(updates) > 0 { mergeAccountExtra(account, updates) - if resetAt != nil { - account.RateLimitResetAt = resetAt - } if usage.UpdatedAt == nil { usage.UpdatedAt = &now } @@ -594,26 +590,26 @@ func (s *AccountUsageService) shouldProbeOpenAICodexSnapshot(accountID int64, no return true } -func (s *AccountUsageService) probeOpenAICodexSnapshot(ctx context.Context, account *Account) (map[string]any, *time.Time, error) { +func (s *AccountUsageService) probeOpenAICodexSnapshot(ctx context.Context, account *Account) (map[string]any, error) { if account == nil || !account.IsOAuth() { - return nil, nil, nil + return nil, nil } accessToken := account.GetOpenAIAccessToken() if accessToken == "" { - return nil, nil, fmt.Errorf("no access token available") + return nil, fmt.Errorf("no access token available") } modelID := openaipkg.DefaultTestModel payload := createOpenAITestPayload(modelID, true) payloadBytes, err := json.Marshal(payload) if err != nil { - return nil, nil, fmt.Errorf("marshal openai probe payload: %w", err) + return nil, fmt.Errorf("marshal openai probe payload: %w", err) } reqCtx, cancel := context.WithTimeout(ctx, 15*time.Second) defer cancel() req, err := http.NewRequestWithContext(reqCtx, http.MethodPost, chatgptCodexURL, bytes.NewReader(payloadBytes)) if err != nil { - return nil, nil, fmt.Errorf("create openai probe request: %w", err) + return nil, fmt.Errorf("create openai probe request: %w", err) } req.Host = "chatgpt.com" req.Header.Set("Content-Type", "application/json") @@ -642,67 +638,51 @@ func (s *AccountUsageService) probeOpenAICodexSnapshot(ctx context.Context, acco ResponseHeaderTimeout: 10 * time.Second, }) if err != nil { - return nil, nil, fmt.Errorf("build openai probe client: %w", err) + return nil, fmt.Errorf("build openai probe client: %w", err) } resp, err := client.Do(req) if err != nil { - return nil, nil, fmt.Errorf("openai codex probe request failed: %w", err) + return nil, fmt.Errorf("openai codex probe request failed: %w", err) } defer func() { _ = resp.Body.Close() }() - updates, resetAt, err := extractOpenAICodexProbeSnapshot(resp) + updates, err := extractOpenAICodexProbeUpdates(resp) if err != nil { - return nil, nil, err + return nil, err } - if len(updates) > 0 || resetAt != nil { - s.persistOpenAICodexProbeSnapshot(account.ID, updates, resetAt) - return updates, resetAt, nil + if len(updates) > 0 { + s.persistOpenAICodexProbeSnapshot(account.ID, updates) + return updates, nil } - return nil, nil, nil + return nil, nil } -func (s *AccountUsageService) persistOpenAICodexProbeSnapshot(accountID int64, updates map[string]any, resetAt *time.Time) { +func (s *AccountUsageService) persistOpenAICodexProbeSnapshot(accountID int64, updates map[string]any) { if s == nil || s.accountRepo == nil || accountID <= 0 { return } - if len(updates) == 0 && resetAt == nil { + if len(updates) == 0 { return } go func() { updateCtx, updateCancel := context.WithTimeout(context.Background(), 5*time.Second) defer updateCancel() - if len(updates) > 0 { - _ = s.accountRepo.UpdateExtra(updateCtx, accountID, updates) - } - if resetAt != nil { - _ = s.accountRepo.SetRateLimited(updateCtx, accountID, *resetAt) - } + _ = s.accountRepo.UpdateExtra(updateCtx, accountID, updates) }() } -func extractOpenAICodexProbeSnapshot(resp *http.Response) (map[string]any, *time.Time, error) { +func extractOpenAICodexProbeUpdates(resp *http.Response) (map[string]any, error) { if resp == nil { - return nil, nil, nil + return nil, nil } if snapshot := ParseCodexRateLimitHeaders(resp.Header); snapshot != nil { - baseTime := time.Now() - updates := buildCodexUsageExtraUpdates(snapshot, baseTime) - resetAt := codexRateLimitResetAtFromSnapshot(snapshot, baseTime) - if len(updates) > 0 { - return updates, resetAt, nil - } - return nil, resetAt, nil + return buildCodexUsageExtraUpdates(snapshot, time.Now()), nil } if resp.StatusCode < 200 || resp.StatusCode >= 300 { - return nil, nil, fmt.Errorf("openai codex probe returned status %d", resp.StatusCode) + return nil, fmt.Errorf("openai codex probe returned status %d", resp.StatusCode) } - return nil, nil, nil -} - -func extractOpenAICodexProbeUpdates(resp *http.Response) (map[string]any, error) { - updates, _, err := extractOpenAICodexProbeSnapshot(resp) - return updates, err + return nil, nil } func mergeAccountExtra(account *Account, updates map[string]any) { diff --git a/backend/internal/service/account_usage_service_test.go b/backend/internal/service/account_usage_service_test.go index fe2552251d..28b49838a5 100644 --- a/backend/internal/service/account_usage_service_test.go +++ b/backend/internal/service/account_usage_service_test.go @@ -92,30 +92,7 @@ func TestExtractOpenAICodexProbeUpdatesAccepts429WithCodexHeaders(t *testing.T) } } -func TestExtractOpenAICodexProbeSnapshotAccepts429WithResetAt(t *testing.T) { - t.Parallel() - - headers := make(http.Header) - headers.Set("x-codex-primary-used-percent", "100") - headers.Set("x-codex-primary-reset-after-seconds", "604800") - headers.Set("x-codex-primary-window-minutes", "10080") - headers.Set("x-codex-secondary-used-percent", "100") - headers.Set("x-codex-secondary-reset-after-seconds", "18000") - headers.Set("x-codex-secondary-window-minutes", "300") - - updates, resetAt, err := extractOpenAICodexProbeSnapshot(&http.Response{StatusCode: http.StatusTooManyRequests, Header: headers}) - if err != nil { - t.Fatalf("extractOpenAICodexProbeSnapshot() error = %v", err) - } - if len(updates) == 0 { - t.Fatal("expected codex probe updates from 429 headers") - } - if resetAt == nil { - t.Fatal("expected resetAt from exhausted codex headers") - } -} - -func TestAccountUsageService_PersistOpenAICodexProbeSnapshotSetsRateLimit(t *testing.T) { +func TestAccountUsageService_PersistOpenAICodexProbeSnapshotOnlyUpdatesExtra(t *testing.T) { t.Parallel() repo := &accountUsageCodexProbeRepo{ @@ -123,12 +100,10 @@ func TestAccountUsageService_PersistOpenAICodexProbeSnapshotSetsRateLimit(t *tes rateLimitCh: make(chan time.Time, 1), } svc := &AccountUsageService{accountRepo: repo} - resetAt := time.Now().Add(2 * time.Hour).UTC().Truncate(time.Second) - svc.persistOpenAICodexProbeSnapshot(321, map[string]any{ "codex_7d_used_percent": 100.0, - "codex_7d_reset_at": resetAt.Format(time.RFC3339), - }, &resetAt) + "codex_7d_reset_at": time.Now().Add(2 * time.Hour).UTC().Truncate(time.Second).Format(time.RFC3339), + }) select { case updates := <-repo.updateExtraCh: @@ -136,16 +111,49 @@ func TestAccountUsageService_PersistOpenAICodexProbeSnapshotSetsRateLimit(t *tes t.Fatalf("codex_7d_used_percent = %v, want 100", got) } case <-time.After(2 * time.Second): - t.Fatal("waiting for codex probe extra persistence timed out") + t.Fatal("等待 codex 探测快照写入 extra 超时") } select { case got := <-repo.rateLimitCh: - if got.Before(resetAt.Add(-time.Second)) || got.After(resetAt.Add(time.Second)) { - t.Fatalf("rate limit resetAt = %v, want around %v", got, resetAt) - } - case <-time.After(2 * time.Second): - t.Fatal("waiting for codex probe rate limit persistence timed out") + t.Fatalf("不应将探测快照写入运行时限流状态: %v", got) + case <-time.After(200 * time.Millisecond): + } +} + +func TestAccountUsageService_GetOpenAIUsage_DoesNotPromoteCodexExtraToRateLimit(t *testing.T) { + t.Parallel() + + resetAt := time.Now().Add(6 * 24 * time.Hour).UTC().Truncate(time.Second) + repo := &accountUsageCodexProbeRepo{ + rateLimitCh: make(chan time.Time, 1), + } + svc := &AccountUsageService{accountRepo: repo} + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{ + "codex_5h_used_percent": 1.0, + "codex_5h_reset_at": time.Now().Add(2 * time.Hour).UTC().Truncate(time.Second).Format(time.RFC3339), + "codex_7d_used_percent": 100.0, + "codex_7d_reset_at": resetAt.Format(time.RFC3339), + }, + } + + usage, err := svc.getOpenAIUsage(context.Background(), account) + if err != nil { + t.Fatalf("getOpenAIUsage() error = %v", err) + } + if usage.SevenDay == nil || usage.SevenDay.Utilization != 100.0 { + t.Fatalf("预期 7 天用量仍然可见,实际为 %#v", usage.SevenDay) + } + if account.RateLimitResetAt != nil { + t.Fatalf("不应让已耗尽的 codex extra 改写运行时限流状态: %v", account.RateLimitResetAt) + } + select { + case got := <-repo.rateLimitCh: + t.Fatalf("不应将已耗尽的 codex extra 持久化为运行时限流状态: %v", got) + case <-time.After(200 * time.Millisecond): } } diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index 864725256f..701f3659ee 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -1491,10 +1491,6 @@ func (s *adminServiceImpl) ListAccounts(ctx context.Context, page, pageSize int, if err != nil { return nil, 0, err } - now := time.Now() - for i := range accounts { - syncOpenAICodexRateLimitFromExtra(ctx, s.accountRepo, &accounts[i], now) - } return accounts, result.Total, nil } diff --git a/backend/internal/service/admin_service_apikey_test.go b/backend/internal/service/admin_service_apikey_test.go index f9fd67423e..419ddbc329 100644 --- a/backend/internal/service/admin_service_apikey_test.go +++ b/backend/internal/service/admin_service_apikey_test.go @@ -65,14 +65,14 @@ func (s *userRepoStubForGroupUpdate) ExistsByEmail(context.Context, string) (boo func (s *userRepoStubForGroupUpdate) RemoveGroupFromAllowedGroups(context.Context, int64) (int64, error) { panic("unexpected") } -func (s *userRepoStubForGroupUpdate) RemoveGroupFromUserAllowedGroups(context.Context, int64, int64) error { - panic("unexpected") -} func (s *userRepoStubForGroupUpdate) UpdateTotpSecret(context.Context, int64, *string) error { panic("unexpected") } func (s *userRepoStubForGroupUpdate) EnableTotp(context.Context, int64) error { panic("unexpected") } func (s *userRepoStubForGroupUpdate) DisableTotp(context.Context, int64) error { panic("unexpected") } +func (s *userRepoStubForGroupUpdate) RemoveGroupFromUserAllowedGroups(context.Context, int64, int64) error { + panic("unexpected") +} // apiKeyRepoStubForGroupUpdate implements APIKeyRepository for AdminUpdateAPIKeyGroupID tests. type apiKeyRepoStubForGroupUpdate struct { @@ -131,9 +131,6 @@ func (s *apiKeyRepoStubForGroupUpdate) SearchAPIKeys(context.Context, int64, str func (s *apiKeyRepoStubForGroupUpdate) ClearGroupIDByGroupID(context.Context, int64) (int64, error) { panic("unexpected") } -func (s *apiKeyRepoStubForGroupUpdate) UpdateGroupIDByUserAndGroup(context.Context, int64, int64, int64) (int64, error) { - panic("unexpected") -} func (s *apiKeyRepoStubForGroupUpdate) CountByGroupID(context.Context, int64) (int64, error) { panic("unexpected") } @@ -158,6 +155,9 @@ func (s *apiKeyRepoStubForGroupUpdate) ResetRateLimitWindows(context.Context, in func (s *apiKeyRepoStubForGroupUpdate) GetRateLimitData(context.Context, int64) (*APIKeyRateLimitData, error) { panic("unexpected") } +func (s *apiKeyRepoStubForGroupUpdate) UpdateGroupIDByUserAndGroup(context.Context, int64, int64, int64) (int64, error) { + panic("unexpected") +} // groupRepoStubForGroupUpdate implements GroupRepository for AdminUpdateAPIKeyGroupID tests. type groupRepoStubForGroupUpdate struct { diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index 5aed323a73..c9f32b3b0c 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -191,6 +191,9 @@ func (s *BillingService) initFallbackPricing() { // Claude 4.6 Opus (与4.5同价) s.fallbackPrices["claude-opus-4.6"] = s.fallbackPrices["claude-opus-4.5"] + // Claude 4.7 Opus (暂与4.6同价,待官方定价更新) + s.fallbackPrices["claude-opus-4.7"] = s.fallbackPrices["claude-opus-4.6"] + // Gemini 3.1 Pro s.fallbackPrices["gemini-3.1-pro"] = &ModelPricing{ InputPricePerToken: 2e-6, // $2 per MTok @@ -278,6 +281,9 @@ func (s *BillingService) getFallbackPricing(model string) *ModelPricing { // 按模型系列匹配 if strings.Contains(modelLower, "opus") { + if strings.Contains(modelLower, "4.7") || strings.Contains(modelLower, "4-7") { + return s.fallbackPrices["claude-opus-4.7"] + } if strings.Contains(modelLower, "4.6") || strings.Contains(modelLower, "4-6") { return s.fallbackPrices["claude-opus-4.6"] } diff --git a/backend/internal/service/channel.go b/backend/internal/service/channel.go index 09aa0fc11c..93beb97277 100644 --- a/backend/internal/service/channel.go +++ b/backend/internal/service/channel.go @@ -37,9 +37,10 @@ type Channel struct { Name string Description string Status string - BillingModelSource string // "requested", "upstream", or "channel_mapped" - RestrictModels bool // 是否限制模型(仅允许定价列表中的模型) - Features string // 渠道特性描述(JSON 数组),用于支付页面展示 + BillingModelSource string // "requested", "upstream", or "channel_mapped" + RestrictModels bool // 是否限制模型(仅允许定价列表中的模型) + Features string // 渠道特性描述(JSON 数组),用于支付页面展示 + FeaturesConfig map[string]any // 渠道功能配置(如 web search emulation) CreatedAt time.Time UpdatedAt time.Time @@ -49,8 +50,6 @@ type Channel struct { ModelPricing []ChannelModelPricing // 渠道级模型映射(按平台分组:platform → {src→dst}) ModelMapping map[string]map[string]string - // 渠道特性配置(如 {"web_search_emulation": {"anthropic": true}}) - FeaturesConfig map[string]any // 账号统计定价 ApplyPricingToAccountStats bool // 是否应用渠道模型定价到账号统计 @@ -72,19 +71,6 @@ type AccountStatsPricingRule struct { UpdatedAt time.Time } -// IsWebSearchEmulationEnabled 返回该渠道是否为指定平台启用了 web search 模拟。 -func (c *Channel) IsWebSearchEmulationEnabled(platform string) bool { - if c == nil || c.FeaturesConfig == nil { - return false - } - wse, ok := c.FeaturesConfig[featureKeyWebSearchEmulation].(map[string]any) - if !ok { - return false - } - enabled, ok := wse[platform].(bool) - return ok && enabled -} - // ChannelModelPricing 渠道模型定价条目 type ChannelModelPricing struct { ID int64 @@ -237,6 +223,19 @@ func (c *Channel) Clone() *Channel { return &cp } +// IsWebSearchEmulationEnabled 返回该渠道是否为指定平台启用了 web search 模拟。 +func (c *Channel) IsWebSearchEmulationEnabled(platform string) bool { + if c == nil || c.FeaturesConfig == nil { + return false + } + wse, ok := c.FeaturesConfig[featureKeyWebSearchEmulation].(map[string]any) + if !ok { + return false + } + enabled, ok := wse[platform].(bool) + return ok && enabled +} + // deepCopyFeaturesConfig creates a deep copy of FeaturesConfig to prevent cache pollution. func deepCopyFeaturesConfig(src map[string]any) map[string]any { dst := make(map[string]any, len(src)) diff --git a/backend/internal/service/domain_constants.go b/backend/internal/service/domain_constants.go index 771e6411a7..cb452efbee 100644 --- a/backend/internal/service/domain_constants.go +++ b/backend/internal/service/domain_constants.go @@ -250,10 +250,6 @@ const ( // SettingKeyEnableCCHSigning 是否对 billing header 中的 cch 进行 xxHash64 签名(默认 false) SettingKeyEnableCCHSigning = "enable_cch_signing" - // Web Search Emulation - // SettingKeyWebSearchEmulationConfig 全局 web search 模拟配置(JSON) - SettingKeyWebSearchEmulationConfig = "web_search_emulation_config" - // Balance Low Notification SettingKeyBalanceLowNotifyEnabled = "balance_low_notify_enabled" // 全局开关 SettingKeyBalanceLowNotifyThreshold = "balance_low_notify_threshold" // 默认阈值(USD) @@ -262,6 +258,9 @@ const ( // Account Quota Notification SettingKeyAccountQuotaNotifyEnabled = "account_quota_notify_enabled" // 全局开关 SettingKeyAccountQuotaNotifyEmails = "account_quota_notify_emails" // 管理员通知邮箱列表(JSON 数组) + + // Web Search Emulation + SettingKeyWebSearchEmulationConfig = "web_search_emulation_config" // JSON 配置 ) // AdminAPIKeyPrefix is the prefix for admin API keys (distinct from user "sk-" keys). diff --git a/backend/internal/service/email_service.go b/backend/internal/service/email_service.go index b01a2ef774..9a03ea30d4 100644 --- a/backend/internal/service/email_service.go +++ b/backend/internal/service/email_service.go @@ -50,7 +50,7 @@ type EmailCache interface { IsPasswordResetEmailInCooldown(ctx context.Context, email string) bool SetPasswordResetEmailCooldown(ctx context.Context, email string, ttl time.Duration) error - // User-level rate limiting for notify email verification codes + // Notify code rate limiting per user IncrNotifyCodeUserRate(ctx context.Context, userID int64, window time.Duration) (int64, error) GetNotifyCodeUserRate(ctx context.Context, userID int64) (int64, error) } diff --git a/backend/internal/service/openai_account_scheduler_ws_snapshot_test.go b/backend/internal/service/openai_account_scheduler_ws_snapshot_test.go new file mode 100644 index 0000000000..c5de820341 --- /dev/null +++ b/backend/internal/service/openai_account_scheduler_ws_snapshot_test.go @@ -0,0 +1,62 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +func TestOpenAIGatewayService_SelectAccountWithScheduler_UsesWSPassthroughSnapshotFlags(t *testing.T) { + ctx := context.Background() + groupID := int64(10105) + account := &Account{ + ID: 35001, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 10, + Extra: map[string]any{ + "openai_oauth_responses_websockets_v2_mode": OpenAIWSIngressModePassthrough, + }, + } + + snapshotCache := &openAISnapshotCacheStub{ + snapshotAccounts: []*Account{account}, + accountsByID: map[int64]*Account{account.ID: account}, + } + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.Enabled = true + cfg.Gateway.OpenAIWS.OAuthEnabled = true + cfg.Gateway.OpenAIWS.APIKeyEnabled = true + cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true + cfg.Gateway.OpenAIWS.ModeRouterV2Enabled = true + cfg.Gateway.OpenAIWS.IngressModeDefault = OpenAIWSIngressModeCtxPool + + svc := &OpenAIGatewayService{ + accountRepo: stubOpenAIAccountRepo{accounts: []Account{*account}}, + cache: &stubGatewayCache{}, + cfg: cfg, + schedulerSnapshot: &SchedulerSnapshotService{cache: snapshotCache}, + concurrencyService: NewConcurrencyService(stubConcurrencyCache{}), + } + + selection, decision, err := svc.SelectAccountWithScheduler( + ctx, + &groupID, + "", + "session_hash_ws_passthrough", + "gpt-5.1", + nil, + OpenAIUpstreamTransportResponsesWebsocketV2, + ) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, account.ID, selection.Account.ID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) +} diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index a72b9bbf4b..2a0a72eb95 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -121,6 +121,28 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( } } + // For API key accounts (including OpenAI-compatible upstream gateways), + // ensure promptCacheKey is also propagated via the request body so that + // upstreams using the Responses API can derive a stable session identifier + // from prompt_cache_key. This makes our Anthropic /v1/messages compatibility + // path behave more like a native Responses client. + if account.Type == AccountTypeAPIKey { + if trimmedKey := strings.TrimSpace(promptCacheKey); trimmedKey != "" { + var reqBody map[string]any + if err := json.Unmarshal(responsesBody, &reqBody); err != nil { + return nil, fmt.Errorf("unmarshal for prompt cache key injection: %w", err) + } + if existing, ok := reqBody["prompt_cache_key"].(string); !ok || strings.TrimSpace(existing) == "" { + reqBody["prompt_cache_key"] = trimmedKey + updated, err := json.Marshal(reqBody) + if err != nil { + return nil, fmt.Errorf("remarshal after prompt cache key injection: %w", err) + } + responsesBody = updated + } + } + } + // 5. Get access token token, _, err := s.GetAccessToken(ctx, account) if err != nil { diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 4ef9eb8a53..10b19c24b0 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -1681,7 +1681,6 @@ func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Co if err != nil || latest == nil { return nil } - syncOpenAICodexRateLimitFromExtra(ctx, s.accountRepo, latest, time.Now()) if !latest.IsSchedulable() || !latest.IsOpenAI() { return nil } @@ -1704,7 +1703,6 @@ func (s *OpenAIGatewayService) getSchedulableAccount(ctx context.Context, accoun if err != nil || account == nil { return account, err } - syncOpenAICodexRateLimitFromExtra(ctx, s.accountRepo, account, time.Now()) return account, nil } @@ -4748,69 +4746,6 @@ func buildCodexUsageExtraUpdates(snapshot *OpenAICodexUsageSnapshot, fallbackNow return updates } -func codexUsagePercentExhausted(value *float64) bool { - return value != nil && *value >= 100-1e-9 -} - -func codexRateLimitResetAtFromSnapshot(snapshot *OpenAICodexUsageSnapshot, fallbackNow time.Time) *time.Time { - if snapshot == nil { - return nil - } - normalized := snapshot.Normalize() - if normalized == nil { - return nil - } - baseTime := codexSnapshotBaseTime(snapshot, fallbackNow) - if codexUsagePercentExhausted(normalized.Used7dPercent) && normalized.Reset7dSeconds != nil { - resetAt := baseTime.Add(time.Duration(*normalized.Reset7dSeconds) * time.Second) - return &resetAt - } - if codexUsagePercentExhausted(normalized.Used5hPercent) && normalized.Reset5hSeconds != nil { - resetAt := baseTime.Add(time.Duration(*normalized.Reset5hSeconds) * time.Second) - return &resetAt - } - return nil -} - -func codexRateLimitResetAtFromExtra(extra map[string]any, now time.Time) *time.Time { - if len(extra) == 0 { - return nil - } - if progress := buildCodexUsageProgressFromExtra(extra, "7d", now); progress != nil && codexUsagePercentExhausted(&progress.Utilization) && progress.ResetsAt != nil && now.Before(*progress.ResetsAt) { - resetAt := progress.ResetsAt.UTC() - return &resetAt - } - if progress := buildCodexUsageProgressFromExtra(extra, "5h", now); progress != nil && codexUsagePercentExhausted(&progress.Utilization) && progress.ResetsAt != nil && now.Before(*progress.ResetsAt) { - resetAt := progress.ResetsAt.UTC() - return &resetAt - } - return nil -} - -func applyOpenAICodexRateLimitFromExtra(account *Account, now time.Time) (*time.Time, bool) { - if account == nil || !account.IsOpenAI() { - return nil, false - } - resetAt := codexRateLimitResetAtFromExtra(account.Extra, now) - if resetAt == nil { - return nil, false - } - if account.RateLimitResetAt != nil && now.Before(*account.RateLimitResetAt) && !account.RateLimitResetAt.Before(*resetAt) { - return account.RateLimitResetAt, false - } - account.RateLimitResetAt = resetAt - return resetAt, true -} - -func syncOpenAICodexRateLimitFromExtra(ctx context.Context, repo AccountRepository, account *Account, now time.Time) *time.Time { - resetAt, changed := applyOpenAICodexRateLimitFromExtra(account, now) - if !changed || resetAt == nil || repo == nil || account == nil || account.ID <= 0 { - return resetAt - } - _ = repo.SetRateLimited(ctx, account.ID, *resetAt) - return resetAt -} - // updateCodexUsageSnapshot saves the Codex usage snapshot to account's Extra field func (s *OpenAIGatewayService) updateCodexUsageSnapshot(ctx context.Context, accountID int64, snapshot *OpenAICodexUsageSnapshot) { if snapshot == nil { @@ -4822,24 +4757,17 @@ func (s *OpenAIGatewayService) updateCodexUsageSnapshot(ctx context.Context, acc now := time.Now() updates := buildCodexUsageExtraUpdates(snapshot, now) - resetAt := codexRateLimitResetAtFromSnapshot(snapshot, now) - if len(updates) == 0 && resetAt == nil { + if len(updates) == 0 { return } - shouldPersistUpdates := len(updates) > 0 && s.getCodexSnapshotThrottle().Allow(accountID, now) - if !shouldPersistUpdates && resetAt == nil { + if !s.getCodexSnapshotThrottle().Allow(accountID, now) { return } go func() { updateCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() - if shouldPersistUpdates { - _ = s.accountRepo.UpdateExtra(updateCtx, accountID, updates) - } - if resetAt != nil { - _ = s.accountRepo.SetRateLimited(updateCtx, accountID, *resetAt) - } + _ = s.accountRepo.UpdateExtra(updateCtx, accountID, updates) }() } diff --git a/backend/internal/service/openai_ws_ratelimit_signal_test.go b/backend/internal/service/openai_ws_ratelimit_signal_test.go index 6313d0c08e..4ee85a3a09 100644 --- a/backend/internal/service/openai_ws_ratelimit_signal_test.go +++ b/backend/internal/service/openai_ws_ratelimit_signal_test.go @@ -345,7 +345,7 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_ErrorEventUsageL } } -func TestOpenAIGatewayService_UpdateCodexUsageSnapshot_ExhaustedSnapshotSetsRateLimit(t *testing.T) { +func TestOpenAIGatewayService_UpdateCodexUsageSnapshot_ExhaustedSnapshotDoesNotSetRateLimit(t *testing.T) { repo := &openAICodexSnapshotAsyncRepo{ updateExtraCh: make(chan map[string]any, 1), rateLimitCh: make(chan time.Time, 1), @@ -359,7 +359,6 @@ func TestOpenAIGatewayService_UpdateCodexUsageSnapshot_ExhaustedSnapshotSetsRate SecondaryResetAfterSeconds: ptrIntWS(1200), SecondaryWindowMinutes: ptrIntWS(300), } - before := time.Now() svc.updateCodexUsageSnapshot(context.Background(), 601, snapshot) select { @@ -371,9 +370,8 @@ func TestOpenAIGatewayService_UpdateCodexUsageSnapshot_ExhaustedSnapshotSetsRate select { case resetAt := <-repo.rateLimitCh: - require.WithinDuration(t, before.Add(time.Hour), resetAt, 2*time.Second) + t.Fatalf("不应因仅写入快照而生成运行时限流时间: %v", resetAt) case <-time.After(2 * time.Second): - t.Fatal("等待 codex 100% 自动切换限流超时") } } @@ -401,7 +399,7 @@ func TestOpenAIGatewayService_UpdateCodexUsageSnapshot_NonExhaustedSnapshotDoesN select { case resetAt := <-repo.rateLimitCh: - t.Fatalf("unexpected rate limit reset at: %v", resetAt) + t.Fatalf("不应写入运行时限流时间: %v", resetAt) case <-time.After(200 * time.Millisecond): } } @@ -409,7 +407,6 @@ func TestOpenAIGatewayService_UpdateCodexUsageSnapshot_NonExhaustedSnapshotDoesN func TestOpenAIGatewayService_UpdateCodexUsageSnapshot_ThrottlesExtraWrites(t *testing.T) { repo := &openAICodexSnapshotAsyncRepo{ updateExtraCh: make(chan map[string]any, 2), - rateLimitCh: make(chan time.Time, 2), } svc := &OpenAIGatewayService{ accountRepo: repo, @@ -443,7 +440,7 @@ func TestOpenAIGatewayService_UpdateCodexUsageSnapshot_ThrottlesExtraWrites(t *t func ptrFloat64WS(v float64) *float64 { return &v } func ptrIntWS(v int) *int { return &v } -func TestOpenAIGatewayService_GetSchedulableAccount_ExhaustedCodexExtraSetsRateLimit(t *testing.T) { +func TestOpenAIGatewayService_GetSchedulableAccount_ExhaustedCodexExtraDoesNotSetRateLimit(t *testing.T) { resetAt := time.Now().Add(6 * 24 * time.Hour) account := Account{ ID: 701, @@ -463,17 +460,15 @@ func TestOpenAIGatewayService_GetSchedulableAccount_ExhaustedCodexExtraSetsRateL fresh, err := svc.getSchedulableAccount(context.Background(), account.ID) require.NoError(t, err) require.NotNil(t, fresh) - require.NotNil(t, fresh.RateLimitResetAt) - require.WithinDuration(t, resetAt.UTC(), *fresh.RateLimitResetAt, time.Second) + require.Nil(t, fresh.RateLimitResetAt) select { case persisted := <-repo.rateLimitCh: - require.WithinDuration(t, resetAt.UTC(), persisted, time.Second) + t.Fatalf("不应将已耗尽的 codex extra 提升为运行时限流状态: %v", persisted) case <-time.After(2 * time.Second): - t.Fatal("等待旧快照补写限流状态超时") } } -func TestAdminService_ListAccounts_ExhaustedCodexExtraReturnsRateLimitedAccount(t *testing.T) { +func TestAdminService_ListAccounts_ExhaustedCodexExtraDoesNotSetRateLimit(t *testing.T) { resetAt := time.Now().Add(4 * 24 * time.Hour) repo := &openAICodexExtraListRepo{ stubOpenAIAccountRepo: stubOpenAIAccountRepo{accounts: []Account{{ @@ -496,13 +491,11 @@ func TestAdminService_ListAccounts_ExhaustedCodexExtraReturnsRateLimitedAccount( require.NoError(t, err) require.Equal(t, int64(1), total) require.Len(t, accounts, 1) - require.NotNil(t, accounts[0].RateLimitResetAt) - require.WithinDuration(t, resetAt.UTC(), *accounts[0].RateLimitResetAt, time.Second) + require.Nil(t, accounts[0].RateLimitResetAt) select { case persisted := <-repo.rateLimitCh: - require.WithinDuration(t, resetAt.UTC(), persisted, time.Second) + t.Fatalf("不应在账号列表查询时将 codex extra 持久化为运行时限流状态: %v", persisted) case <-time.After(2 * time.Second): - t.Fatal("等待列表补写限流状态超时") } } diff --git a/backend/internal/service/payment_config_providers.go b/backend/internal/service/payment_config_providers.go index f949c8b475..b915d8f3fc 100644 --- a/backend/internal/service/payment_config_providers.go +++ b/backend/internal/service/payment_config_providers.go @@ -231,10 +231,18 @@ func (s *PaymentConfigService) UpdateProviderInstance(ctx context.Context, id in } } if req.AllowUserRefund != nil { - // Only allow enabling when refund_enabled is true + // Only allow enabling when refund_enabled is (or will be) true if *req.AllowUserRefund { - inst, err := s.entClient.PaymentProviderInstance.Get(ctx, id) - if err == nil && inst.RefundEnabled { + refundEnabled := false + if req.RefundEnabled != nil { + refundEnabled = *req.RefundEnabled + } else { + inst, err := s.entClient.PaymentProviderInstance.Get(ctx, id) + if err == nil { + refundEnabled = inst.RefundEnabled + } + } + if refundEnabled { u.SetAllowUserRefund(true) } } else { @@ -251,8 +259,8 @@ func (s *PaymentConfigService) UpdateProviderInstance(ctx context.Context, id in func (s *PaymentConfigService) GetUserRefundEligibleInstanceIDs(ctx context.Context) ([]string, error) { instances, err := s.entClient.PaymentProviderInstance.Query(). Where( - paymentproviderinstance.AllowUserRefundEQ(true), paymentproviderinstance.RefundEnabledEQ(true), + paymentproviderinstance.AllowUserRefundEQ(true), ).Select(paymentproviderinstance.FieldID).All(ctx) if err != nil { return nil, err diff --git a/backend/internal/service/payment_refund.go b/backend/internal/service/payment_refund.go index 685ef158ef..c5bda763cd 100644 --- a/backend/internal/service/payment_refund.go +++ b/backend/internal/service/payment_refund.go @@ -2,6 +2,7 @@ package service import ( "context" + "errors" "fmt" "log/slog" "math" @@ -93,15 +94,16 @@ func (s *PaymentService) PrepareRefund(ctx context.Context, oid int64, amt float // Check provider instance allows admin refund inst, instErr := s.getOrderProviderInstance(ctx, o) if instErr != nil { - slog.Warn("refund: provider instance not found", "orderID", oid, "error", instErr) + slog.Warn("refund: provider instance lookup failed", "orderID", oid, "error", instErr) + return nil, nil, infraerrors.InternalServer("PROVIDER_LOOKUP_FAILED", "failed to look up payment provider for this order") } - if inst != nil && !inst.RefundEnabled { - return nil, nil, infraerrors.Forbidden("REFUND_DISABLED", "refund is not enabled for this provider") - } - if inst == nil && instErr == nil { + if inst == nil { // Legacy order without provider_instance_id — block refund return nil, nil, infraerrors.Forbidden("REFUND_DISABLED", "refund is not available for this order") } + if !inst.RefundEnabled { + return nil, nil, infraerrors.Forbidden("REFUND_DISABLED", "refund is not enabled for this provider") + } if math.IsNaN(amt) || math.IsInf(amt, 0) { return nil, nil, infraerrors.BadRequest("INVALID_AMOUNT", "invalid refund amount") } @@ -179,11 +181,17 @@ func (s *PaymentService) ExecuteRefund(ctx context.Context, p *RefundPlan) (*Ref if !s.hasAuditLog(ctx, p.OrderID, "REFUND_ROLLBACK_FAILED") { _, err := s.subscriptionSvc.ExtendSubscription(ctx, p.SubscriptionID, -p.SubDaysToDeduct) if err != nil { - // If deducting would expire the subscription, revoke it entirely - slog.Info("subscription deduction would expire, revoking", "orderID", p.OrderID, "subID", p.SubscriptionID, "days", p.SubDaysToDeduct) - if revokeErr := s.subscriptionSvc.RevokeSubscription(ctx, p.SubscriptionID); revokeErr != nil { + if errors.Is(err, ErrAdjustWouldExpire) { + // Deduction would expire the subscription — revoke it entirely + slog.Info("subscription deduction would expire, revoking", "orderID", p.OrderID, "subID", p.SubscriptionID, "days", p.SubDaysToDeduct) + if revokeErr := s.subscriptionSvc.RevokeSubscription(ctx, p.SubscriptionID); revokeErr != nil { + s.restoreStatus(ctx, p) + return nil, fmt.Errorf("revoke subscription: %w", revokeErr) + } + } else { + // Other errors (DB failure, not found) — abort refund s.restoreStatus(ctx, p) - return nil, fmt.Errorf("revoke subscription: %w", revokeErr) + return nil, fmt.Errorf("deduct subscription days: %w", err) } } } else { diff --git a/backend/internal/service/pricing_service.go b/backend/internal/service/pricing_service.go index 3b3f31c309..2bf48702aa 100644 --- a/backend/internal/service/pricing_service.go +++ b/backend/internal/service/pricing_service.go @@ -656,65 +656,95 @@ func (s *PricingService) extractBaseName(model string) string { // matchByModelFamily 基于模型系列匹配 func (s *PricingService) matchByModelFamily(model string) *LiteLLMModelPricing { - // Claude模型系列匹配规则 - familyPatterns := map[string][]string{ - "opus-4.6": {"claude-opus-4.6", "claude-opus-4-6"}, - "opus-4.5": {"claude-opus-4.5", "claude-opus-4-5"}, - "opus-4": {"claude-opus-4", "claude-3-opus"}, - "sonnet-4.5": {"claude-sonnet-4.5", "claude-sonnet-4-5"}, - "sonnet-4": {"claude-sonnet-4", "claude-3-5-sonnet"}, - "sonnet-3.5": {"claude-3-5-sonnet", "claude-3.5-sonnet"}, - "sonnet-3": {"claude-3-sonnet"}, - "haiku-3.5": {"claude-3-5-haiku", "claude-3.5-haiku"}, - "haiku-3": {"claude-3-haiku"}, + // modelFamily 定义一个模型系列的匹配和定价查找规则。 + type modelFamily struct { + name string // 系列名称 + match []string // 用于将模型归类到此系列的模式(strings.Contains 匹配) + pricing []string // 用于在定价数据中查找价格的模式(nil 则复用 match;可包含低版本 fallback) } - // 确定模型属于哪个系列 - var matchedFamily string - for family, patterns := range familyPatterns { - for _, pattern := range patterns { + // 按特异性降序排列:高版本号在前,避免 "claude-opus-4"(opus-4 系列) + // 因子串关系误匹配 "claude-opus-4-7"(opus-4.7 系列)。 + // 注意:原 map 实现存在 Go map 迭代随机性导致的同类 bug,此处改为有序切片修复。 + families := []modelFamily{ + {name: "opus-4.7", match: []string{"claude-opus-4-7", "claude-opus-4.7"}, pricing: []string{"claude-opus-4-7", "claude-opus-4.7", "claude-opus-4-6"}}, + {name: "opus-4.6", match: []string{"claude-opus-4-6", "claude-opus-4.6"}}, + {name: "opus-4.5", match: []string{"claude-opus-4-5", "claude-opus-4.5"}}, + {name: "opus-4", match: []string{"claude-opus-4", "claude-3-opus"}}, + {name: "sonnet-4.5", match: []string{"claude-sonnet-4-5", "claude-sonnet-4.5"}}, + {name: "sonnet-4", match: []string{"claude-sonnet-4", "claude-3-5-sonnet"}}, + {name: "sonnet-3.5", match: []string{"claude-3-5-sonnet", "claude-3.5-sonnet"}}, + {name: "sonnet-3", match: []string{"claude-3-sonnet"}}, + {name: "haiku-3.5", match: []string{"claude-3-5-haiku", "claude-3.5-haiku"}}, + {name: "haiku-3", match: []string{"claude-3-haiku"}}, + } + + // Phase 1: 按有序切片归类(最具体的系列优先匹配) + var matched *modelFamily + for i := range families { + for _, pattern := range families[i].match { if strings.Contains(model, pattern) || strings.Contains(model, strings.ReplaceAll(pattern, "-", "")) { - matchedFamily = family + matched = &families[i] break } } - if matchedFamily != "" { + if matched != nil { break } } - if matchedFamily == "" { - // 简单的系列匹配 - if strings.Contains(model, "opus") { - if strings.Contains(model, "4.5") || strings.Contains(model, "4-5") { - matchedFamily = "opus-4.5" - } else { - matchedFamily = "opus-4" + // Phase 2: 二次兜底——当模型 ID 不含已知模式串时,按关键字粗分 + if matched == nil { + var fallbackName string + switch { + case strings.Contains(model, "opus"): + switch { + case strings.Contains(model, "4.7") || strings.Contains(model, "4-7"): + fallbackName = "opus-4.7" + case strings.Contains(model, "4.6") || strings.Contains(model, "4-6"): + fallbackName = "opus-4.6" + case strings.Contains(model, "4.5") || strings.Contains(model, "4-5"): + fallbackName = "opus-4.5" + default: + fallbackName = "opus-4" } - } else if strings.Contains(model, "sonnet") { - if strings.Contains(model, "4.5") || strings.Contains(model, "4-5") { - matchedFamily = "sonnet-4.5" - } else if strings.Contains(model, "3-5") || strings.Contains(model, "3.5") { - matchedFamily = "sonnet-3.5" - } else { - matchedFamily = "sonnet-4" + case strings.Contains(model, "sonnet"): + switch { + case strings.Contains(model, "4.5") || strings.Contains(model, "4-5"): + fallbackName = "sonnet-4.5" + case strings.Contains(model, "3-5") || strings.Contains(model, "3.5"): + fallbackName = "sonnet-3.5" + default: + fallbackName = "sonnet-4" } - } else if strings.Contains(model, "haiku") { - if strings.Contains(model, "3-5") || strings.Contains(model, "3.5") { - matchedFamily = "haiku-3.5" - } else { - matchedFamily = "haiku-3" + case strings.Contains(model, "haiku"): + switch { + case strings.Contains(model, "3-5") || strings.Contains(model, "3.5"): + fallbackName = "haiku-3.5" + default: + fallbackName = "haiku-3" + } + } + if fallbackName != "" { + for i := range families { + if families[i].name == fallbackName { + matched = &families[i] + break + } } } } - if matchedFamily == "" { + if matched == nil { return nil } - // 在价格数据中查找该系列的模型 - patterns := familyPatterns[matchedFamily] - for _, pattern := range patterns { + // Phase 3: 在定价数据中查找该系列的价格 + lookups := matched.pricing + if lookups == nil { + lookups = matched.match + } + for _, pattern := range lookups { for key, pricing := range s.pricingData { keyLower := strings.ToLower(key) if strings.Contains(keyLower, pattern) { diff --git a/backend/internal/service/ratelimit_service.go b/backend/internal/service/ratelimit_service.go index 4d8009b7bd..53581574bd 100644 --- a/backend/internal/service/ratelimit_service.go +++ b/backend/internal/service/ratelimit_service.go @@ -152,6 +152,11 @@ func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Acc msg := "Credit balance exhausted (400): " + upstreamMsg s.handleAuthError(ctx, account, msg) shouldDisable = true + } else if strings.Contains(strings.ToLower(upstreamMsg), "identity verification is required") { + // KYC 身份验证要求 → 永久禁用,账号需完成身份验证后才能恢复 + msg := "Identity verification required (400): " + upstreamMsg + s.handleAuthError(ctx, account, msg) + shouldDisable = true } // 其他 400 错误(如参数问题)不处理,不禁用账号 case 401: diff --git a/backend/internal/service/scheduler_snapshot_service.go b/backend/internal/service/scheduler_snapshot_service.go index d1330abb87..62b6993d0e 100644 --- a/backend/internal/service/scheduler_snapshot_service.go +++ b/backend/internal/service/scheduler_snapshot_service.go @@ -20,6 +20,14 @@ var ( const outboxEventTimeout = 2 * time.Minute +// batchSeenKey tracks which (groupID, platform) bucket sets have already been +// rebuilt within a single pollOutbox call, to avoid redundant work when multiple +// account_changed events share the same groups. +type batchSeenKey struct { + groupID int64 + platform string +} + type SchedulerSnapshotService struct { cache SchedulerCache outboxRepo SchedulerOutboxRepository @@ -244,9 +252,10 @@ func (s *SchedulerSnapshotService) pollOutbox() { } watermarkForCheck := watermark + seen := make(map[batchSeenKey]struct{}) for _, event := range events { eventCtx, cancel := context.WithTimeout(context.Background(), outboxEventTimeout) - err := s.handleOutboxEvent(eventCtx, event) + err := s.handleOutboxEvent(eventCtx, event, seen) cancel() if err != nil { logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] outbox handle failed: id=%d type=%s err=%v", event.ID, event.EventType, err) @@ -255,8 +264,20 @@ func (s *SchedulerSnapshotService) pollOutbox() { } lastID := events[len(events)-1].ID - if err := s.cache.SetOutboxWatermark(ctx, lastID); err != nil { - logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] outbox watermark write failed: %v", err) + var wmErr error + for i := range 3 { + wmCtx, wmCancel := context.WithTimeout(context.Background(), 5*time.Second) + wmErr = s.cache.SetOutboxWatermark(wmCtx, lastID) + wmCancel() + if wmErr == nil { + break + } + if i < 2 { + time.Sleep(200 * time.Millisecond) + } + } + if wmErr != nil { + logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] outbox watermark write failed: %v", wmErr) } else { watermarkForCheck = lastID } @@ -264,18 +285,18 @@ func (s *SchedulerSnapshotService) pollOutbox() { s.checkOutboxLag(ctx, events[0], watermarkForCheck) } -func (s *SchedulerSnapshotService) handleOutboxEvent(ctx context.Context, event SchedulerOutboxEvent) error { +func (s *SchedulerSnapshotService) handleOutboxEvent(ctx context.Context, event SchedulerOutboxEvent, seen map[batchSeenKey]struct{}) error { switch event.EventType { case SchedulerOutboxEventAccountLastUsed: return s.handleLastUsedEvent(ctx, event.Payload) case SchedulerOutboxEventAccountBulkChanged: - return s.handleBulkAccountEvent(ctx, event.Payload) + return s.handleBulkAccountEvent(ctx, event.Payload, seen) case SchedulerOutboxEventAccountGroupsChanged: - return s.handleAccountEvent(ctx, event.AccountID, event.Payload) + return s.handleAccountEvent(ctx, event.AccountID, event.Payload, seen) case SchedulerOutboxEventAccountChanged: - return s.handleAccountEvent(ctx, event.AccountID, event.Payload) + return s.handleAccountEvent(ctx, event.AccountID, event.Payload, seen) case SchedulerOutboxEventGroupChanged: - return s.handleGroupEvent(ctx, event.GroupID) + return s.handleGroupEvent(ctx, event.GroupID, seen) case SchedulerOutboxEventFullRebuild: return s.triggerFullRebuild("outbox") default: @@ -309,7 +330,7 @@ func (s *SchedulerSnapshotService) handleLastUsedEvent(ctx context.Context, payl return s.cache.UpdateLastUsed(ctx, updates) } -func (s *SchedulerSnapshotService) handleBulkAccountEvent(ctx context.Context, payload map[string]any) error { +func (s *SchedulerSnapshotService) handleBulkAccountEvent(ctx context.Context, payload map[string]any, seen map[batchSeenKey]struct{}) error { if payload == nil { return nil } @@ -323,15 +344,15 @@ func (s *SchedulerSnapshotService) handleBulkAccountEvent(ctx context.Context, p } ids := make([]int64, 0, len(rawIDs)) - seen := make(map[int64]struct{}, len(rawIDs)) + seenIDs := make(map[int64]struct{}, len(rawIDs)) for _, id := range rawIDs { if id <= 0 { continue } - if _, exists := seen[id]; exists { + if _, exists := seenIDs[id]; exists { continue } - seen[id] = struct{}{} + seenIDs[id] = struct{}{} ids = append(ids, id) } if len(ids) == 0 { @@ -384,10 +405,10 @@ func (s *SchedulerSnapshotService) handleBulkAccountEvent(ctx context.Context, p for gid := range rebuildGroupSet { rebuildGroupIDs = append(rebuildGroupIDs, gid) } - return s.rebuildByGroupIDs(ctx, rebuildGroupIDs, "account_bulk_change") + return s.rebuildByGroupIDs(ctx, rebuildGroupIDs, "account_bulk_change", seen) } -func (s *SchedulerSnapshotService) handleAccountEvent(ctx context.Context, accountID *int64, payload map[string]any) error { +func (s *SchedulerSnapshotService) handleAccountEvent(ctx context.Context, accountID *int64, payload map[string]any, seen map[batchSeenKey]struct{}) error { if accountID == nil || *accountID <= 0 { return nil } @@ -408,7 +429,7 @@ func (s *SchedulerSnapshotService) handleAccountEvent(ctx context.Context, accou return err } } - return s.rebuildByGroupIDs(ctx, groupIDs, "account_miss") + return s.rebuildByGroupIDs(ctx, groupIDs, "account_miss", seen) } return err } @@ -420,18 +441,18 @@ func (s *SchedulerSnapshotService) handleAccountEvent(ctx context.Context, accou if len(groupIDs) == 0 { groupIDs = account.GroupIDs } - return s.rebuildByAccount(ctx, account, groupIDs, "account_change") + return s.rebuildByAccount(ctx, account, groupIDs, "account_change", seen) } -func (s *SchedulerSnapshotService) handleGroupEvent(ctx context.Context, groupID *int64) error { +func (s *SchedulerSnapshotService) handleGroupEvent(ctx context.Context, groupID *int64, seen map[batchSeenKey]struct{}) error { if groupID == nil || *groupID <= 0 { return nil } groupIDs := []int64{*groupID} - return s.rebuildByGroupIDs(ctx, groupIDs, "group_change") + return s.rebuildByGroupIDs(ctx, groupIDs, "group_change", seen) } -func (s *SchedulerSnapshotService) rebuildByAccount(ctx context.Context, account *Account, groupIDs []int64, reason string) error { +func (s *SchedulerSnapshotService) rebuildByAccount(ctx context.Context, account *Account, groupIDs []int64, reason string, seen map[batchSeenKey]struct{}) error { if account == nil { return nil } @@ -441,21 +462,21 @@ func (s *SchedulerSnapshotService) rebuildByAccount(ctx context.Context, account } var firstErr error - if err := s.rebuildBucketsForPlatform(ctx, account.Platform, groupIDs, reason); err != nil && firstErr == nil { + if err := s.rebuildBucketsForPlatform(ctx, account.Platform, groupIDs, reason, seen); err != nil && firstErr == nil { firstErr = err } if account.Platform == PlatformAntigravity && account.IsMixedSchedulingEnabled() { - if err := s.rebuildBucketsForPlatform(ctx, PlatformAnthropic, groupIDs, reason); err != nil && firstErr == nil { + if err := s.rebuildBucketsForPlatform(ctx, PlatformAnthropic, groupIDs, reason, seen); err != nil && firstErr == nil { firstErr = err } - if err := s.rebuildBucketsForPlatform(ctx, PlatformGemini, groupIDs, reason); err != nil && firstErr == nil { + if err := s.rebuildBucketsForPlatform(ctx, PlatformGemini, groupIDs, reason, seen); err != nil && firstErr == nil { firstErr = err } } return firstErr } -func (s *SchedulerSnapshotService) rebuildByGroupIDs(ctx context.Context, groupIDs []int64, reason string) error { +func (s *SchedulerSnapshotService) rebuildByGroupIDs(ctx context.Context, groupIDs []int64, reason string, seen map[batchSeenKey]struct{}) error { groupIDs = s.normalizeGroupIDs(groupIDs) if len(groupIDs) == 0 { return nil @@ -463,19 +484,30 @@ func (s *SchedulerSnapshotService) rebuildByGroupIDs(ctx context.Context, groupI platforms := []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity} var firstErr error for _, platform := range platforms { - if err := s.rebuildBucketsForPlatform(ctx, platform, groupIDs, reason); err != nil && firstErr == nil { + if err := s.rebuildBucketsForPlatform(ctx, platform, groupIDs, reason, seen); err != nil && firstErr == nil { firstErr = err } } return firstErr } -func (s *SchedulerSnapshotService) rebuildBucketsForPlatform(ctx context.Context, platform string, groupIDs []int64, reason string) error { +func (s *SchedulerSnapshotService) rebuildBucketsForPlatform(ctx context.Context, platform string, groupIDs []int64, reason string, seen map[batchSeenKey]struct{}) error { if platform == "" { return nil } var firstErr error for _, gid := range groupIDs { + // Within a single poll batch, skip (groupID, platform) pairs that were + // already rebuilt. The first rebuild loads fresh DB data for all accounts + // in the group, so subsequent rebuilds for the same group+platform within + // the same batch are redundant. + if seen != nil { + key := batchSeenKey{gid, platform} + if _, exists := seen[key]; exists { + continue + } + seen[key] = struct{}{} + } if err := s.rebuildBucket(ctx, SchedulerBucket{GroupID: gid, Platform: platform, Mode: SchedulerModeSingle}, reason); err != nil && firstErr == nil { firstErr = err } diff --git a/backend/internal/service/settings_view.go b/backend/internal/service/settings_view.go index 9ed7440bc1..ab2eb274fd 100644 --- a/backend/internal/service/settings_view.go +++ b/backend/internal/service/settings_view.go @@ -107,8 +107,8 @@ type SystemSettings struct { EnableMetadataPassthrough bool // 是否透传客户端原始 metadata(默认 false) EnableCCHSigning bool // 是否对 billing header cch 进行签名(默认 false) - // Web Search Emulation (read-only quick check; full config via dedicated API) - WebSearchEmulationEnabled bool + // Web Search Emulation + WebSearchEmulationEnabled bool // 是否启用 web search 模拟 // Balance low notification BalanceLowNotifyEnabled bool diff --git a/backend/internal/service/user_service_test.go b/backend/internal/service/user_service_test.go index 6d752f3ed5..a998d5f443 100644 --- a/backend/internal/service/user_service_test.go +++ b/backend/internal/service/user_service_test.go @@ -46,12 +46,12 @@ func (m *mockUserRepo) RemoveGroupFromAllowedGroups(context.Context, int64) (int return 0, nil } func (m *mockUserRepo) AddGroupToAllowedGroups(context.Context, int64, int64) error { return nil } +func (m *mockUserRepo) UpdateTotpSecret(context.Context, int64, *string) error { return nil } +func (m *mockUserRepo) EnableTotp(context.Context, int64) error { return nil } +func (m *mockUserRepo) DisableTotp(context.Context, int64) error { return nil } func (m *mockUserRepo) RemoveGroupFromUserAllowedGroups(context.Context, int64, int64) error { return nil } -func (m *mockUserRepo) UpdateTotpSecret(context.Context, int64, *string) error { return nil } -func (m *mockUserRepo) EnableTotp(context.Context, int64) error { return nil } -func (m *mockUserRepo) DisableTotp(context.Context, int64) error { return nil } // --- mock: APIKeyAuthCacheInvalidator --- diff --git a/frontend/src/__tests__/setup.ts b/frontend/src/__tests__/setup.ts index decb2a370a..0cb4921915 100644 --- a/frontend/src/__tests__/setup.ts +++ b/frontend/src/__tests__/setup.ts @@ -36,6 +36,22 @@ class MockResizeObserver { globalThis.ResizeObserver = MockResizeObserver as unknown as typeof ResizeObserver +// Mock matchMedia (jsdom doesn't implement it). +// Default matches=true so desktop viewport queries pass and components that +// only lazy-load on mobile render content immediately in tests. +if (typeof window !== 'undefined' && !window.matchMedia) { + window.matchMedia = (query: string): MediaQueryList => ({ + matches: true, + media: query, + onchange: null, + addListener: vi.fn(), + removeListener: vi.fn(), + addEventListener: vi.fn(), + removeEventListener: vi.fn(), + dispatchEvent: vi.fn() + }) as MediaQueryList +} + // Vue Test Utils 全局配置 config.global.stubs = { // 可以在这里添加全局 stub diff --git a/frontend/src/api/payment.ts b/frontend/src/api/payment.ts index 013306d950..5cedb107ec 100644 --- a/frontend/src/api/payment.ts +++ b/frontend/src/api/payment.ts @@ -77,6 +77,7 @@ export const paymentAPI = { return apiClient.post(`/payment/orders/${id}/refund-request`, data) }, + /** Get provider instance IDs that allow user refund */ getRefundEligibleProviders() { return apiClient.get<{ provider_instance_ids: string[] }>('/payment/orders/refund-eligible-providers') } diff --git a/frontend/src/components/account/AccountUsageCell.vue b/frontend/src/components/account/AccountUsageCell.vue index c76dc4964c..1c023fb312 100644 --- a/frontend/src/components/account/AccountUsageCell.vue +++ b/frontend/src/components/account/AccountUsageCell.vue @@ -105,7 +105,7 @@
-
- +