diff --git a/backend/cmd/server/wire.go b/backend/cmd/server/wire.go index 8bd9ac5b16..cf5984372d 100644 --- a/backend/cmd/server/wire.go +++ b/backend/cmd/server/wire.go @@ -88,6 +88,7 @@ func provideCleanup( schedulerSnapshot *service.SchedulerSnapshotService, tokenRefresh *service.TokenRefreshService, accountExpiry *service.AccountExpiryService, + cnProviderBalanceCheck *service.CNProviderBalanceCheckService, codexVersionSync *service.OpenAICodexVersionSyncService, proxyExpiry *service.ProxyExpiryService, subscriptionExpiry *service.SubscriptionExpiryService, @@ -238,6 +239,12 @@ func provideCleanup( accountExpiry.Stop() return nil }}, + {"CNProviderBalanceCheckService", func() error { + if cnProviderBalanceCheck != nil { + cnProviderBalanceCheck.Stop() + } + return nil + }}, {"OpenAICodexVersionSyncService", func() error { codexVersionSync.Stop() return nil diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index ab9f4530b8..2b0caf982c 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -218,6 +218,9 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { antigravityOAuthHandler := admin.NewAntigravityOAuthHandler(antigravityOAuthService) tokenRefreshService := service.ProvideTokenRefreshService(accountRepository, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, compositeTokenCacheInvalidator, schedulerCache, configConfig, tempUnschedCache, privacyClientFactory, proxyRepository, oAuthRefreshAPI, openAIGatewayService) grokOAuthHandler := admin.NewGrokOAuthHandler(grokOAuthService, adminService, grokQuotaService, tokenRefreshService) + cnProviderQuotaService := service.ProvideCNProviderQuotaService(accountRepository, proxyRepository, httpUpstream, configConfig) + cnProviderBalanceService := service.ProvideCNProviderBalanceService(accountRepository, proxyRepository, httpUpstream, configConfig) + cnProviderHandler := admin.NewCNProviderHandler(cnProviderQuotaService, cnProviderBalanceService) proxyHandler := admin.NewProxyHandler(adminService) adminRedeemHandler := admin.NewRedeemHandler(adminService, redeemService) promoHandler := admin.NewPromoHandler(promoService) @@ -277,7 +280,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { auditLogHandler := admin.NewAuditLogHandler(auditLogService, totpService) upstreamBillingProbeService := service.ProvideUpstreamBillingProbeService(accountRepository, accountTestService, settingService, leaderLockCache, db) ollamaCloudUsageService := service.ProvideOllamaCloudUsageService(accountRepository, httpUpstream, settingService, secretEncryptor, configConfig, leaderLockCache, db) - adminHandlers := handler.ProvideAdminHandlers(dashboardHandler, adminUserHandler, groupHandler, accountHandler, adminAnnouncementHandler, dataManagementHandler, backupHandler, oAuthHandler, openAIOAuthHandler, geminiOAuthHandler, antigravityOAuthHandler, grokOAuthHandler, proxyHandler, adminRedeemHandler, promoHandler, settingHandler, opsHandler, systemHandler, adminSubscriptionHandler, adminUsageHandler, userAttributeHandler, errorPassthroughHandler, tlsFingerprintProfileHandler, adminAPIKeyHandler, scheduledTestHandler, channelHandler, channelMonitorHandler, channelMonitorRequestTemplateHandler, contentModerationHandler, promptAdminHandler, paymentHandler, affiliateHandler, complianceHandler, auditLogHandler, upstreamBillingProbeService, ollamaCloudUsageService) + adminHandlers := handler.ProvideAdminHandlers(dashboardHandler, adminUserHandler, groupHandler, accountHandler, adminAnnouncementHandler, dataManagementHandler, backupHandler, oAuthHandler, openAIOAuthHandler, geminiOAuthHandler, antigravityOAuthHandler, grokOAuthHandler, cnProviderHandler, proxyHandler, adminRedeemHandler, promoHandler, settingHandler, opsHandler, systemHandler, adminSubscriptionHandler, adminUsageHandler, userAttributeHandler, errorPassthroughHandler, tlsFingerprintProfileHandler, adminAPIKeyHandler, scheduledTestHandler, channelHandler, channelMonitorHandler, channelMonitorRequestTemplateHandler, contentModerationHandler, promptAdminHandler, paymentHandler, affiliateHandler, complianceHandler, auditLogHandler, upstreamBillingProbeService, ollamaCloudUsageService) usageRecordWorkerPool := service.NewUsageRecordWorkerPool(configConfig) userMsgQueueCache := repository.NewUserMsgQueueCache(redisClient) userMessageQueueService := service.ProvideUserMessageQueueService(userMsgQueueCache, rpmCache, configConfig) @@ -327,6 +330,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { opsScheduledReportService := service.ProvideOpsScheduledReportService(opsService, userService, emailService, redisClient, configConfig) opsIngressRejectAggregator := service.ProvideOpsIngressRejectAggregator(opsRepository, opsService) accountExpiryService := service.ProvideAccountExpiryService(accountRepository) + cnProviderBalanceCheckService := service.ProvideCNProviderBalanceCheckService(accountRepository, cnProviderBalanceService, cnProviderQuotaService, configConfig) openAICodexVersionSyncService := service.ProvideOpenAICodexVersionSyncService(settingRepository, settingService, gitHubReleaseClient) proxyExpiryService := service.ProvideProxyExpiryService(proxyRepository) subscriptionExpiryService := service.ProvideSubscriptionExpiryService(userSubscriptionRepository, settingRepository, notificationEmailService, leaderLockCache, db) @@ -336,7 +340,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { channelMonitorRunner := service.ProvideChannelMonitorRunner(channelMonitorService, settingService) channelMonitorV2Aggregator := service.ProvideChannelMonitorV2Aggregator(channelMonitorV2Repository, db, settingService) userPlatformQuotaUsageFlusher := service.ProvideUserPlatformQuotaUsageFlusher(configConfig, billingCache, serviceUserPlatformQuotaRepository, timingWheelService) - v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, opsService, opsIngressRejectAggregator, apiKeyService, authCacheInvalidationWorker, schedulerSnapshotService, tokenRefreshService, accountExpiryService, openAICodexVersionSyncService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, batchImageCleanupService, batchImageWorkerRuntime, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, channelMonitorV2Aggregator, userPlatformQuotaUsageFlusher, upstreamBillingProbeService, ollamaCloudUsageService, auditLogService, promptService) + v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, opsService, opsIngressRejectAggregator, apiKeyService, authCacheInvalidationWorker, schedulerSnapshotService, tokenRefreshService, accountExpiryService, cnProviderBalanceCheckService, openAICodexVersionSyncService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, batchImageCleanupService, batchImageWorkerRuntime, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, channelMonitorV2Aggregator, userPlatformQuotaUsageFlusher, upstreamBillingProbeService, ollamaCloudUsageService, auditLogService, promptService) application := &Application{ Server: httpServer, PromptAudit: promptService, @@ -380,6 +384,7 @@ func provideCleanup( schedulerSnapshot *service.SchedulerSnapshotService, tokenRefresh *service.TokenRefreshService, accountExpiry *service.AccountExpiryService, + cnProviderBalanceCheck *service.CNProviderBalanceCheckService, codexVersionSync *service.OpenAICodexVersionSyncService, proxyExpiry *service.ProxyExpiryService, subscriptionExpiry *service.SubscriptionExpiryService, @@ -529,6 +534,12 @@ func provideCleanup( accountExpiry.Stop() return nil }}, + {"CNProviderBalanceCheckService", func() error { + if cnProviderBalanceCheck != nil { + cnProviderBalanceCheck.Stop() + } + return nil + }}, {"OpenAICodexVersionSyncService", func() error { codexVersionSync.Stop() return nil diff --git a/backend/cmd/server/wire_gen_test.go b/backend/cmd/server/wire_gen_test.go index dff3a9129e..41e41bce13 100644 --- a/backend/cmd/server/wire_gen_test.go +++ b/backend/cmd/server/wire_gen_test.go @@ -66,6 +66,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) { schedulerSnapshotSvc, tokenRefreshSvc, accountExpirySvc, + nil, // cnProviderBalanceCheck codexVersionSyncSvc, proxyExpirySvc, subscriptionExpirySvc, diff --git a/backend/ent/schema/user_platform_quota.go b/backend/ent/schema/user_platform_quota.go index a0b5598600..123d7c0bfa 100644 --- a/backend/ent/schema/user_platform_quota.go +++ b/backend/ent/schema/user_platform_quota.go @@ -41,7 +41,8 @@ func (UserPlatformQuota) Fields() []ent.Field { // 注意:平台列表的单一权威源为 service.AllowedQuotaPlatforms; // 此处为 ent 构建期约束,需与 service.AllowedQuotaPlatforms 保持同步。 switch s { - case "anthropic", "openai", "gemini", "antigravity", "grok": + case "anthropic", "openai", "gemini", "antigravity", "grok", + "kimi", "zhipu", "deepseek": return nil default: return fmt.Errorf("platform %q is not allowed", s) diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index 2c5d5cf545..e437b11150 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -1025,6 +1025,10 @@ type GatewayConfig struct { // Grok: Grok/xAI gateway scheduling and free-tier soft-gate settings. Grok GatewayGrokConfig `mapstructure:"grok"` + + // CNProviders: 国产 OpenAI 兼容供应商(kimi/zhipu/deepseek)的余额检测配置。 + // 仅作用于 payg(按量付费)账号:周期探测余额,低于阈值则临时停调。 + CNProviders GatewayCNProvidersConfig `mapstructure:"cn_providers"` } // GatewayGrokConfig holds Grok-specific gateway scheduling knobs. @@ -1057,6 +1061,18 @@ type GatewayGrokConfig struct { FreeQuotaStatsCacheSeconds int `mapstructure:"free_quota_stats_cache_seconds"` } +// GatewayCNProvidersConfig 国产 OpenAI 兼容供应商(kimi/zhipu/deepseek)的余额检测配置。 +// +// 仅作用于 payg(按量付费)账号(kimi/deepseek 有公开余额端点;zhipu 无,仅靠响应式 429/402)。 +// - balance_check_enabled: 是否启用周期余额检测(默认 true) +// - balance_threshold: 余额低于此值(账户货币单位,默认 0.5)触发临时停调 +// - balance_check_interval_minutes: 余额检测周期(分钟,默认 10) +type GatewayCNProvidersConfig struct { + BalanceCheckEnabled bool `mapstructure:"balance_check_enabled"` + BalanceThreshold float64 `mapstructure:"balance_threshold"` + BalanceCheckIntervalMinutes int `mapstructure:"balance_check_interval_minutes"` +} + type GatewayLiveConfig struct { // MaxSessionDurationSeconds 是 Live 会话的硬上限。 MaxSessionDurationSeconds int `mapstructure:"max_session_duration_seconds"` @@ -2355,6 +2371,10 @@ func setDefaults() { viper.SetDefault("gateway.grok.free_quota_soft_gate_percent", 95) viper.SetDefault("gateway.grok.free_quota_window_hours", 24) viper.SetDefault("gateway.grok.free_quota_stats_cache_seconds", 60) + // 国产供应商余额检测(kimi/deepseek payg;zhipu 无余额端点,仅靠响应式 429/402)。 + viper.SetDefault("gateway.cn_providers.balance_check_enabled", true) + viper.SetDefault("gateway.cn_providers.balance_threshold", 0.5) + viper.SetDefault("gateway.cn_providers.balance_check_interval_minutes", 10) viper.SetDefault("gateway.image_concurrency.enabled", false) viper.SetDefault("gateway.image_concurrency.max_concurrent_requests", 0) viper.SetDefault("gateway.image_concurrency.overflow_mode", ImageConcurrencyOverflowModeReject) diff --git a/backend/internal/domain/constants.go b/backend/internal/domain/constants.go index 148c1cd5d3..3640d1b82e 100644 --- a/backend/internal/domain/constants.go +++ b/backend/internal/domain/constants.go @@ -23,7 +23,27 @@ const ( PlatformGemini = "gemini" PlatformAntigravity = "antigravity" PlatformGrok = "grok" - PlatformComposite = "composite" + // 国产 OpenAI 兼容供应商(经 OpenAI 网关转发,按 Chat Completions 协议)。 + PlatformKimi = "kimi" // Kimi (月之暗面 / Moonshot) + PlatformZhipu = "zhipu" // 智谱 GLM (bigmodel) + PlatformDeepseek = "deepseek" // DeepSeek + PlatformComposite = "composite" +) + +// Account mode constants 区分国产供应商的「按量付费(余额)」与「Coding Plan」两种接入方式。 +// 存储于 credentials["account_mode"],决定 base_url 预设与额度监控方式。 +const ( + AccountModePayG = "payg" // 按量付费:消耗余额,做余额检测冷却 + AccountModeCoding = "coding" // Coding Plan:滚动用量窗口冷却(5h / weekly) +) + +// API protocol constants 国产供应商的上游 API 协议维度。存储于 +// credentials["api_protocol"],与 account_mode 正交:协议决定转发端点与格式, +// 模式决定额度监控方式。同协议请求零转换直通;跨协议组合才走转换链。 +const ( + APIProtocolChatCompletions = "chat_completions" // OpenAI Chat Completions(默认) + APIProtocolAnthropic = "anthropic" // 原生 Anthropic /v1/messages(适配 Claude Code) + APIProtocolResponses = "responses" // OpenAI Responses(仅 deepseek,适配 Codex) ) // Account type constants diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index 30fc6b9828..465c6a7816 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -1025,7 +1025,8 @@ func (h *AccountHandler) Update(c *gin.Context) { // 当前请求。探测错误仅记录日志,不向上下文传播:探测失败时标记保持缺失, // 网关会按"现状即证据"默认走 Responses。 func (h *AccountHandler) scheduleOpenAIResponsesProbe(account *service.Account) { - if account == nil || account.Platform != service.PlatformOpenAI || account.Type != service.AccountTypeAPIKey { + if account == nil || account.Type != service.AccountTypeAPIKey || + (account.Platform != service.PlatformOpenAI && !service.IsCNProvider(account.Platform)) { return } if h.accountTestService == nil { diff --git a/backend/internal/handler/admin/cn_provider_handler.go b/backend/internal/handler/admin/cn_provider_handler.go new file mode 100644 index 0000000000..7c8d67848a --- /dev/null +++ b/backend/internal/handler/admin/cn_provider_handler.go @@ -0,0 +1,69 @@ +package admin + +import ( + "strconv" + + "github.com/Wei-Shaw/sub2api/internal/pkg/response" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" +) + +// CNProviderHandler 暴露国产供应商(kimi/zhipu/deepseek)的额度与余额查询端点。 +// +// - GET /admin/cn-providers/accounts/:id/quota Coding Plan 滚动窗口用量(kimi/zhipu) +// - GET /admin/cn-providers/accounts/:id/balance payg 账号余额(kimi/deepseek) +// +// 智谱(zhipu)无余额端点,故同一账号仅 quota 或 balance 其一可用:服务端按账号 +// platform + account_mode 校验并返回明确错误(见 CNProvider*Service 的 load*Account)。 +type CNProviderHandler struct { + quotaService *service.CNProviderQuotaService + balanceService *service.CNProviderBalanceService +} + +func NewCNProviderHandler( + quotaService *service.CNProviderQuotaService, + balanceService *service.CNProviderBalanceService, +) *CNProviderHandler { + return &CNProviderHandler{ + quotaService: quotaService, + balanceService: balanceService, + } +} + +// QueryQuota 查询 Coding Plan 滚动窗口用量(5h + weekly)。 +func (h *CNProviderHandler) QueryQuota(c *gin.Context) { + accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + response.BadRequest(c, "Invalid account ID") + return + } + if h == nil || h.quotaService == nil { + response.BadRequest(c, "cn provider quota service is not enabled") + return + } + result, err := h.quotaService.QueryUsage(c.Request.Context(), accountID) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, result) +} + +// QueryBalance 查询 payg 账号余额。 +func (h *CNProviderHandler) QueryBalance(c *gin.Context) { + accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + response.BadRequest(c, "Invalid account ID") + return + } + if h == nil || h.balanceService == nil { + response.BadRequest(c, "cn provider balance service is not enabled") + return + } + result, err := h.balanceService.QueryBalance(c.Request.Context(), accountID) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, result) +} diff --git a/backend/internal/handler/admin/user_platform_quota_admin_test.go b/backend/internal/handler/admin/user_platform_quota_admin_test.go index 0211480b82..f1d1b5b605 100644 --- a/backend/internal/handler/admin/user_platform_quota_admin_test.go +++ b/backend/internal/handler/admin/user_platform_quota_admin_test.go @@ -112,12 +112,13 @@ func TestUpdateUserPlatformQuotas_Success(t *testing.T) { if len(repo.upsertCalls) != 1 { t.Fatalf("UpsertForUser should be called once, got %d", len(repo.upsertCalls)) } - if repo.upsertCalls[0].userID != 42 || len(repo.upsertCalls[0].records) != len(service.AllowedQuotaPlatforms) { + // upsert 记录数 = 请求体中给出的平台数(未给出的平台不落库)。 + if repo.upsertCalls[0].userID != 42 || len(repo.upsertCalls[0].records) != 5 { t.Errorf("unexpected upsert call: %+v", repo.upsertCalls[0]) } - // 缓存失效:按全部允许平台统一失效。 - if len(cache.deleteCalls) != 5 { - t.Errorf("expected 5 cache delete calls, got %d: %+v", len(cache.deleteCalls), cache.deleteCalls) + // 缓存失效:按全部允许平台统一失效(含 kimi/zhipu/deepseek)。 + if len(cache.deleteCalls) != len(service.AllowedQuotaPlatforms) { + t.Errorf("expected %d cache delete calls, got %d: %+v", len(service.AllowedQuotaPlatforms), len(cache.deleteCalls), cache.deleteCalls) } } diff --git a/backend/internal/handler/handler.go b/backend/internal/handler/handler.go index 359aff0071..7c0d3921eb 100644 --- a/backend/internal/handler/handler.go +++ b/backend/internal/handler/handler.go @@ -19,6 +19,7 @@ type AdminHandlers struct { GeminiOAuth *admin.GeminiOAuthHandler AntigravityOAuth *admin.AntigravityOAuthHandler GrokOAuth *admin.GrokOAuthHandler + CNProvider *admin.CNProviderHandler Proxy *admin.ProxyHandler Redeem *admin.RedeemHandler Promo *admin.PromoHandler diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 1695eccc1b..28515d98f0 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -164,13 +164,11 @@ func wrapUsageRecordTaskContext(parent context.Context, task service.UsageRecord func openAICompatibleRequestPlatform(ctx context.Context, apiKey *service.APIKey) string { if platform, ok := service.ResolvedTargetPlatformFromContext(ctx); ok { - if platform == service.PlatformGrok { - return service.PlatformGrok - } - return service.PlatformOpenAI + // 保留 grok 与国产供应商原值,其他归一为 openai(与调度器精确匹配语义一致)。 + return service.NormalizeOpenAICompatiblePlatform(platform) } - if apiKey != nil && apiKey.Group != nil && apiKey.Group.Platform == service.PlatformGrok { - return service.PlatformGrok + if apiKey != nil && apiKey.Group != nil { + return service.NormalizeOpenAICompatiblePlatform(apiKey.Group.Platform) } return service.PlatformOpenAI } @@ -203,7 +201,9 @@ func allowOpenAICompatibleMessagesDispatch(apiKey *service.APIKey) bool { } func openAICompatibleTextTargetAllowed(c *gin.Context, apiKey *service.APIKey, model string) bool { - return compositeTargetPlatformAllowed(c, apiKey, model, service.PlatformOpenAI, service.PlatformGrok) + return compositeTargetPlatformAllowed(c, apiKey, model, + service.PlatformOpenAI, service.PlatformGrok, + service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek) } // NewOpenAIGatewayHandler creates a new OpenAIGatewayHandler diff --git a/backend/internal/handler/wire.go b/backend/internal/handler/wire.go index 36fdb85280..1eef8c0019 100644 --- a/backend/internal/handler/wire.go +++ b/backend/internal/handler/wire.go @@ -23,6 +23,7 @@ func ProvideAdminHandlers( geminiOAuthHandler *admin.GeminiOAuthHandler, antigravityOAuthHandler *admin.AntigravityOAuthHandler, grokOAuthHandler *admin.GrokOAuthHandler, + cnProviderHandler *admin.CNProviderHandler, proxyHandler *admin.ProxyHandler, redeemHandler *admin.RedeemHandler, promoHandler *admin.PromoHandler, @@ -63,6 +64,7 @@ func ProvideAdminHandlers( GeminiOAuth: geminiOAuthHandler, AntigravityOAuth: antigravityOAuthHandler, GrokOAuth: grokOAuthHandler, + CNProvider: cnProviderHandler, Proxy: proxyHandler, Redeem: redeemHandler, Promo: promoHandler, @@ -253,6 +255,7 @@ var ProviderSet = wire.NewSet( admin.NewGeminiOAuthHandler, admin.NewAntigravityOAuthHandler, admin.NewGrokOAuthHandler, + admin.NewCNProviderHandler, admin.NewProxyHandler, admin.NewRedeemHandler, admin.NewPromoHandler, diff --git a/backend/internal/repository/user_platform_quota_repo_integration_test.go b/backend/internal/repository/user_platform_quota_repo_integration_test.go index 4a8d1b5896..1bd4b08786 100644 --- a/backend/internal/repository/user_platform_quota_repo_integration_test.go +++ b/backend/internal/repository/user_platform_quota_repo_integration_test.go @@ -98,6 +98,35 @@ func TestUserPlatformQuotaRepository_BulkInsertInitial_GrokAllowed(t *testing.T) require.InDelta(t, 9.0, *rec.DailyLimitUSD, 1e-9) } +// TestUserPlatformQuotaRepository_BulkInsertInitial_CNProvidersAllowed 回归迁移 224: +// kimi/zhipu/deepseek 平台必须能写入 user_platform_quotas(CHECK 约束已含国产供应商)。 +// 历史 bug:三个平台不在约束内 → 注册预填充 8 平台默认配额时整条多行 INSERT 中止 → +// fail-open 吞错 → 新用户拿到零条配额记录(缺失配额行 = 无限额)。 +func TestUserPlatformQuotaRepository_BulkInsertInitial_CNProvidersAllowed(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + txCtx := dbent.NewTxContext(ctx, tx) + client := tx.Client() + + userID := mustCreateUserForQuota(t, client) + repo := NewUserPlatformQuotaRepository(client) + + daily := 12.0 + records := []UserPlatformQuotaRecord{ + {UserID: userID, Platform: "kimi", DailyLimitUSD: &daily}, + {UserID: userID, Platform: "zhipu"}, + {UserID: userID, Platform: "deepseek"}, + } + require.NoError(t, repo.BulkInsertInitial(txCtx, records), + "kimi/zhipu/deepseek 平台应可写入(迁移 224 后 CHECK 约束已含国产供应商)") + + for _, platform := range []string{"kimi", "zhipu", "deepseek"} { + rec, err := repo.GetByUserPlatform(txCtx, userID, platform) + require.NoError(t, err) + require.NotNil(t, rec, "%s 配额行应已写入", platform) + } +} + func TestUserPlatformQuotaRepository_GetByUserPlatform(t *testing.T) { ctx := context.Background() tx := testEntTx(t) diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index 79dfb919a6..b0b1df7bc0 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -861,7 +861,7 @@ func TestAPIContracts(t *testing.T) { "force_email_on_third_party_signup": false, "default_concurrency": 5, "default_balance": 1.25, - "default_platform_quotas": {"anthropic":{"daily":null,"weekly":null,"monthly":null},"antigravity":{"daily":null,"weekly":null,"monthly":null},"gemini":{"daily":null,"weekly":null,"monthly":null},"grok":{"daily":null,"weekly":null,"monthly":null},"openai":{"daily":null,"weekly":null,"monthly":null}}, + "default_platform_quotas": {"anthropic":{"daily":null,"weekly":null,"monthly":null},"antigravity":{"daily":null,"weekly":null,"monthly":null},"deepseek":{"daily":null,"weekly":null,"monthly":null},"gemini":{"daily":null,"weekly":null,"monthly":null},"grok":{"daily":null,"weekly":null,"monthly":null},"kimi":{"daily":null,"weekly":null,"monthly":null},"openai":{"daily":null,"weekly":null,"monthly":null},"zhipu":{"daily":null,"weekly":null,"monthly":null}}, "auth_source_default_email_platform_quotas": null, "auth_source_default_github_platform_quotas": null, "auth_source_default_google_platform_quotas": null, @@ -1175,7 +1175,7 @@ func TestAPIContracts(t *testing.T) { "purchase_subscription_url": "", "table_default_page_size": 20, "table_page_size_options": [10, 20, 50], - "default_platform_quotas": {"anthropic":{"daily":null,"weekly":null,"monthly":null},"antigravity":{"daily":null,"weekly":null,"monthly":null},"gemini":{"daily":null,"weekly":null,"monthly":null},"grok":{"daily":null,"weekly":null,"monthly":null},"openai":{"daily":null,"weekly":null,"monthly":null}}, + "default_platform_quotas": {"anthropic":{"daily":null,"weekly":null,"monthly":null},"antigravity":{"daily":null,"weekly":null,"monthly":null},"deepseek":{"daily":null,"weekly":null,"monthly":null},"gemini":{"daily":null,"weekly":null,"monthly":null},"grok":{"daily":null,"weekly":null,"monthly":null},"kimi":{"daily":null,"weekly":null,"monthly":null},"openai":{"daily":null,"weekly":null,"monthly":null},"zhipu":{"daily":null,"weekly":null,"monthly":null}}, "auth_source_default_email_platform_quotas": null, "auth_source_default_github_platform_quotas": null, "auth_source_default_google_platform_quotas": null, diff --git a/backend/internal/server/routes/admin.go b/backend/internal/server/routes/admin.go index 149d4f95bc..60d6516bd0 100644 --- a/backend/internal/server/routes/admin.go +++ b/backend/internal/server/routes/admin.go @@ -58,6 +58,9 @@ func RegisterAdminRoutes( // Grok OAuth registerGrokOAuthRoutes(admin, h) + // 国产供应商(kimi/zhipu/deepseek)额度与余额 + registerCNProviderRoutes(admin, h) + // 代理管理 registerProxyRoutes(admin, h, stepUpAuth) @@ -482,6 +485,17 @@ func registerGrokOAuthRoutes(admin *gin.RouterGroup, h *handler.Handlers) { } } +// registerCNProviderRoutes 注册国产供应商(kimi/zhipu/deepseek)的额度与余额查询端点。 +func registerCNProviderRoutes(admin *gin.RouterGroup, h *handler.Handlers) { + cn := admin.Group("/cn-providers") + { + // Coding Plan 滚动窗口用量(kimi/zhipu coding 账号)。 + cn.GET("/accounts/:id/quota", h.Admin.CNProvider.QueryQuota) + // payg 账号余额(kimi/deepseek;zhipu 无余额端点)。 + cn.GET("/accounts/:id/balance", h.Admin.CNProvider.QueryBalance) + } +} + func registerProxyRoutes(admin *gin.RouterGroup, h *handler.Handlers, stepUpAuth middleware.StepUpAuthMiddleware) { proxies := admin.Group("/proxies") { diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index 6d99d2c7ea..3c4b33f519 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -47,7 +47,9 @@ func RegisterGatewayRoutes( isOpenAIResponsesCompatibleGatewayPlatform := func(c *gin.Context) bool { switch getGroupPlatform(c) { - case service.PlatformOpenAI, service.PlatformGrok: + case service.PlatformOpenAI, service.PlatformGrok, + service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek: + // 国产 OpenAI 兼容供应商(kimi/zhipu/deepseek)与 openai/grok 一样经 OpenAI 网关转发。 return true default: return false @@ -58,7 +60,7 @@ func RegisterGatewayRoutes( } countTokensHandler := func(c *gin.Context) { switch getGroupPlatform(c) { - case service.PlatformOpenAI: + case service.PlatformOpenAI, service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek: h.OpenAIGateway.CountTokens(c) case service.PlatformGrok: h.OpenAIGateway.GrokCountTokens(c) diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index cb74aaf50f..cd66ef89ce 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -271,8 +271,30 @@ func (a *Account) IsGrokOAuth() bool { return a.IsGrok() && a.Type == AccountTypeOAuth } +// IsKimi / IsZhipu / IsDeepseek 标识国产 OpenAI 兼容供应商账号。 +func (a *Account) IsKimi() bool { + return a.Platform == PlatformKimi +} + +func (a *Account) IsZhipu() bool { + return a.Platform == PlatformZhipu +} + +func (a *Account) IsDeepseek() bool { + return a.Platform == PlatformDeepseek +} + +// IsCNProvider 报告是否为国产 OpenAI 兼容供应商(kimi/zhipu/deepseek)。 +func (a *Account) IsCNProvider() bool { + return a != nil && IsCNProvider(a.Platform) +} + +// IsOpenAICompatible 报告账号是否走 OpenAI 网关(OpenAI 协议族)。 +// openai/grok 原生走 OpenAI 网关;kimi/zhipu/deepseek 同为 OpenAI Chat Completions +// 兼容上游,也经 OpenAI 网关转发。 func (a *Account) IsOpenAICompatible() bool { - return a != nil && (a.Platform == PlatformOpenAI || a.Platform == PlatformGrok) + return a != nil && (a.Platform == PlatformOpenAI || a.Platform == PlatformGrok || + a.Platform == PlatformKimi || a.Platform == PlatformZhipu || a.Platform == PlatformDeepseek) } func (a *Account) GeminiOAuthType() string { @@ -1281,17 +1303,161 @@ func (a *Account) IsOpenAIApiKey() bool { return a.IsOpenAI() && a.Type == AccountTypeAPIKey } +// GetOpenAIBaseURL 解析 OpenAI 协议族账号的上游 base_url。 +// 适用 openai 与国产 OpenAI 兼容供应商(kimi/zhipu/deepseek);grok 走 GetGrokBaseURL, +// 此处对 grok 返回 "" 以保持原有行为。 func (a *Account) GetOpenAIBaseURL() string { - if !a.IsOpenAI() { + if !a.IsOpenAI() && !a.IsCNProvider() { return "" } - if a.Type == AccountTypeAPIKey { - baseURL := a.GetCredential("base_url") - if baseURL != "" { + if a.Type == AccountTypeAPIKey || a.Type == AccountTypeUpstream { + if baseURL := strings.TrimSpace(a.GetCredential("base_url")); baseURL != "" { return baseURL } } - return "https://api.openai.com" + // 平台默认 base_url:CN 供应商按 account_mode 选择 payg / coding 默认值。 + switch a.Platform { + case PlatformKimi: + if a.GetAccountMode() == AccountModeCoding { + return DefaultKimiCodingBaseURL + } + return DefaultKimiPayGBaseURL + case PlatformZhipu: + if a.GetAccountMode() == AccountModeCoding { + return DefaultZhipuCodingBaseURL + } + return DefaultZhipuPayGBaseURL + case PlatformDeepseek: + return DefaultDeepseekBaseURL + default: + return "https://api.openai.com" + } +} + +// GetAccountMode 返回国产供应商账号的接入模式(payg / coding);非国产供应商或未设置时 +// 返回空串。存储于 credentials["account_mode"]。 +func (a *Account) GetAccountMode() string { + if a == nil { + return "" + } + mode := strings.TrimSpace(a.GetCredential("account_mode")) + if mode == AccountModePayG || mode == AccountModeCoding { + return mode + } + return "" +} + +// IsCodingPlan 报告账号是否为 Coding Plan 模式(用于滚动用量窗口冷却)。 +func (a *Account) IsCodingPlan() bool { + return a.GetAccountMode() == AccountModeCoding +} + +// GetAPIProtocol 返回国产供应商账号的上游 API 协议。存储于 +// credentials["api_protocol"];缺失或与平台不匹配时回退 chat_completions +// (与既有行为完全一致)。responses 协议仅 deepseek 支持(官方原生 /responses +// 端点,适配 Codex);kimi/zhipu 无此端点。 +func (a *Account) GetAPIProtocol() string { + if a == nil || !a.IsCNProvider() { + return APIProtocolChatCompletions + } + switch strings.TrimSpace(a.GetCredential("api_protocol")) { + case APIProtocolAnthropic: + return APIProtocolAnthropic + case APIProtocolResponses: + if a.Platform == PlatformDeepseek { + return APIProtocolResponses + } + case APIProtocolChatCompletions: + return APIProtocolChatCompletions + } + return APIProtocolChatCompletions +} + +// IsAnthropicProtocol 报告账号是否以原生 Anthropic 协议接入上游 +// (/v1/messages 直通,适配 Claude Code 等客户端)。 +func (a *Account) IsAnthropicProtocol() bool { + return a.GetAPIProtocol() == APIProtocolAnthropic +} + +// GetAnthropicProtocolBaseURL 返回 Anthropic 协议账号的上游 base_url +// (上游路径为 {base}/v1/messages)。优先取凭证 base_url,缺失时按 +// 供应商 × 接入模式返回默认端点。非 Anthropic 协议账号返回空串。 +func (a *Account) GetAnthropicProtocolBaseURL() string { + if a == nil || !a.IsAnthropicProtocol() { + return "" + } + if a.Type == AccountTypeAPIKey || a.Type == AccountTypeUpstream { + if baseURL := strings.TrimSpace(a.GetCredential("base_url")); baseURL != "" { + return baseURL + } + } + switch a.Platform { + case PlatformKimi: + if a.GetAccountMode() == AccountModeCoding { + return DefaultKimiCodingAnthropicBaseURL + } + return DefaultKimiPayGAnthropicBaseURL + case PlatformZhipu: + return DefaultZhipuAnthropicBaseURL + case PlatformDeepseek: + return DefaultDeepseekAnthropicBaseURL + default: + return "" + } +} + +// GetOpenAIFormatBaseURL 返回供 OpenAI 格式端点(/v1/models、/v1/chat/completions +// 等)使用的 base。chat_completions / responses 协议下与 GetOpenAIBaseURL +// 一致(凭证 base_url 或平台默认);anthropic 协议下凭证 base_url 指向 Anthropic +// 端点,不能拿来拼 OpenAI 路径,此时返回该供应商 × 模式的 Chat Completions +// 默认 base(模型同步等协议族共用路径仍可用)。 +func (a *Account) GetOpenAIFormatBaseURL() string { + if a == nil || !a.IsAnthropicProtocol() { + return a.GetOpenAIBaseURL() + } + switch a.Platform { + case PlatformKimi: + if a.GetAccountMode() == AccountModeCoding { + return DefaultKimiCodingBaseURL + } + return DefaultKimiPayGBaseURL + case PlatformZhipu: + if a.GetAccountMode() == AccountModeCoding { + return DefaultZhipuCodingBaseURL + } + return DefaultZhipuPayGBaseURL + case PlatformDeepseek: + return DefaultDeepseekBaseURL + default: + return a.GetOpenAIBaseURL() + } +} + +// GetCNAPIKey 返回国产 OpenAI 兼容供应商账号的 api_key 凭据(kimi/zhipu/deepseek)。 +// 与 openai 的 GetOpenAIApiKey 区分:后者仅对 openai 平台返回。 +func (a *Account) GetCNAPIKey() string { + if a == nil || !a.IsCNProvider() { + return "" + } + return a.GetCredential("api_key") +} + +// GetCodingPlanProvider 根据 base_url 识别 Coding Plan 供应商(kimi / zhipu), +// 用于路由到对应的额度查询端点。非 coding 模式或无法识别时返回空串。 +// 判定规则与 cc-switch coding_plan.rs::detect_provider 保持一致。 +func (a *Account) GetCodingPlanProvider() string { + if a == nil || a.GetAccountMode() != AccountModeCoding { + return "" + } + baseURL := strings.ToLower(a.GetOpenAIBaseURL()) + switch { + case strings.Contains(baseURL, "api.kimi.com/coding"): + return PlatformKimi + case strings.Contains(baseURL, "bigmodel.cn"), strings.Contains(baseURL, "api.z.ai"): + return PlatformZhipu + default: + return "" + } } func (a *Account) GetOpenAIAccessToken() string { @@ -1404,6 +1570,23 @@ func (a *Account) GetOpenAIApiKey() string { return a.GetCredential("api_key") } +// GetOpenAIProtocolAPIKey 返回 OpenAI 协议族 APIKey 账号的密钥。 +// 覆盖 openai 原生账号与国产 OpenAI 兼容供应商(kimi/zhipu/deepseek)账号, +// 供转发鉴权、模型列表同步等协议族共用路径使用。注意 IsOpenAIApiKey 语义上 +// 仅指 openai 平台账号,调度倍率/WS 能力门控继续以其为准,不受本方法影响。 +func (a *Account) GetOpenAIProtocolAPIKey() string { + if a == nil { + return "" + } + if a.IsCNProvider() { + if a.Type != AccountTypeAPIKey { + return "" + } + return a.GetCredential("api_key") + } + return a.GetOpenAIApiKey() +} + func (a *Account) GetOpenAIUserAgent() string { if !a.IsOpenAI() { return "" diff --git a/backend/internal/service/account_scheduling_threshold_eval.go b/backend/internal/service/account_scheduling_threshold_eval.go index b8a9dae80d..7dde1fd163 100644 --- a/backend/internal/service/account_scheduling_threshold_eval.go +++ b/backend/internal/service/account_scheduling_threshold_eval.go @@ -58,6 +58,10 @@ func EvaluateAccountSchedulingThreshold(account *Account, thresholds map[string] winner = pickLatestResetSchedulingCandidate(anthropicThresholdCandidates(account), threshold, now) case PlatformGrok: winner = pickLatestResetSchedulingCandidate(grokThresholdCandidates(account), threshold, now) + case PlatformKimi: + winner = pickLatestResetSchedulingCandidate(cnProviderThresholdCandidates(account, PlatformKimi), threshold, now) + case PlatformZhipu: + winner = pickLatestResetSchedulingCandidate(cnProviderThresholdCandidates(account, PlatformZhipu), threshold, now) default: return decision } @@ -303,6 +307,45 @@ func grokThresholdCandidates(account *Account) []*accountSchedulingThresholdCand } } +// cnProviderThresholdCandidates 读取国产供应商 Coding Plan 账号的 5h / weekly 滚动窗口 +// 用量快照(由 CNProviderQuotaService 写入 account.Extra,键形如 +// _5h_used_percent / _weekly_reset_at)。payg 账号无此快照, +// 候选为空 → 不触发阈值停调(余额型走余额检测)。与 openai 的快照驱动停调一致: +// 仅当用量超阈值且窗口尚未重置时才停调。 +func cnProviderThresholdCandidates(account *Account, provider string) []*accountSchedulingThresholdCandidate { + if account == nil || len(account.Extra) == 0 { + return nil + } + return []*accountSchedulingThresholdCandidate{ + cnThresholdCandidate(account.Extra, provider, "5h"), + cnThresholdCandidate(account.Extra, provider, "weekly"), + } +} + +func cnThresholdCandidate(extra map[string]any, provider, window string) *accountSchedulingThresholdCandidate { + var usedKey, resetKey string + switch window { + case "5h": + usedKey = cnExtraKey(provider, cnExtraSuffix5hUsed) + resetKey = cnExtraKey(provider, cnExtraSuffix5hReset) + case "weekly": + usedKey = cnExtraKey(provider, cnExtraSuffixWeeklyUsed) + resetKey = cnExtraKey(provider, cnExtraSuffixWeeklyReset) + default: + return nil + } + usedPercent, ok := extra[usedKey] + if !ok { + return nil + } + return &accountSchedulingThresholdCandidate{ + window: window, + scope: provider, + usedPercent: schedulingPercentValue(usedPercent), + until: parseSchedulingResetAt(extra[resetKey]), + } +} + func pickLatestResetSchedulingCandidate(candidates []*accountSchedulingThresholdCandidate, threshold int, now time.Time) *accountSchedulingThresholdCandidate { var winner *accountSchedulingThresholdCandidate for _, candidate := range candidates { diff --git a/backend/internal/service/account_service.go b/backend/internal/service/account_service.go index 65d4d3f268..be199cb256 100644 --- a/backend/internal/service/account_service.go +++ b/backend/internal/service/account_service.go @@ -509,6 +509,9 @@ func (s *AccountService) TestCredentials(ctx context.Context, id int64) error { case PlatformGrok: // Grok OAuth credentials are validated via token exchange/refresh and request-path probes. return nil + case PlatformKimi, PlatformZhipu, PlatformDeepseek: + // 国产 OpenAI 兼容供应商:凭证为 API Key,实际可用性经余额/额度探测与转发路径验证。 + return nil default: return fmt.Errorf("unsupported platform: %s", account.Platform) } diff --git a/backend/internal/service/anthropic_apikey_auth.go b/backend/internal/service/anthropic_apikey_auth.go index 7752d6fe3a..044d5be96c 100644 --- a/backend/internal/service/anthropic_apikey_auth.go +++ b/backend/internal/service/anthropic_apikey_auth.go @@ -14,9 +14,14 @@ const ( // GetAnthropicAPIKeyAuthScheme returns the upstream authentication scheme for // Anthropic API-key accounts. Missing or invalid values keep the historical -// x-api-key behavior. +// x-api-key behavior. CN providers using their native Anthropic endpoints +// (api_protocol=anthropic) share the same override knob — Kimi/DeepSeek default +// to x-api-key, Zhipu can opt into Authorization: Bearer. func (a *Account) GetAnthropicAPIKeyAuthScheme() string { - if a == nil || a.Platform != PlatformAnthropic || a.Type != AccountTypeAPIKey { + if a == nil || a.Type != AccountTypeAPIKey { + return AnthropicAPIKeyAuthSchemeXAPIKey + } + if a.Platform != PlatformAnthropic && !a.IsCNProvider() { return AnthropicAPIKeyAuthSchemeXAPIKey } diff --git a/backend/internal/service/bedrock_stream.go b/backend/internal/service/bedrock_stream.go index 98196d27ec..9ce6b87e24 100644 --- a/backend/internal/service/bedrock_stream.go +++ b/backend/internal/service/bedrock_stream.go @@ -139,7 +139,7 @@ func (s *GatewayService) handleBedrockStreamingResponse( sseData = transformBedrockInvocationMetrics(sseData) // 解析 SSE 事件数据提取 usage - s.parseSSEUsagePassthrough(string(sseData), usage) + parseSSEUsagePassthrough(string(sseData), usage) // 确定 SSE event type eventType := gjson.GetBytes(sseData, "type").String() diff --git a/backend/internal/service/cn_provider_balance_check_service.go b/backend/internal/service/cn_provider_balance_check_service.go new file mode 100644 index 0000000000..d5ddd6712e --- /dev/null +++ b/backend/internal/service/cn_provider_balance_check_service.go @@ -0,0 +1,279 @@ +package service + +import ( + "context" + "fmt" + "log" + "strings" + "sync" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" +) + +// cnQuotaProber 抽象额度探测(*CNProviderQuotaService 实现,测试可替换)。 +type cnQuotaProber interface { + QueryUsage(ctx context.Context, accountID int64) (*CNProviderQuotaProbeResult, error) +} + +// cnQuotaProbeConcurrency 周期任务并发探测额度账号的并发度。 +const cnQuotaProbeConcurrency = 4 + +// CNProviderBalanceCheckService 周期性探测国产供应商账号: +// - payg(按量付费):余额低于阈值则临时停调,恢复则清除(仅清除本服务写入的停调); +// - coding plan:调用 CNProviderQuotaService 探测 5h/weekly 滚动窗口并落 extra 快照, +// 调度阈值评估(cnProviderThresholdCandidates)据此自动停调/恢复。 +// +// 克隆自 AccountExpiryService 的 Start/Stop/runOnce + ticker 骨架。 +// 余额探测仅覆盖有公开余额端点的 kimi / deepseek;智谱无余额端点,仅靠响应式 429/402。 +// 额度探测覆盖 kimi / zhipu 的 coding plan 账号(deepseek 无 coding 套餐)。 +type CNProviderBalanceCheckService struct { + accountRepo AccountRepository + balanceService *CNProviderBalanceService + quotaService cnQuotaProber + cfg *config.Config + interval time.Duration + stopCh chan struct{} + stopOnce sync.Once + wg sync.WaitGroup +} + +// NewCNProviderBalanceCheckService 构造周期余额/额度检测服务。 +// interval <= 0 时 Start() 直接返回(不启动),便于通过配置关闭。 +func NewCNProviderBalanceCheckService( + accountRepo AccountRepository, + balanceService *CNProviderBalanceService, + quotaService *CNProviderQuotaService, + cfg *config.Config, + interval time.Duration, +) *CNProviderBalanceCheckService { + return &CNProviderBalanceCheckService{ + accountRepo: accountRepo, + balanceService: balanceService, + quotaService: quotaService, + cfg: cfg, + interval: interval, + stopCh: make(chan struct{}), + } +} + +func (s *CNProviderBalanceCheckService) Start() { + if s == nil || s.accountRepo == nil || s.balanceService == nil || s.cfg == nil { + return + } + if !s.cfg.Gateway.CNProviders.BalanceCheckEnabled { + return + } + if s.interval <= 0 { + return + } + log.Printf("[CNBalance] started (interval=%s threshold=%.2f)", s.interval, s.cfg.Gateway.CNProviders.BalanceThreshold) + s.wg.Add(1) + go func() { + defer s.wg.Done() + ticker := time.NewTicker(s.interval) + defer ticker.Stop() + + // 启动后先等待一个周期再首次探测,避免与进程启动峰重叠。 + for { + select { + case <-ticker.C: + s.runOnce() + case <-s.stopCh: + return + } + } + }() +} + +func (s *CNProviderBalanceCheckService) Stop() { + if s == nil { + return + } + s.stopOnce.Do(func() { + close(s.stopCh) + }) + s.wg.Wait() +} + +func (s *CNProviderBalanceCheckService) runOnce() { + // 收集 coding 探测目标(kimi/deepseek + 智谱)与 payg 检查队列。 + // coding 探测统一在收集完成后按 4 并发执行:单账号探测 15-20s,串行 × + // 多账号会耗尽整体预算(120s 上限),排在后面的账号快照会饥饿, + // 连锁影响阈值停调的新鲜度判定。 + type quotaTarget struct { + id int64 + platform string + } + var quotaTargets []quotaTarget + var paygTargets []*Account + collect := func(platform string, accounts []Account) { + for i := range accounts { + account := &accounts[i] + if !account.IsActive() { + continue + } + // coding 账号:探测滚动窗口并落快照(不要求 Schedulable——已被 + // 阈值停调的账号也需要新鲜快照决定是否续停)。 + if account.IsCodingPlan() { + quotaTargets = append(quotaTargets, quotaTarget{id: account.ID, platform: account.Platform}) + continue + } + // payg 余额探测仅 kimi/deepseek(智谱无公开余额端点,payg 账号 + // 依赖响应式 402/429 处理)。 + if platform != PlatformZhipu && account.Schedulable { + paygTargets = append(paygTargets, account) + } + } + } + for _, platform := range s.platforms() { + accounts, err := s.accountRepo.ListByPlatform(context.Background(), platform) + if err != nil { + log.Printf("[CNBalance] list %s accounts failed: %v", platform, err) + continue + } + collect(platform, accounts) + } + // 智谱无余额端点,仅进额度探测。 + if s.quotaService != nil { + accounts, err := s.accountRepo.ListByPlatform(context.Background(), PlatformZhipu) + if err != nil { + log.Printf("[CNBalance] list %s accounts failed: %v", PlatformZhipu, err) + } else { + collect(PlatformZhipu, accounts) + } + } + + // 预算按工作量放大:4 并发 × 15s/批 + payg 每账号 5s,下限 30s 上限 300s。 + batches := (len(quotaTargets) + cnQuotaProbeConcurrency - 1) / cnQuotaProbeConcurrency + timeout := 30*time.Second + time.Duration(batches)*15*time.Second + time.Duration(len(paygTargets))*5*time.Second + if timeout > 300*time.Second { + timeout = 300 * time.Second + } + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + + threshold := s.cfg.Gateway.CNProviders.BalanceThreshold + paused, cleared := 0, 0 + for _, account := range paygTargets { + switch s.checkOne(ctx, account, threshold) { + case cnBalancePaused: + paused++ + case cnBalanceCleared: + cleared++ + } + } + + if len(quotaTargets) > 0 && s.quotaService != nil { + sem := make(chan struct{}, cnQuotaProbeConcurrency) + var wg sync.WaitGroup + for _, target := range quotaTargets { + wg.Add(1) + go func(t quotaTarget) { + defer wg.Done() + sem <- struct{}{} + defer func() { <-sem }() + s.probeQuota(ctx, t.id, t.platform) + }(target) + } + wg.Wait() + } + + if paused > 0 || cleared > 0 { + log.Printf("[CNBalance] paused=%d cleared=%d (threshold=%.2f)", paused, cleared, threshold) + } +} + +// probeQuota 探测单个 coding plan 账号的滚动窗口用量并落 extra 快照。 +// 不在此处做停调/恢复决策:调度阈值评估读取快照统一判定(含暂停账号的续停)。 +func (s *CNProviderBalanceCheckService) probeQuota(ctx context.Context, accountID int64, platform string) { + if s.quotaService == nil { + return + } + result, err := s.quotaService.QueryUsage(ctx, accountID) + if err != nil { + log.Printf("[CNBalance] quota probe account %d (%s) failed: %v", accountID, platform, err) + return + } + if result != nil && !result.Success && result.Error != "" { + log.Printf("[CNBalance] quota probe account %d (%s) error: %s", accountID, platform, result.Error) + } +} + +type cnBalanceCheckOutcome int + +const ( + cnBalanceNoChange cnBalanceCheckOutcome = iota + cnBalancePaused + cnBalanceCleared +) + +// checkOne 探测单账号余额并决定停调/恢复。探测失败时不动现状(避免瞬时网络抖动 +// 误解除或误停调)。 +func (s *CNProviderBalanceCheckService) checkOne(ctx context.Context, account *Account, threshold float64) cnBalanceCheckOutcome { + result, err := s.balanceService.QueryBalance(ctx, account.ID) + if err != nil || result == nil || !result.Success { + return cnBalanceNoChange + } + + // 双币种(deepseek CNY+USD)任一币种余额达标即可继续调度;仅当全部低于 + // 阈值(或不可用)才停调。 + low := !result.Available || allCNBalancesBelowThreshold(result, threshold) + if low { + // 已被(任何来源)停调时不覆盖其 reason。 + if !account.IsSchedulable() { + return cnBalanceNoChange + } + reason := cnBalanceLowReason(fmt.Sprintf("余额 %.4g %s 低于阈值 %.2f", result.Balance, result.Currency, threshold)) + if err := s.accountRepo.SetTempUnschedulable(ctx, account.ID, time.Now().Add(s.cooldown()), reason); err != nil { + log.Printf("[CNBalance] pause account %d failed: %v", account.ID, err) + return cnBalanceNoChange + } + log.Printf("[CNBalance] paused account %d (%s): balance=%.4g %s", account.ID, account.Platform, result.Balance, result.Currency) + return cnBalancePaused + } + + // 余额健康:仅清除「本服务写入」的临时停调(reason 前缀匹配),不触碰其他子系统。 + if account.TempUnschedulableUntil != nil && strings.HasPrefix(account.TempUnschedulableReason, cnBalanceLowReasonPrefix) { + if err := s.accountRepo.ClearTempUnschedulable(ctx, account.ID); err != nil { + log.Printf("[CNBalance] clear account %d failed: %v", account.ID, err) + return cnBalanceNoChange + } + log.Printf("[CNBalance] reactivated account %d (%s): balance=%.4g %s", account.ID, account.Platform, result.Balance, result.Currency) + return cnBalanceCleared + } + return cnBalanceNoChange +} + +func (s *CNProviderBalanceCheckService) platforms() []string { + return []string{PlatformKimi, PlatformDeepseek} +} + +// allCNBalancesBelowThreshold 判断全部币种余额是否均低于阈值。 +// 无明细时退回主币种判定(与旧行为一致)。 +func allCNBalancesBelowThreshold(result *CNProviderBalanceResult, threshold float64) bool { + if len(result.Balances) == 0 { + return result.Balance < threshold + } + for _, entry := range result.Balances { + if entry.Balance >= threshold { + return false + } + } + return true +} + +// cooldown 返回临时停调持续时长(= 2× 检测周期),与响应式 402/429 路径一致。 +func (s *CNProviderBalanceCheckService) cooldown() time.Duration { + minutes := 10 + if s.cfg != nil { + if cfgMin := s.cfg.Gateway.CNProviders.BalanceCheckIntervalMinutes; cfgMin > 0 { + minutes = cfgMin + } + } + cooldown := time.Duration(minutes) * time.Minute * 2 + if cooldown < time.Minute { + cooldown = 10 * time.Minute + } + return cooldown +} diff --git a/backend/internal/service/cn_provider_balance_check_service_test.go b/backend/internal/service/cn_provider_balance_check_service_test.go new file mode 100644 index 0000000000..26f7b69ea2 --- /dev/null +++ b/backend/internal/service/cn_provider_balance_check_service_test.go @@ -0,0 +1,101 @@ +package service + +import ( + "context" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +// 周期任务对 coding plan 账号的额度探测行为(runOnce 集成路径): +// - kimi coding 账号(含已被阈值停调的)→ 额度探测被调用; +// - 智谱 coding 账号 → 额度探测被调用(智谱不进 kimi/deepseek 余额循环); +// - payg 账号不经过额度探测(走余额路径,本测试不放 payg 账号避免真实网络); +// - 非激活账号完全跳过。 + +type fakeCNQuotaProber struct { + probed []int64 +} + +func (f *fakeCNQuotaProber) QueryUsage(ctx context.Context, accountID int64) (*CNProviderQuotaProbeResult, error) { + f.probed = append(f.probed, accountID) + return &CNProviderQuotaProbeResult{Success: true, Persisted: true}, nil +} + +type fakeCNCheckRepo struct { + AccountRepository + byPlatform map[string][]Account +} + +func (r *fakeCNCheckRepo) ListByPlatform(ctx context.Context, platform string) ([]Account, error) { + return r.byPlatform[platform], nil +} + +func TestCNProviderBalanceCheckRunOnceProbesCodingPlanQuota(t *testing.T) { + kimiActive := Account{ID: 1, Platform: PlatformKimi, Type: AccountTypeAPIKey, Status: StatusActive, + Credentials: map[string]any{"account_mode": "coding"}} + // 已被阈值停调的 coding 账号也要刷新快照(决定是否续停)。 + kimiPaused := Account{ID: 2, Platform: PlatformKimi, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: false, + Credentials: map[string]any{"account_mode": "coding"}} + // 非激活账号跳过。 + kimiInactive := Account{ID: 3, Platform: PlatformKimi, Type: AccountTypeAPIKey, Status: StatusDisabled, + Credentials: map[string]any{"account_mode": "coding"}} + zhipuCoding := Account{ID: 4, Platform: PlatformZhipu, Type: AccountTypeAPIKey, Status: StatusActive, + Credentials: map[string]any{"account_mode": "coding"}} + + repo := &fakeCNCheckRepo{byPlatform: map[string][]Account{ + PlatformKimi: {kimiActive, kimiPaused, kimiInactive}, + PlatformZhipu: {zhipuCoding}, + }} + prober := &fakeCNQuotaProber{} + svc := &CNProviderBalanceCheckService{ + accountRepo: repo, + quotaService: prober, + cfg: &config.Config{}, + } + + svc.runOnce() + + require.ElementsMatch(t, []int64{1, 2, 4}, prober.probed) +} + +// runOnceZhipuQuota 在 quotaService 缺失时安全跳过(Start 门控不启动的老部署路径)。 +func TestCNProviderBalanceCheckRunOnceWithoutQuotaService(t *testing.T) { + repo := &fakeCNCheckRepo{byPlatform: map[string][]Account{ + PlatformZhipu: {{ID: 4, Platform: PlatformZhipu, Type: AccountTypeAPIKey, Status: StatusActive, + Credentials: map[string]any{"account_mode": "coding"}}}, + }} + svc := &CNProviderBalanceCheckService{accountRepo: repo, cfg: &config.Config{}} + require.NotPanics(t, func() { svc.runOnce() }) +} + +// 双币种(deepseek CNY+USD)停调判定:任一币种达标即不停调,全部低于阈值才停; +// 无明细时退回主币种(兼容旧结果)。 +func TestAllCNBalancesBelowThreshold(t *testing.T) { + dualLow := &CNProviderBalanceResult{ + Balance: 1.0, + Currency: "CNY", + Balances: []CNProviderBalanceEntry{ + {Currency: "CNY", Balance: 1.0}, + {Currency: "USD", Balance: 0.5}, + }, + } + require.True(t, allCNBalancesBelowThreshold(dualLow, 5.0)) + + dualMixed := &CNProviderBalanceResult{ + Balance: 1.0, + Currency: "CNY", + Balances: []CNProviderBalanceEntry{ + {Currency: "CNY", Balance: 1.0}, + {Currency: "USD", Balance: 20.0}, + }, + } + require.False(t, allCNBalancesBelowThreshold(dualMixed, 5.0)) + + // 无明细:按主币种判定(旧行为)。 + singleLow := &CNProviderBalanceResult{Balance: 1.0, Currency: "CNY"} + require.True(t, allCNBalancesBelowThreshold(singleLow, 5.0)) + singleOK := &CNProviderBalanceResult{Balance: 10.0, Currency: "CNY"} + require.False(t, allCNBalancesBelowThreshold(singleOK, 5.0)) +} diff --git a/backend/internal/service/cn_provider_balance_service.go b/backend/internal/service/cn_provider_balance_service.go new file mode 100644 index 0000000000..d01c9b6d08 --- /dev/null +++ b/backend/internal/service/cn_provider_balance_service.go @@ -0,0 +1,277 @@ +package service + +import ( + "context" + "fmt" + "io" + "log/slog" + "net/http" + "strconv" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/tidwall/gjson" + "golang.org/x/sync/singleflight" +) + +// 国产供应商 payg(按量付费)账号余额探测服务。 +// +// 仅覆盖有公开余额端点的供应商: +// - Kimi/Moonshot:GET https://api.moonshot.cn/v1/users/me/balance (Bearer) → data.available_balance +// - DeepSeek: GET https://api.deepseek.com/user/balance (Bearer) → balance_infos[].total_balance + is_available +// +// 智谱(zhipu)无公开余额端点(OpenAPI 规格验证),仅靠响应式 429/402(见 +// ratelimit_cn_providers.go)。解析逻辑对齐 cc-switch services/balance.rs::query_deepseek。 +const ( + cnBalanceUpstreamTimeout = 15 * time.Second + cnBalanceMaxBodyBytes = 256 * 1024 + + // Extra 余额快照键后缀(加 provider 前缀)。 + cnBalanceExtraSuffixBalance = "balance" + cnBalanceExtraSuffixCurrency = "balance_currency" + cnBalanceExtraSuffixAvailable = "balance_available" // deepseek is_available 健康标记 + cnBalanceExtraSuffixUpdated = "balance_updated_at" + cnBalanceExtraSuffixBalances = "balances" // 多币种明细(deepseek USD+CNY) +) + +// CNProviderBalanceEntry 是单一币种的余额明细。 +type CNProviderBalanceEntry struct { + Currency string `json:"currency"` + Balance float64 `json:"balance"` +} + +// CNProviderBalanceResult 是余额探测的返回结构(管理端 + UI 消费)。 +type CNProviderBalanceResult struct { + Provider string `json:"provider"` + Success bool `json:"success"` + // Balance/Currency 为主币种(balance_infos 首条,兼容单币种消费方); + // 完整明细见 Balances(deepseek 双币种账号含 CNY + USD 两条)。 + Balance float64 `json:"balance"` + Currency string `json:"currency,omitempty"` + Balances []CNProviderBalanceEntry `json:"balances,omitempty"` + Available bool `json:"available"` // 健康标记(deepseek is_available;kimi 无此概念恒 true) + StatusCode int `json:"status_code,omitempty"` + FetchedAt int64 `json:"fetched_at"` + Persisted bool `json:"persisted"` + Error string `json:"error,omitempty"` +} + +// CNProviderBalanceService 探测 Kimi / DeepSeek payg 账号的账户余额。 +type CNProviderBalanceService struct { + accountRepo AccountRepository + proxyRepo ProxyRepository + httpUpstream HTTPUpstream + cfg *config.Config + flight singleflight.Group +} + +// NewCNProviderBalanceService 构造余额探测服务。 +func NewCNProviderBalanceService( + accountRepo AccountRepository, + proxyRepo ProxyRepository, + httpUpstream HTTPUpstream, + cfg *config.Config, +) *CNProviderBalanceService { + return &CNProviderBalanceService{ + accountRepo: accountRepo, + proxyRepo: proxyRepo, + httpUpstream: httpUpstream, + cfg: cfg, + } +} + +// QueryBalance 探测指定 payg 账号的余额并落 Extra 快照。 +func (s *CNProviderBalanceService) QueryBalance(ctx context.Context, accountID int64) (*CNProviderBalanceResult, error) { + if s == nil || s.accountRepo == nil || s.httpUpstream == nil { + return nil, infraerrors.New(http.StatusInternalServerError, "CN_BALANCE_NOT_CONFIGURED", "cn provider balance service is not configured") + } + key := "cn_balance:" + strconv.FormatInt(accountID, 10) + resultCh := s.flight.DoChan(key, func() (any, error) { + probeCtx, cancel := context.WithTimeout(context.Background(), cnBalanceUpstreamTimeout+5*time.Second) + defer cancel() + return s.queryBalance(probeCtx, accountID) + }) + select { + case <-ctx.Done(): + return nil, ctx.Err() + case flightResult := <-resultCh: + if flightResult.Err != nil { + return nil, flightResult.Err + } + result, ok := flightResult.Val.(*CNProviderBalanceResult) + if !ok || result == nil { + return nil, infraerrors.New(http.StatusInternalServerError, "CN_BALANCE_PROBE_RESULT_INVALID", "invalid cn provider balance probe result") + } + cloned := *result + return &cloned, nil + } +} + +func (s *CNProviderBalanceService) queryBalance(ctx context.Context, accountID int64) (*CNProviderBalanceResult, error) { + account, err := s.loadPayGAccount(ctx, accountID) + if err != nil { + return nil, err + } + provider := account.Platform + if provider != PlatformKimi && provider != PlatformDeepseek { + return nil, infraerrors.New(http.StatusBadRequest, "CN_BALANCE_NO_ENDPOINT", "account provider has no balance endpoint") + } + + apiKey := strings.TrimSpace(account.GetCNAPIKey()) + if apiKey == "" { + return nil, infraerrors.New(http.StatusBadRequest, "CN_BALANCE_NO_APIKEY", "account api_key is empty") + } + + targetURL := cnBalanceURL(account) + // 探测发起前过出站 URL 安全策略(与网关转发/Grok 探测同一套校验): + // DeepSeek 端点由账号 base_url 衍生,不得把 API key 发往策略外主机。 + validatedURL, err := cnValidateProbeURL(s.cfg, targetURL) + if err != nil { + return nil, infraerrors.New(http.StatusForbidden, "CN_BALANCE_URL_REJECTED", err.Error()) + } + targetURL = validatedURL + proxyURL := s.resolveProxyURL(ctx, account) + callCtx, cancel := context.WithTimeout(ctx, cnBalanceUpstreamTimeout) + defer cancel() + req, err := http.NewRequestWithContext(callCtx, http.MethodGet, targetURL, nil) + if err != nil { + return nil, infraerrors.Newf(http.StatusInternalServerError, "CN_BALANCE_REQUEST_BUILD_FAILED", "build request: %v", err) + } + req.Header.Set("Authorization", "Bearer "+apiKey) + req.Header.Set("Accept", "application/json") + account.ApplyHeaderOverrides(req.Header) + + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, maxInt(account.Concurrency, 1)) + if err != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "CN_BALANCE_REQUEST_FAILED", "upstream request failed: %v", err) + } + defer func() { _ = resp.Body.Close() }() + bodyBytes, _ := io.ReadAll(io.LimitReader(resp.Body, cnBalanceMaxBodyBytes)) + + now := time.Now().UTC() + result := &CNProviderBalanceResult{ + Provider: provider, + FetchedAt: now.Unix(), + StatusCode: resp.StatusCode, + Available: true, + } + + if resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden { + result.Error = fmt.Sprintf("Authentication failed (HTTP %d)", resp.StatusCode) + return result, nil + } + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + result.Error = fmt.Sprintf("API error (HTTP %d): %s", resp.StatusCode, truncate(strings.TrimSpace(string(bodyBytes)), 240)) + return result, nil + } + + var entries []CNProviderBalanceEntry + available := true + switch provider { + case PlatformKimi: + // Moonshot:code==0 成功;data.available_balance(number),单币种 CNY。 + balance, _ := cnParseF64(gjson.GetBytes(bodyBytes, "data.available_balance").Value()) + entries = append(entries, CNProviderBalanceEntry{Currency: "CNY", Balance: balance}) + case PlatformDeepseek: + // is_available 缺省视为 true(健康);显式存在时取其值。 + if v := gjson.GetBytes(bodyBytes, "is_available"); v.Exists() { + available = v.Bool() + } + // balance_infos 逐条解析:双币种账号同时返回 CNY + USD(数组顺序即 + // 主次序,首条为主币种)。 + gjson.GetBytes(bodyBytes, "balance_infos").ForEach(func(_, info gjson.Result) bool { + currency := strings.ToUpper(strings.TrimSpace(info.Get("currency").String())) + balance, _ := cnParseF64(info.Get("total_balance").Value()) + if currency == "" { + currency = "CNY" + } + entries = append(entries, CNProviderBalanceEntry{Currency: currency, Balance: balance}) + return true + }) + if len(entries) == 0 { + entries = append(entries, CNProviderBalanceEntry{Currency: "CNY"}) + } + } + result.Balances = entries + result.Balance = entries[0].Balance + result.Currency = entries[0].Currency + result.Available = available + result.Success = true + + balanceUpdates := make([]any, 0, len(entries)) + for _, entry := range entries { + balanceUpdates = append(balanceUpdates, map[string]any{ + "currency": entry.Currency, + "balance": entry.Balance, + }) + } + updates := map[string]any{ + cnExtraKey(provider, cnBalanceExtraSuffixBalance): result.Balance, + cnExtraKey(provider, cnBalanceExtraSuffixCurrency): result.Currency, + cnExtraKey(provider, cnBalanceExtraSuffixAvailable): available, + cnExtraKey(provider, cnBalanceExtraSuffixUpdated): now.Format(time.RFC3339), + cnExtraKey(provider, cnBalanceExtraSuffixBalances): balanceUpdates, + // 余额探测成功即清除响应式 402/429 写下的 balance_low 标记。 + cnExtraKey(provider, cnBalanceExtraSuffixLow): false, + } + if err := s.accountRepo.UpdateExtra(ctx, account.ID, updates); err != nil { + slog.Warn("cn_balance_persist_failed", "account_id", account.ID, "provider", provider, "error", err) + } else { + result.Persisted = true + } + return result, nil +} + +// loadPayGAccount 加载 payg 模式的国产供应商账号(余额仅对 payg 有意义;coding 走额度)。 +func (s *CNProviderBalanceService) loadPayGAccount(ctx context.Context, accountID int64) (*Account, error) { + account, err := s.accountRepo.GetByID(ctx, accountID) + if err != nil { + return nil, infraerrors.Newf(http.StatusNotFound, "CN_BALANCE_ACCOUNT_NOT_FOUND", "account not found: %v", err) + } + if account == nil { + return nil, infraerrors.New(http.StatusNotFound, "CN_BALANCE_ACCOUNT_NOT_FOUND", "account not found") + } + if !account.IsCNProvider() { + return nil, infraerrors.New(http.StatusBadRequest, "CN_BALANCE_INVALID_PLATFORM", "account is not a CN provider account") + } + // coding 账号走额度探测,余额端点不适用。 + if account.IsCodingPlan() { + return nil, infraerrors.New(http.StatusBadRequest, "CN_BALANCE_CODING_PLAN", "coding plan account has no balance endpoint; use quota probe") + } + return account, nil +} + +func (s *CNProviderBalanceService) resolveProxyURL(ctx context.Context, account *Account) string { + if account == nil || account.ProxyID == nil { + return "" + } + if account.Proxy != nil { + return account.Proxy.URL() + } + if s != nil && s.proxyRepo != nil { + if proxy, err := s.proxyRepo.GetByID(ctx, *account.ProxyID); err == nil && proxy != nil { + account.Proxy = proxy + return proxy.URL() + } + } + return "" +} + +// cnBalanceURL 解析账号的余额端点。 +// +// - Kimi:固定 https://api.moonshot.cn/v1/users/me/balance(与 base_url 无关,Moonshot 仅此一处) +// - DeepSeek:基于 base_url 拼接 /user/balance(支持自定义域名) +func cnBalanceURL(account *Account) string { + switch account.Platform { + case PlatformKimi: + return "https://api.moonshot.cn/v1/users/me/balance" + case PlatformDeepseek: + // Anthropic 协议账号的凭证 base_url 指向 /anthropic 端点,余额探测需回退 + // 到 OpenAI 格式 base(协议感知)再拼接 /user/balance。 + return strings.TrimRight(account.GetOpenAIFormatBaseURL(), "/") + "/user/balance" + default: + return "" + } +} diff --git a/backend/internal/service/cn_provider_probe_url.go b/backend/internal/service/cn_provider_probe_url.go new file mode 100644 index 0000000000..95b0ff783f --- /dev/null +++ b/backend/internal/service/cn_provider_probe_url.go @@ -0,0 +1,48 @@ +package service + +// CN 供应商探测端点的出站 URL 安全策略校验(配额/余额探测共用)。 +// +// 背景(review B4):这两条探测路径会把账号 API key 发往 base_url 衍生端点, +// 此前完全绕过 security.url_allowlist——在加固部署里构成任意外发与内网探测 +// 面(本项目此前发生过账号测试 SSRF 生产事件)。与网关转发 +// (validateUpstreamBaseURL)、Grok 探测(grokOperatorPolicyValidator)一致, +// 探测发起前必须过同一套运营者策略。 + +import ( + "errors" + "fmt" + "strings" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" +) + +// cnValidateProbeURL 按全局出站 URL 安全策略校验探测端点,返回规范化 URL。 +// 白名单开启时强制 UpstreamHosts(阻断私网与未列名主机);关闭时仅做格式 +// 校验(HTTP 允许与否跟随配置);cfg 为 nil 时退化为纯格式校验。 +func cnValidateProbeURL(cfg *config.Config, raw string) (string, error) { + trimmed := strings.TrimSpace(raw) + if trimmed == "" { + return "", errors.New("probe url is required") + } + if cfg != nil && cfg.Security.URLAllowlist.Enabled { + normalized, err := urlvalidator.ValidateHTTPSURL(trimmed, urlvalidator.ValidationOptions{ + AllowedHosts: cfg.Security.URLAllowlist.UpstreamHosts, + RequireAllowlist: true, + AllowPrivate: cfg.Security.URLAllowlist.AllowPrivateHosts, + }) + if err != nil { + return "", fmt.Errorf("probe target rejected by URL security policy: %w", err) + } + return normalized, nil + } + var allowInsecureHTTP bool + if cfg != nil { + allowInsecureHTTP = cfg.Security.URLAllowlist.AllowInsecureHTTP + } + normalized, err := urlvalidator.ValidateURLFormat(trimmed, allowInsecureHTTP) + if err != nil { + return "", fmt.Errorf("probe target rejected by URL security policy: %w", err) + } + return normalized, nil +} diff --git a/backend/internal/service/cn_provider_probe_url_test.go b/backend/internal/service/cn_provider_probe_url_test.go new file mode 100644 index 0000000000..f285088400 --- /dev/null +++ b/backend/internal/service/cn_provider_probe_url_test.go @@ -0,0 +1,129 @@ +package service + +// CN 供应商探测端点 URL 安全策略回归测试(review B4): +// 配额/余额探测不得绕过 security.url_allowlist——base_url 衍生的探测端点 +// 必须先过运营者策略,被拒绝时不得发起任何上游请求(API key 不出站)。 + +import ( + "context" + "net/http" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + "github.com/stretchr/testify/require" +) + +func cnProbeAllowlistConfig(hosts ...string) *config.Config { + return &config.Config{ + Security: config.SecurityConfig{ + URLAllowlist: config.URLAllowlistConfig{ + Enabled: true, + UpstreamHosts: hosts, + }, + }, + } +} + +func TestCNValidateProbeURL_AllowlistPolicy(t *testing.T) { + cfg := cnProbeAllowlistConfig("api.moonshot.cn", "api.deepseek.com") + + // 白名单内主机放行(保留完整路径)。 + ok, err := cnValidateProbeURL(cfg, "https://api.moonshot.cn/v1/users/me/balance") + require.NoError(t, err) + require.Equal(t, "https://api.moonshot.cn/v1/users/me/balance", ok) + + // 白名单外主机拒绝。 + _, err = cnValidateProbeURL(cfg, "https://relay.attacker.example/v1/usages") + require.Error(t, err) + require.Contains(t, err.Error(), "rejected by URL security policy") + + // 私网主机拒绝(内网探测面)。 + _, err = cnValidateProbeURL(cfg, "http://169.254.169.254/latest/meta-data") + require.Error(t, err) + + // 白名单关闭:仅格式校验,任意 https 主机放行。 + formatOnly, err := cnValidateProbeURL(&config.Config{ + Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}, + }, "https://relay.attacker.example/v1/usages") + require.NoError(t, err) + require.Equal(t, "https://relay.attacker.example/v1/usages", formatOnly) +} + +// recordingHTTPUpstream 断言探测被策略拒绝时没有任何上游请求发出。 +type recordingHTTPUpstream struct{ calls int } + +func (u *recordingHTTPUpstream) Do(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) { + u.calls++ + return nil, context.DeadlineExceeded +} + +func (u *recordingHTTPUpstream) DoWithTLS(req *http.Request, proxyURL string, accountID int64, accountConcurrency int, profile *tlsfingerprint.Profile) (*http.Response, error) { + u.calls++ + return nil, context.DeadlineExceeded +} + +type fakeCNProbeAccountRepo struct { + AccountRepository + account *Account +} + +func (r *fakeCNProbeAccountRepo) GetByID(ctx context.Context, id int64) (*Account, error) { + return r.account, nil +} + +// kimi coding 账号的 base_url 指向中转(含 api.kimi.com/coding 路径段即可被识别 +// 为 kimi coding plan)→ 衍生额度端点落在中转主机上,白名单未列名必须拒绝。 +func TestCNProviderQuotaService_RejectsURLBlockedByPolicy(t *testing.T) { + repo := &fakeCNProbeAccountRepo{account: &Account{ + ID: 1, Platform: PlatformKimi, Type: AccountTypeAPIKey, Status: StatusActive, + Credentials: map[string]any{ + "account_mode": "coding", + "api_key": "sk-test", + "base_url": "https://relay.attacker.example/api.kimi.com/coding", + }, + }} + upstream := &recordingHTTPUpstream{} + svc := NewCNProviderQuotaService(repo, nil, upstream, cnProbeAllowlistConfig("api.kimi.com")) + + _, err := svc.QueryUsage(context.Background(), 1) + require.Error(t, err) + require.Contains(t, err.Error(), "CN_QUOTA_URL_REJECTED") + require.Zero(t, upstream.calls, "probe must not issue any upstream request when URL policy rejects the target") +} + +// deepseek payg 账号自定义 base_url → 余额端点落在中转主机上,必须先过策略。 +func TestCNProviderBalanceService_RejectsURLBlockedByPolicy(t *testing.T) { + repo := &fakeCNProbeAccountRepo{account: &Account{ + ID: 2, Platform: PlatformDeepseek, Type: AccountTypeAPIKey, Status: StatusActive, + Credentials: map[string]any{ + "account_mode": "payg", + "api_key": "sk-test", + "base_url": "https://relay.attacker.example", + }, + }} + upstream := &recordingHTTPUpstream{} + svc := NewCNProviderBalanceService(repo, nil, upstream, cnProbeAllowlistConfig("api.deepseek.com")) + + _, err := svc.QueryBalance(context.Background(), 2) + require.Error(t, err) + require.Contains(t, err.Error(), "CN_BALANCE_URL_REJECTED") + require.Zero(t, upstream.calls, "probe must not issue any upstream request when URL policy rejects the target") +} + +// 白名单包含官方主机的正常路径:URL 校验通过后才发出上游请求(此处允许到达 +// httpUpstream 层即视为通过校验,不发真实网络)。 +func TestCNProviderBalanceService_OfficialHostPassesValidation(t *testing.T) { + repo := &fakeCNProbeAccountRepo{account: &Account{ + ID: 3, Platform: PlatformDeepseek, Type: AccountTypeAPIKey, Status: StatusActive, + Credentials: map[string]any{ + "account_mode": "payg", + "api_key": "sk-test", + }, + }} + upstream := &recordingHTTPUpstream{} + svc := NewCNProviderBalanceService(repo, nil, upstream, cnProbeAllowlistConfig("api.deepseek.com")) + + _, _ = svc.QueryBalance(context.Background(), 3) + require.Equal(t, 1, upstream.calls, "official host must pass URL policy and reach the upstream layer") +} diff --git a/backend/internal/service/cn_provider_quota_service.go b/backend/internal/service/cn_provider_quota_service.go new file mode 100644 index 0000000000..f4cda8d6d1 --- /dev/null +++ b/backend/internal/service/cn_provider_quota_service.go @@ -0,0 +1,565 @@ +package service + +import ( + "context" + "encoding/json" + "fmt" + "io" + "log/slog" + "net/http" + "sort" + "strconv" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/tidwall/gjson" + "golang.org/x/sync/singleflight" +) + +// 国产供应商 Coding Plan 滚动窗口额度探测服务(Kimi For Coding / 智谱 GLM Coding Plan)。 +// +// 与 grok_quota_service 不同:CN 供应商走数据面 API Key(无 OAuth token provider), +// 额度端点为只读 GET,解析 5h + weekly 两档滚动窗口并落 account.Extra 快照, +// 供账号调度阈值评估(account_scheduling_threshold_eval.go)做主动停调。 +// +// 解析逻辑对齐 cc-switch(farion1231/cc-switch)services/coding_plan.rs 的 +// query_kimi / query_zhipu,包括智谱 unit 字段优先分类与 reset 兜底启发式。 +const ( + cnQuotaUpstreamTimeout = 15 * time.Second + cnQuotaMaxBodyBytes = 256 * 1024 + + // Extra 快照键后缀(加 provider 前缀,如 kimi_5h_used_percent)。 + cnExtraSuffix5hUsed = "5h_used_percent" + cnExtraSuffix5hReset = "5h_reset_at" + cnExtraSuffixWeeklyUsed = "weekly_used_percent" + cnExtraSuffixWeeklyReset = "weekly_reset_at" + cnExtraSuffixUsageUpdated = "usage_updated_at" +) + +// cnExtraKey 拼接 provider 维度的 extra 键。 +func cnExtraKey(provider, suffix string) string { return provider + "_" + suffix } + +// CNQuotaTier 表示一个滚动用量窗口档位(5h / weekly)。 +type CNQuotaTier struct { + Window string `json:"window"` // "5h" | "weekly" + UsedPercent float64 `json:"used_percent"` // 已用百分比(0-100+,不做裁剪) + ResetAt string `json:"reset_at,omitempty"` // RFC3339,空表示无重置时间 +} + +// CNProviderQuotaProbeResult 是 Coding Plan 额度探测的返回结构(管理端 + UI 消费)。 +type CNProviderQuotaProbeResult struct { + Provider string `json:"provider"` + Source string `json:"source"` + Success bool `json:"success"` + CredentialValid bool `json:"credential_valid"` // false = 401/403 鉴权失败 + Tiers []CNQuotaTier `json:"tiers,omitempty"` + PlanLevel string `json:"plan_level,omitempty"` // 智谱套餐等级 + StatusCode int `json:"status_code,omitempty"` + FetchedAt int64 `json:"fetched_at"` + Persisted bool `json:"persisted"` + Error string `json:"error,omitempty"` +} + +// CNProviderQuotaService 探测 Kimi / Zhipu Coding Plan 的滚动窗口用量。 +type CNProviderQuotaService struct { + accountRepo AccountRepository + proxyRepo ProxyRepository + httpUpstream HTTPUpstream + cfg *config.Config + flight singleflight.Group +} + +// NewCNProviderQuotaService 构造 Coding Plan 额度探测服务。 +func NewCNProviderQuotaService( + accountRepo AccountRepository, + proxyRepo ProxyRepository, + httpUpstream HTTPUpstream, + cfg *config.Config, +) *CNProviderQuotaService { + return &CNProviderQuotaService{ + accountRepo: accountRepo, + proxyRepo: proxyRepo, + httpUpstream: httpUpstream, + cfg: cfg, + } +} + +// QueryUsage 探测指定账号的 Coding Plan 滚动窗口用量并落 Extra 快照。 +// 同一账号的并发探测会被 singleflight 合并。 +func (s *CNProviderQuotaService) QueryUsage(ctx context.Context, accountID int64) (*CNProviderQuotaProbeResult, error) { + if s == nil || s.accountRepo == nil || s.httpUpstream == nil { + return nil, infraerrors.New(http.StatusInternalServerError, "CN_QUOTA_NOT_CONFIGURED", "cn provider quota service is not configured") + } + key := "cn_quota:" + strconv.FormatInt(accountID, 10) + resultCh := s.flight.DoChan(key, func() (any, error) { + probeCtx, cancel := context.WithTimeout(context.Background(), cnQuotaUpstreamTimeout+5*time.Second) + defer cancel() + return s.queryUsage(probeCtx, accountID) + }) + select { + case <-ctx.Done(): + return nil, ctx.Err() + case flightResult := <-resultCh: + if flightResult.Err != nil { + return nil, flightResult.Err + } + result, ok := flightResult.Val.(*CNProviderQuotaProbeResult) + if !ok || result == nil { + return nil, infraerrors.New(http.StatusInternalServerError, "CN_QUOTA_PROBE_RESULT_INVALID", "invalid cn provider quota probe result") + } + cloned := *result + return &cloned, nil + } +} + +func (s *CNProviderQuotaService) queryUsage(ctx context.Context, accountID int64) (*CNProviderQuotaProbeResult, error) { + account, err := s.loadCodingPlanAccount(ctx, accountID) + if err != nil { + return nil, err + } + + provider := account.GetCodingPlanProvider() + if provider != PlatformKimi && provider != PlatformZhipu { + return nil, infraerrors.New(http.StatusBadRequest, "CN_QUOTA_NOT_CODING_PLAN", "account is not a kimi/zhipu coding plan account") + } + + apiKey := strings.TrimSpace(account.GetCNAPIKey()) + if apiKey == "" { + return nil, infraerrors.New(http.StatusBadRequest, "CN_QUOTA_NO_APIKEY", "account api_key is empty") + } + + baseURL := account.GetOpenAIBaseURL() + var ( + targetURL string + authHeader string + ) + switch provider { + case PlatformKimi: + targetURL = kimiQuotaURL(baseURL) + authHeader = "Bearer " + apiKey + case PlatformZhipu: + targetURL = zhipuQuotaURL(baseURL) + authHeader = apiKey // 智谱额度端点鉴权不加 Bearer 前缀 + } + + // 探测发起前过出站 URL 安全策略(与网关转发/Grok 探测同一套校验): + // 端点多由账号 base_url 衍生,不得把 API key 发往策略外主机。 + validatedURL, err := cnValidateProbeURL(s.cfg, targetURL) + if err != nil { + return nil, infraerrors.New(http.StatusForbidden, "CN_QUOTA_URL_REJECTED", err.Error()) + } + targetURL = validatedURL + + proxyURL := s.resolveProxyURL(ctx, account) + callCtx, cancel := context.WithTimeout(ctx, cnQuotaUpstreamTimeout) + defer cancel() + req, err := http.NewRequestWithContext(callCtx, http.MethodGet, targetURL, nil) + if err != nil { + return nil, infraerrors.Newf(http.StatusInternalServerError, "CN_QUOTA_REQUEST_BUILD_FAILED", "build request: %v", err) + } + req.Header.Set("Authorization", authHeader) + req.Header.Set("Accept", "application/json") + if provider == PlatformZhipu { + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept-Language", "en-US,en") + } + // 探测与真实转发保持同一套账号级请求头覆写,避免探测通过但转发失败。 + account.ApplyHeaderOverrides(req.Header) + + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, maxInt(account.Concurrency, 1)) + if err != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "CN_QUOTA_REQUEST_FAILED", "upstream request failed: %v", err) + } + defer func() { _ = resp.Body.Close() }() + bodyBytes, _ := io.ReadAll(io.LimitReader(resp.Body, cnQuotaMaxBodyBytes)) + + now := time.Now().UTC() + result := &CNProviderQuotaProbeResult{ + Provider: provider, + Source: "coding_plan", + FetchedAt: now.Unix(), + StatusCode: resp.StatusCode, + } + + if resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden { + // 鉴权失败:不落快照(不覆盖之前的有效值),仅返回失败结果供前端提示。 + result.Error = fmt.Sprintf("Authentication failed (HTTP %d)", resp.StatusCode) + return result, nil + } + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + result.Error = fmt.Sprintf("API error (HTTP %d): %s", resp.StatusCode, truncate(strings.TrimSpace(string(bodyBytes)), 240)) + return result, nil + } + + // 智谱业务级错误(HTTP 2xx 但 success=false)。 + if provider == PlatformZhipu { + if success := gjson.GetBytes(bodyBytes, "success"); success.Exists() && !success.Bool() { + msg := strings.TrimSpace(gjson.GetBytes(bodyBytes, "msg").String()) + if msg == "" { + msg = "unknown zhipu quota error" + } + result.Error = "API error: " + msg + return result, nil + } + } + + var tiers []CNQuotaTier + switch provider { + case PlatformKimi: + tiers = parseKimiUsageTiers(bodyBytes) + case PlatformZhipu: + tiers = parseZhipuTokenTiers(gjson.GetBytes(bodyBytes, "data")) + result.PlanLevel = strings.TrimSpace(gjson.GetBytes(bodyBytes, "data.level").String()) + } + result.Tiers = tiers + result.Success = true + result.CredentialValid = true + + updates := cnQuotaExtraUpdates(provider, tiers, now) + if err := s.accountRepo.UpdateExtra(ctx, account.ID, updates); err != nil { + slog.Warn("cn_quota_persist_failed", "account_id", account.ID, "provider", provider, "error", err) + } else { + result.Persisted = true + } + return result, nil +} + +func (s *CNProviderQuotaService) loadCodingPlanAccount(ctx context.Context, accountID int64) (*Account, error) { + account, err := s.accountRepo.GetByID(ctx, accountID) + if err != nil { + return nil, infraerrors.Newf(http.StatusNotFound, "CN_QUOTA_ACCOUNT_NOT_FOUND", "account not found: %v", err) + } + if account == nil { + return nil, infraerrors.New(http.StatusNotFound, "CN_QUOTA_ACCOUNT_NOT_FOUND", "account not found") + } + if !account.IsCNProvider() { + return nil, infraerrors.New(http.StatusBadRequest, "CN_QUOTA_INVALID_PLATFORM", "account is not a CN provider account") + } + if !account.IsCodingPlan() { + return nil, infraerrors.New(http.StatusBadRequest, "CN_QUOTA_NOT_CODING_PLAN", "account is not a coding plan account") + } + return account, nil +} + +func (s *CNProviderQuotaService) resolveProxyURL(ctx context.Context, account *Account) string { + if account == nil || account.ProxyID == nil { + return "" + } + if account.Proxy != nil { + return account.Proxy.URL() + } + if s != nil && s.proxyRepo != nil { + if proxy, err := s.proxyRepo.GetByID(ctx, *account.ProxyID); err == nil && proxy != nil { + account.Proxy = proxy + return proxy.URL() + } + } + return "" +} + +// zhipuQuotaURL 根据 base_url 解析智谱额度端点(与数据面推理域名同主机)。 +func zhipuQuotaURL(baseURL string) string { + return zhipuQuotaHost(baseURL) + "/api/monitor/usage/quota/limit" +} + +// kimiQuotaURL 根据 base_url 解析 Kimi For Coding 额度端点。 +// cc-switch query_kimi 固定探测 https://api.kimi.com/coding/v1/usages +// (实测 /coding/usages 无 /v1 → 404)。coding/v1(CC 协议默认)与 +// coding(Anthropic 协议默认)两种 base 统一剥掉尾部后拼回 /v1/usages, +// 协议切换不影响额度探测端点。 +func kimiQuotaURL(baseURL string) string { + base := strings.TrimSuffix(strings.TrimRight(baseURL, "/"), "/v1") + return base + "/v1/usages" +} + +func zhipuQuotaHost(baseURL string) string { + switch u := strings.ToLower(baseURL); { + case strings.Contains(u, "bigmodel.cn"): + return "https://open.bigmodel.cn" + case strings.Contains(u, "z.ai"): + return "https://api.z.ai" + default: + // 国产优先:未知域名回落国内站(与前端 zhipu 预设一致)。 + return "https://open.bigmodel.cn" + } +} + +// parseKimiUsageTiers 解析 Kimi For Coding 的 /usages 响应。 +// +// - limits[].detail.{limit,remaining,resetTime} → 5h 窗口(取首个 detail) +// - usage.{limit,remaining,resetTime} → 周窗口 +// +// utilization = (limit-remaining)/limit*100。 +func parseKimiUsageTiers(body []byte) []CNQuotaTier { + var tiers []CNQuotaTier + + if limits := gjson.GetBytes(body, "limits"); limits.IsArray() { + limits.ForEach(func(_, item gjson.Result) bool { + detail := item.Get("detail") + if !detail.Exists() { + return true + } + limit, _ := cnParseF64(detail.Get("limit").Value()) + remaining, _ := cnParseF64(detail.Get("remaining").Value()) + used := limit - remaining + if used < 0 { + used = 0 + } + var util float64 + if limit > 0 { + util = used / limit * 100 + } + tiers = append(tiers, CNQuotaTier{ + Window: "5h", + UsedPercent: util, + ResetAt: cnNormalizeResetTime(detail.Get("resetTime").Value()), + }) + return false // 取首个 detail 作为 5h 窗口 + }) + } + + if usage := gjson.GetBytes(body, "usage"); usage.Exists() { + limit, _ := cnParseF64(usage.Get("limit").Value()) + remaining, _ := cnParseF64(usage.Get("remaining").Value()) + used := limit - remaining + if used < 0 { + used = 0 + } + var util float64 + if limit > 0 { + util = used / limit * 100 + } + tiers = append(tiers, CNQuotaTier{ + Window: "weekly", + UsedPercent: util, + ResetAt: cnNormalizeResetTime(usage.Get("resetTime").Value()), + }) + } + + return tiers +} + +// cnZhipuWindow 标识智谱 TOKENS_LIMIT 条目所属窗口。 +type cnZhipuWindow int + +const ( + cnZhipuWindowUnknown cnZhipuWindow = iota + cnZhipuWindow5h + cnZhipuWindowWeekly +) + +// classifyZhipuWindowUnit 按 unit 字段判定窗口类型(3=5h,6=weekly)。 +// unit 缺失或未识别时返回 Unknown,由调用方走 reset 时间启发式兜底。 +func classifyZhipuWindowUnit(unit int64) cnZhipuWindow { + switch unit { + case 3: + return cnZhipuWindow5h + case 6: + return cnZhipuWindowWeekly + default: + return cnZhipuWindowUnknown + } +} + +// parseZhipuTokenTiers 解析智谱额度响应 data.limits 为 5h + weekly 两档。 +// +// 分类优先级(对齐 cc-switch parse_zhipu_token_tiers,issue #3036): +// 1. 显式 unit 字段(3=5h / 6=weekly)——不能用 reset 排序代替,周期末尾 +// 周窗口会比 5h 更早重置,时间排序必然标反。 +// 2. unit 缺失/未识别:无 nextResetTime 的条目优先归 5h(0% 状态下 5h 桶可能 +// 没有 reset),其余按 reset 升序依次填入仍空缺的槽位。 +// +// CREDIT_LIMIT(信用额度)与 TOKENS_LIMIT(token 窗口)度量不同:两者同时返回时 +// 只让 TOKENS_LIMIT 参与 5h/weekly 槽位竞争,避免信用额度百分比污染阈值停调 +// 快照;仅当无任何 TOKENS_LIMIT 条目时才降级用 CREDIT_LIMIT 展示。 +// 老套餐只回 1 条 TOKENS_LIMIT,自然降级为仅 5h;新套餐回 2 条。 +func parseZhipuTokenTiers(data gjson.Result) []CNQuotaTier { + type entry struct { + resetMs int64 + hasReset bool + percentage float64 + resetISO string + } + var ( + fiveHour entry + fiveHourSet bool + weekly entry + weeklySet bool + unclassified []entry + ) + + classify := func(item gjson.Result, e entry) { + switch classifyZhipuWindowUnit(item.Get("unit").Int()) { + case cnZhipuWindow5h: + if !fiveHourSet { + fiveHour, fiveHourSet = e, true + } else { + unclassified = append(unclassified, e) + } + case cnZhipuWindowWeekly: + if !weeklySet { + weekly, weeklySet = e, true + } else { + unclassified = append(unclassified, e) + } + default: + unclassified = append(unclassified, e) + } + } + var creditFallback []entry + hasTokensLimit := false + + data.Get("limits").ForEach(func(_, item gjson.Result) bool { + limitType := strings.ToUpper(strings.TrimSpace(item.Get("type").String())) + if limitType != "TOKENS_LIMIT" && limitType != "CREDIT_LIMIT" { + return true + } + percentage := 0.0 + if p, ok := cnParseF64(item.Get("percentage").Value()); ok { + percentage = p + } + var ( + resetMs int64 + hasReset bool + resetISO string + ) + if nr := item.Get("nextResetTime"); nr.Exists() { + switch nr.Type { + case gjson.Number: + resetMs = nr.Int() + hasReset = resetMs > 0 + resetISO = cnMillisToRFC3339(resetMs) + case gjson.String: + resetISO = cnNormalizeResetTime(nr.String()) + hasReset = resetISO != "" + } + } + e := entry{resetMs: resetMs, hasReset: hasReset, percentage: percentage, resetISO: resetISO} + if limitType == "TOKENS_LIMIT" { + hasTokensLimit = true + classify(item, e) + } else { + creditFallback = append(creditFallback, e) + } + return true + }) + + // 无任何 TOKENS_LIMIT 条目(部分套餐只报信用额度):降级用 CREDIT_LIMIT 展示。 + if !hasTokensLimit { + unclassified = append(unclassified, creditFallback...) + } + + // 无 reset 的条目排前,再按 reset 升序,依次填入仍空缺的槽位。 + sort.SliceStable(unclassified, func(i, j int) bool { + if unclassified[i].hasReset != unclassified[j].hasReset { + return !unclassified[i].hasReset + } + return unclassified[i].resetMs < unclassified[j].resetMs + }) + for _, e := range unclassified { + switch { + case !fiveHourSet: + fiveHour, fiveHourSet = e, true + case !weeklySet: + weekly, weeklySet = e, true + } + } + + var tiers []CNQuotaTier + if fiveHourSet { + tiers = append(tiers, CNQuotaTier{Window: "5h", UsedPercent: fiveHour.percentage, ResetAt: fiveHour.resetISO}) + } + if weeklySet { + tiers = append(tiers, CNQuotaTier{Window: "weekly", UsedPercent: weekly.percentage, ResetAt: weekly.resetISO}) + } + return tiers +} + +// cnQuotaExtraUpdates 根据 tier 列表构造 provider 维度的 Extra 快照更新。 +func cnQuotaExtraUpdates(provider string, tiers []CNQuotaTier, now time.Time) map[string]any { + updates := map[string]any{ + cnExtraKey(provider, cnExtraSuffixUsageUpdated): now.Format(time.RFC3339), + } + for _, t := range tiers { + switch t.Window { + case "5h": + updates[cnExtraKey(provider, cnExtraSuffix5hUsed)] = t.UsedPercent + if t.ResetAt != "" { + updates[cnExtraKey(provider, cnExtraSuffix5hReset)] = t.ResetAt + } + case "weekly": + updates[cnExtraKey(provider, cnExtraSuffixWeeklyUsed)] = t.UsedPercent + if t.ResetAt != "" { + updates[cnExtraKey(provider, cnExtraSuffixWeeklyReset)] = t.ResetAt + } + } + } + return updates +} + +// cnParseF64 把 JSON 数值或字符串解析为 float64(兼容 "100" 与 100)。 +func cnParseF64(raw any) (float64, bool) { + switch v := raw.(type) { + case float64: + return v, true + case float32: + return float64(v), true + case int: + return float64(v), true + case int64: + return float64(v), true + case json.Number: + f, err := v.Float64() + return f, err == nil + case string: + f, err := strconv.ParseFloat(strings.TrimSpace(v), 64) + return f, err == nil + default: + return 0, false + } +} + +// cnNormalizeResetTime 把上游重置时间(ISO8601 字符串 / 秒级 / 毫秒级数字)归一化为 +// RFC3339 字符串;无法识别或非正时间戳返回空串。 +func cnNormalizeResetTime(raw any) string { + switch v := raw.(type) { + case string: + s := strings.TrimSpace(v) + if s == "" { + return "" + } + if ts, err := parseSchedulingTime(s); err == nil { + return ts.UTC().Format(time.RFC3339) + } + return "" + case float64: + return cnMillisToRFC3339(int64(v)) + case int: + return cnMillisToRFC3339(int64(v)) + case int64: + return cnMillisToRFC3339(v) + case json.Number: + if n, err := v.Int64(); err == nil { + return cnMillisToRFC3339(n) + } + return "" + default: + return "" + } +} + +// cnMillisToRFC3339 把秒级(<1e12)或毫秒级时间戳转为 RFC3339 字符串;非正返回空串。 +func cnMillisToRFC3339(n int64) string { + if n <= 0 { + return "" + } + var ms int64 + if n < 1_000_000_000_000 { + ms = n * 1000 + } else { + ms = n + } + return time.UnixMilli(ms).UTC().Format(time.RFC3339) +} diff --git a/backend/internal/service/cn_providers_test.go b/backend/internal/service/cn_providers_test.go new file mode 100644 index 0000000000..24b8bd4405 --- /dev/null +++ b/backend/internal/service/cn_providers_test.go @@ -0,0 +1,625 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" + + "github.com/Wei-Shaw/sub2api/internal/config" +) + +// TestCNExtraKey 验证 provider 维度的 Extra 快照键由前缀 + 后缀拼接。 +func TestCNExtraKey(t *testing.T) { + t.Parallel() + require.Equal(t, "kimi_5h_used_percent", cnExtraKey(PlatformKimi, cnExtraSuffix5hUsed)) + require.Equal(t, "zhipu_weekly_reset_at", cnExtraKey(PlatformZhipu, cnExtraSuffixWeeklyReset)) + require.Equal(t, "deepseek_balance", cnExtraKey(PlatformDeepseek, cnBalanceExtraSuffixBalance)) +} + +// TestCNParseF64 兼容 JSON 数值与字符串(cc-switch 与上游字段类型不一致)。 +func TestCNParseF64(t *testing.T) { + t.Parallel() + cases := []struct { + name string + raw any + want float64 + ok bool + }{ + {"float64", float64(12.5), 12.5, true}, + {"int", 100, 100, true}, + {"numeric string", "33.3", 33.3, true}, + {"trim string", " 7 ", 7, true}, + {"non-numeric string", "abc", 0, false}, + {"nil", nil, 0, false}, + } + for _, tc := range cases { + tc := tc + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + got, ok := cnParseF64(tc.raw) + require.Equal(t, tc.ok, ok) + if ok { + require.InDelta(t, tc.want, got, 1e-9) + } + }) + } +} + +// TestCNMillisToRFC3339 秒级(<1e12)按秒、毫秒级按毫秒处理;非正返回空串。 +func TestCNMillisToRFC3339(t *testing.T) { + t.Parallel() + // 1700000000 秒 = 1700000000000 毫秒 + want := time.UnixMilli(1700000000000).UTC().Format(time.RFC3339) + require.Equal(t, want, cnMillisToRFC3339(1700000000)) // 秒级 + require.Equal(t, want, cnMillisToRFC3339(1700000000000)) // 毫秒级 + require.Equal(t, "", cnMillisToRFC3339(0)) // 非正 + require.Equal(t, "", cnMillisToRFC3339(-1)) +} + +// TestCnnormalizeResetTime 覆盖 ISO8601 字符串 / 数字(秒、毫秒)/ 非法输入。 +func TestCnnormalizeResetTime(t *testing.T) { + t.Parallel() + // ISO8601 字符串归一化为 RFC3339(UTC)。 + require.Equal(t, "2026-08-14T10:00:00Z", cnNormalizeResetTime("2026-08-14T10:00:00Z")) + // 毫秒级 float64。 + require.Equal(t, + time.UnixMilli(1700000000000).UTC().Format(time.RFC3339), + cnNormalizeResetTime(float64(1700000000000))) + // 非法字符串。 + require.Equal(t, "", cnNormalizeResetTime("not-a-time")) + require.Equal(t, "", cnNormalizeResetTime("")) +} + +// TestParseKimiUsageTiers 验证 Kimi For Coding /usages 解析: +// - 首个 limits[].detail → 5h 桶,utilization=(limit-remaining)/limit*100 +// - usage → weekly 桶 +// - 仅取首个 detail(多个 detail 时不应重复产出 5h) +func TestParseKimiUsageTiers(t *testing.T) { + t.Parallel() + body := []byte(`{ + "limits": [ + {"name": "5h", "detail": {"limit": 1000, "remaining": 600, "resetTime": "2026-08-14T15:00:00Z"}}, + {"name": "ignored-second-detail", "detail": {"limit": 999, "remaining": 0, "resetTime": "2026-08-14T20:00:00Z"}} + ], + "usage": {"limit": 10000, "remaining": 4000, "resetTime": "2026-08-18T00:00:00Z"} + }`) + tiers := parseKimiUsageTiers(body) + require.Len(t, tiers, 2) + require.Equal(t, "5h", tiers[0].Window) + require.InDelta(t, 40.0, tiers[0].UsedPercent, 1e-9) // (1000-600)/1000*100 + require.Equal(t, "2026-08-14T15:00:00Z", tiers[0].ResetAt) + require.Equal(t, "weekly", tiers[1].Window) + require.InDelta(t, 60.0, tiers[1].UsedPercent, 1e-9) // (10000-4000)/10000*100 + require.Equal(t, "2026-08-18T00:00:00Z", tiers[1].ResetAt) +} + +// TestParseKimiUsageTiers_LimitZero 不应除零:limit=0 → utilization=0。 +func TestParseKimiUsageTiers_LimitZero(t *testing.T) { + t.Parallel() + body := []byte(`{"limits":[{"detail":{"limit":0,"remaining":0,"resetTime":"2026-08-14T15:00:00Z"}}]}`) + tiers := parseKimiUsageTiers(body) + require.Len(t, tiers, 1) + require.InDelta(t, 0.0, tiers[0].UsedPercent, 1e-9) +} + +// TestParseZhipuTokenTiers_UnitClassification 显式 unit(3=5h / 6=weekly)优先分类, +// 不能被 reset 时间排序覆盖(周期末尾周窗口会更早重置)。 +func TestParseZhipuTokenTiers_UnitClassification(t *testing.T) { + t.Parallel() + // weekly 的 nextResetTime 早于 5h(模拟周期末尾),但 unit 必须胜出。 + data := gjson.Parse(`{ + "limits": [ + {"type":"TOKENS_LIMIT","unit":6,"percentage":70,"nextResetTime":1700000000000}, + {"type":"TOKENS_LIMIT","unit":3,"percentage":20,"nextResetTime":1700000099999} + ] + }`) + tiers := parseZhipuTokenTiers(data) + require.Len(t, tiers, 2) + require.Equal(t, "5h", tiers[0].Window) + require.InDelta(t, 20.0, tiers[0].UsedPercent, 1e-9) + require.Equal(t, "weekly", tiers[1].Window) + require.InDelta(t, 70.0, tiers[1].UsedPercent, 1e-9) +} + +// TestParseZhipuTokenTiers_SingleTierOldPlan 老套餐仅回 1 条 → 降级为仅 5h。 +func TestParseZhipuTokenTiers_SingleTierOldPlan(t *testing.T) { + t.Parallel() + data := gjson.Parse(`{"limits":[{"type":"TOKENS_LIMIT","unit":3,"percentage":15,"nextResetTime":1700000000000}]}`) + tiers := parseZhipuTokenTiers(data) + require.Len(t, tiers, 1) + require.Equal(t, "5h", tiers[0].Window) +} + +// TestParseZhipuTokenTiers_FallbackHeuristic unit 缺失时:无 reset 的条目优先归 5h, +// 其余按 reset 升序填入剩余槽位。 +func TestParseZhipuTokenTiers_FallbackHeuristic(t *testing.T) { + t.Parallel() + // 无 unit:A 无 reset、B 有 reset。A 先填 5h,B 填 weekly。 + data := gjson.Parse(`{ + "limits": [ + {"type":"TOKENS_LIMIT","percentage":50,"nextResetTime":1700000000000}, + {"type":"TOKENS_LIMIT","percentage":10} + ] + }`) + tiers := parseZhipuTokenTiers(data) + require.Len(t, tiers, 2) + require.Equal(t, "5h", tiers[0].Window) + require.InDelta(t, 10.0, tiers[0].UsedPercent, 1e-9) // 无 reset 优先 5h + require.Equal(t, "weekly", tiers[1].Window) + require.InDelta(t, 50.0, tiers[1].UsedPercent, 1e-9) +} + +// TestParseZhipuTokenTiers_IgnoresNonTokenEntries 非 TOKENS_LIMIT/CREDIT_LIMIT 条目跳过。 +func TestParseZhipuTokenTiers_IgnoresNonTokenEntries(t *testing.T) { + t.Parallel() + data := gjson.Parse(`{"limits":[{"type":"OTHER_LIMIT","unit":3,"percentage":99}]}`) + require.Empty(t, parseZhipuTokenTiers(data)) +} + +// TestCNQuotaExtraUpdates 验证 tier 列表落 Extra 快照键的 provider 前缀与窗口映射。 +func TestCNQuotaExtraUpdates(t *testing.T) { + t.Parallel() + now := time.Date(2026, 8, 14, 0, 0, 0, 0, time.UTC) + tiers := []CNQuotaTier{ + {Window: "5h", UsedPercent: 40, ResetAt: "2026-08-14T15:00:00Z"}, + {Window: "weekly", UsedPercent: 60, ResetAt: "2026-08-18T00:00:00Z"}, + } + updates := cnQuotaExtraUpdates(PlatformKimi, tiers, now) + require.Equal(t, 40.0, updates["kimi_5h_used_percent"]) + require.Equal(t, "2026-08-14T15:00:00Z", updates["kimi_5h_reset_at"]) + require.Equal(t, 60.0, updates["kimi_weekly_used_percent"]) + require.Equal(t, "2026-08-18T00:00:00Z", updates["kimi_weekly_reset_at"]) + require.Equal(t, now.Format(time.RFC3339), updates["kimi_usage_updated_at"]) +} + +// TestCNProviderResponseIndicatesInsufficientBalance 覆盖中英文余额不足文案与否定用例。 +func TestCNProviderResponseIndicatesInsufficientBalance(t *testing.T) { + t.Parallel() + positive := []string{ + `{"error":{"message":"余额不足"}}`, + `{"error":{"message":"Insufficient balance"}}`, + `{"code":"insufficient_credit"}`, + `"balance is not enough"`, + `"no enough balance"`, + } + for _, body := range positive { + require.True(t, cnProviderResponseIndicatesInsufficientBalance([]byte(body)), body) + } + negative := []string{ + `{"error":{"message":"rate limit exceeded"}}`, + `{"error":{"message":"quota exhausted"}}`, + ``, + } + for _, body := range negative { + require.False(t, cnProviderResponseIndicatesInsufficientBalance([]byte(body)), body) + } +} + +// TestCNBalanceLowReason 验证稳定前缀(供周期检测任务识别并清除)。 +func TestCNBalanceLowReason(t *testing.T) { + t.Parallel() + require.Equal(t, "cn_balance_low: upstream said x", + cnBalanceLowReason("upstream said x")) + require.Equal(t, "cn_balance_low: 余额不足,账号临时停调", + cnBalanceLowReason(" ")) + require.True(t, len(cnBalanceLowReason("")) > len(cnBalanceLowReasonPrefix)) +} + +// TestZhipuQuotaHost 按域名路由智谱额度端点主机(bigmodel.cn / z.ai / 默认国内站)。 +func TestZhipuQuotaHost(t *testing.T) { + t.Parallel() + require.Equal(t, "https://open.bigmodel.cn", zhipuQuotaHost("https://open.bigmodel.cn/api/paas/v4")) + require.Equal(t, "https://api.z.ai", zhipuQuotaHost("https://api.z.ai/api/paas/v4")) + require.Equal(t, "https://open.bigmodel.cn", zhipuQuotaHost("https://custom.example.com")) // 默认国内站 + require.Equal(t, "https://open.bigmodel.cn/api/monitor/usage/quota/limit", zhipuQuotaURL("https://open.bigmodel.cn/api/paas/v4")) +} + +// TestKimiQuotaURL 两种协议默认 base(coding/v1 与 coding)都归一到 /coding/v1/usages +// (cc-switch 固定端点;无 /v1 的 /coding/usages 实测 404)。 +func TestKimiQuotaURL(t *testing.T) { + t.Parallel() + require.Equal(t, "https://api.kimi.com/coding/v1/usages", kimiQuotaURL("https://api.kimi.com/coding/v1")) + require.Equal(t, "https://api.kimi.com/coding/v1/usages", kimiQuotaURL("https://api.kimi.com/coding")) + require.Equal(t, "https://api.kimi.com/coding/v1/usages", kimiQuotaURL("https://api.kimi.com/coding/")) + require.Equal(t, "https://api.kimi.com/coding/v1/usages", kimiQuotaURL("https://api.kimi.com/coding/v1/")) +} + +// TestCNBalanceURL Kimi 固定端点;DeepSeek 基于 base_url 拼接。 +func TestCNBalanceURL(t *testing.T) { + t.Parallel() + kimi := &Account{Platform: PlatformKimi} + require.Equal(t, "https://api.moonshot.cn/v1/users/me/balance", cnBalanceURL(kimi)) + + deepseek := &Account{ + Platform: PlatformDeepseek, + Credentials: map[string]any{"base_url": "https://api.deepseek.com"}, + } + require.Equal(t, "https://api.deepseek.com/user/balance", cnBalanceURL(deepseek)) +} + +// TestCNProviderThresholdCandidates 从 Extra 快照读取 5h / weekly 候选。 +func TestCNProviderThresholdCandidates(t *testing.T) { + t.Parallel() + account := &Account{ + Platform: PlatformKimi, + Extra: map[string]any{ + "kimi_5h_used_percent": 90.0, + "kimi_5h_reset_at": "2026-08-14T15:00:00Z", + "kimi_weekly_used_percent": 50.0, + "kimi_weekly_reset_at": "2026-08-18T00:00:00Z", + }, + } + cands := cnProviderThresholdCandidates(account, PlatformKimi) + // 仅返回非 nil 候选(两窗口均存在 → 2 条)。 + var present []*accountSchedulingThresholdCandidate + for _, c := range cands { + if c != nil { + present = append(present, c) + } + } + require.Len(t, present, 2) + + // 缺少 used 键的窗口不产生候选。 + partial := &Account{ + Platform: PlatformKimi, + Extra: map[string]any{"kimi_5h_reset_at": "2026-08-14T15:00:00Z"}, // 无 used_percent + } + require.Empty(t, filterNil(cnProviderThresholdCandidates(partial, PlatformKimi))) + + // 空 Extra / nil account 安全返回。 + require.Empty(t, filterNil(cnProviderThresholdCandidates(&Account{Platform: PlatformKimi}, PlatformKimi))) +} + +func filterNil(cands []*accountSchedulingThresholdCandidate) []*accountSchedulingThresholdCandidate { + var out []*accountSchedulingThresholdCandidate + for _, c := range cands { + if c != nil { + out = append(out, c) + } + } + return out +} + +// TestEvaluateAccountSchedulingThreshold_KimiCodingPlan 集成验证:kimi coding 账号 +// 5h 用量超阈值且窗口未重置 → 主动停调至 5h 重置点。 +func TestEvaluateAccountSchedulingThreshold_KimiCodingPlan(t *testing.T) { + t.Parallel() + now := time.Date(2026, 8, 14, 12, 0, 0, 0, time.UTC) + reset := now.Add(3 * time.Hour) + account := &Account{ + Platform: PlatformKimi, + Extra: map[string]any{ + "kimi_5h_used_percent": 90.0, + "kimi_5h_reset_at": reset.Format(time.RFC3339), + "kimi_weekly_used_percent": 30.0, + "kimi_weekly_reset_at": now.Add(7 * 24 * time.Hour).Format(time.RFC3339), + }, + } + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{PlatformKimi: 80}, now) + require.True(t, decision.ShouldPause) + require.Equal(t, PlatformKimi, decision.Platform) + require.Equal(t, "5h", decision.Window) + require.InDelta(t, 90.0, decision.UsedPercent, 1e-9) + require.NotNil(t, decision.Until) + require.True(t, reset.Equal(*decision.Until)) +} + +// TestEvaluateAccountSchedulingThreshold_CNWindowResetSkipped 窗口已重置(reset<=now) +// 或用量低于阈值 → 不停调(candidateMatchesThreshold 要求 until.After(now))。 +func TestEvaluateAccountSchedulingThreshold_CNWindowResetSkipped(t *testing.T) { + t.Parallel() + now := time.Date(2026, 8, 14, 12, 0, 0, 0, time.UTC) + // 重置时间已过。 + expired := &Account{ + Platform: PlatformZhipu, + Extra: map[string]any{ + "zhipu_5h_used_percent": 99.0, + "zhipu_5h_reset_at": now.Add(-1 * time.Hour).Format(time.RFC3339), + }, + } + require.False(t, EvaluateAccountSchedulingThreshold(expired, map[string]int{PlatformZhipu: 80}, now).ShouldPause) + + // 用量低于阈值。 + low := &Account{ + Platform: PlatformZhipu, + Extra: map[string]any{ + "zhipu_5h_used_percent": 20.0, + "zhipu_5h_reset_at": now.Add(3 * time.Hour).Format(time.RFC3339), + }, + } + require.False(t, EvaluateAccountSchedulingThreshold(low, map[string]int{PlatformZhipu: 80}, now).ShouldPause) +} + +// TestCNProviderQuotaSnapshotReset Coding Plan 429 冷却:取快照中最早的「仍在未来」窗口重置点。 +func TestCNProviderQuotaSnapshotReset(t *testing.T) { + t.Parallel() + now := time.Date(2026, 8, 14, 12, 0, 0, 0, time.UTC) + future5h := now.Add(2 * time.Hour) + futureWeekly := now.Add(3 * 24 * time.Hour) + pastWeekly := now.Add(-24 * time.Hour) + + // 5h 在未来、weekly 已过期 → 返回 5h。 + account := &Account{ + Platform: PlatformKimi, + Credentials: map[string]any{"account_mode": AccountModeCoding}, + Extra: map[string]any{ + "kimi_5h_reset_at": future5h.Format(time.RFC3339), + "kimi_weekly_reset_at": pastWeekly.Format(time.RFC3339), + }, + } + got := cnProviderQuotaSnapshotReset(account, now) + require.NotNil(t, got) + require.True(t, future5h.Equal(*got)) + + // 两窗口均在未来 → 取较早者(429 多由 5h 窗口触发,避免冷却到 weekly 重置)。 + both := &Account{ + Platform: PlatformKimi, + Credentials: map[string]any{"account_mode": AccountModeCoding}, + Extra: map[string]any{ + "kimi_5h_reset_at": future5h.Format(time.RFC3339), + "kimi_weekly_reset_at": futureWeekly.Format(time.RFC3339), + }, + } + gotBoth := cnProviderQuotaSnapshotReset(both, now) + require.NotNil(t, gotBoth) + require.True(t, future5h.Equal(*gotBoth)) + + // 两窗口均过期 → nil。 + expired := &Account{ + Platform: PlatformKimi, + Credentials: map[string]any{"account_mode": AccountModeCoding}, + Extra: map[string]any{ + "kimi_5h_reset_at": pastWeekly.Format(time.RFC3339), + "kimi_weekly_reset_at": pastWeekly.Format(time.RFC3339), + }, + } + require.Nil(t, cnProviderQuotaSnapshotReset(expired, now)) + + // payg 账号(非 coding)→ nil(余额型走余额检测)。 + payg := &Account{ + Platform: PlatformKimi, + Credentials: map[string]any{"account_mode": AccountModePayG}, + Extra: map[string]any{"kimi_5h_reset_at": future5h.Format(time.RFC3339)}, + } + require.Nil(t, cnProviderQuotaSnapshotReset(payg, now)) +} + +// TestNormalizeOpenAICompatiblePlatform_SchedulerExactMatch 回归保护: +// grok 与国产供应商原样保留,其余归一为 openai —— 保证 kimi/zhipu/deepseek 分组请求 +// 精确匹配同名账号(与 openai/grok 当前行为一致),不会错误并入 openai 池。 +func TestNormalizeOpenAICompatiblePlatform_SchedulerExactMatch(t *testing.T) { + t.Parallel() + require.Equal(t, PlatformGrok, NormalizeOpenAICompatiblePlatform(PlatformGrok)) + require.Equal(t, PlatformKimi, NormalizeOpenAICompatiblePlatform(PlatformKimi)) + require.Equal(t, PlatformZhipu, NormalizeOpenAICompatiblePlatform(PlatformZhipu)) + require.Equal(t, PlatformDeepseek, NormalizeOpenAICompatiblePlatform(PlatformDeepseek)) + // 其他平台(含空、anthropic、未知)一律归一为 openai。 + require.Equal(t, PlatformOpenAI, NormalizeOpenAICompatiblePlatform("")) + require.Equal(t, PlatformOpenAI, NormalizeOpenAICompatiblePlatform(PlatformAnthropic)) + require.Equal(t, PlatformOpenAI, NormalizeOpenAICompatiblePlatform("something-else")) +} + +// TestGetOpenAIProtocolAPIKey_CNProviders 验证 OpenAI 协议族密钥读取覆盖国产供应商, +// 同时保持 IsOpenAIApiKey 的 openai-only 语义(调度倍率/WS 门控不受影响)。 +func TestGetOpenAIProtocolAPIKey_CNProviders(t *testing.T) { + t.Parallel() + + kimi := &Account{ + Platform: PlatformKimi, + Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-kimi"}, + } + require.Equal(t, "sk-kimi", kimi.GetOpenAIProtocolAPIKey()) + require.False(t, kimi.IsOpenAIApiKey(), "IsOpenAIApiKey stays openai-only for scheduling gates") + + // 非 APIKey 类型的 CN 账号不返回密钥 + notAPIKey := &Account{ + Platform: PlatformDeepseek, + Type: AccountTypeOAuth, + Credentials: map[string]any{"api_key": "sk-leak"}, + } + require.Equal(t, "", notAPIKey.GetOpenAIProtocolAPIKey()) + + // openai 原生账号行为不变 + openai := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-openai"}, + } + require.Equal(t, "sk-openai", openai.GetOpenAIProtocolAPIKey()) +} + +// TestBuildUpstreamModelsRequest_CNProviders 验证“同步上游支持的模型”对国产供应商可用: +// 密钥经 GetOpenAIProtocolAPIKey 读取,/models 端点拼接到账号 base_url(含默认值)。 +func TestBuildUpstreamModelsRequest_CNProviders(t *testing.T) { + t.Parallel() + + svc := &AccountTestService{cfg: &config.Config{}} + cases := []struct { + name string + platform string + mode string + wantURL string + }{ + {"kimi default", PlatformKimi, "", "https://api.moonshot.cn/v1/models"}, + {"kimi coding", PlatformKimi, AccountModeCoding, "https://api.kimi.com/coding/v1/models"}, + {"zhipu default", PlatformZhipu, "", "https://open.bigmodel.cn/api/paas/v4/models"}, + {"zhipu coding", PlatformZhipu, AccountModeCoding, "https://open.bigmodel.cn/api/coding/paas/v4/models"}, + {"deepseek", PlatformDeepseek, "", "https://api.deepseek.com/v1/models"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + creds := map[string]any{"api_key": "sk-test"} + if tc.mode != "" { + creds["account_mode"] = tc.mode + } + account := &Account{ID: 1, Platform: tc.platform, Type: AccountTypeAPIKey, Credentials: creds} + req, err := svc.buildUpstreamModelsRequest(context.Background(), account) + require.NoError(t, err) + require.Equal(t, tc.wantURL, req.URL.String()) + require.Equal(t, "Bearer sk-test", req.Header.Get("Authorization")) + }) + } +} + +// TestGetAPIProtocol 验证协议凭证维度的平台校验矩阵: +// responses 仅 deepseek;缺失/非法值回退 chat_completions(与旧行为一致)。 +func TestGetAPIProtocol(t *testing.T) { + t.Parallel() + + mk := func(platform, protocol string) *Account { + creds := map[string]any{"api_key": "sk-test"} + if protocol != "" { + creds["api_protocol"] = protocol + } + return &Account{Platform: platform, Type: AccountTypeAPIKey, Credentials: creds} + } + + require.Equal(t, APIProtocolChatCompletions, mk(PlatformKimi, "").GetAPIProtocol(), "缺失回退默认") + require.Equal(t, APIProtocolAnthropic, mk(PlatformZhipu, APIProtocolAnthropic).GetAPIProtocol()) + require.Equal(t, APIProtocolAnthropic, mk(PlatformKimi, APIProtocolAnthropic).GetAPIProtocol()) + require.Equal(t, APIProtocolAnthropic, mk(PlatformDeepseek, APIProtocolAnthropic).GetAPIProtocol()) + require.Equal(t, APIProtocolResponses, mk(PlatformDeepseek, APIProtocolResponses).GetAPIProtocol()) + require.Equal(t, APIProtocolChatCompletions, mk(PlatformKimi, APIProtocolResponses).GetAPIProtocol(), "kimi 无 responses 端点") + require.Equal(t, APIProtocolChatCompletions, mk(PlatformZhipu, APIProtocolResponses).GetAPIProtocol(), "zhipu 无 responses 端点") + require.Equal(t, APIProtocolChatCompletions, mk(PlatformKimi, "bogus").GetAPIProtocol(), "非法值回退默认") + require.Equal(t, APIProtocolChatCompletions, (&Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}).GetAPIProtocol(), "非 CN 供应商恒为默认") +} + +// TestAnthropicProtocolBaseURL 验证 Anthropic 协议默认端点与协议感知的 +// OpenAI 格式 base 回退。 +func TestAnthropicProtocolBaseURL(t *testing.T) { + t.Parallel() + + // 默认端点(按供应商 × 模式) + require.Equal(t, "https://api.moonshot.cn/anthropic", (&Account{ + Platform: PlatformKimi, Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_protocol": APIProtocolAnthropic}, + }).GetAnthropicProtocolBaseURL()) + require.Equal(t, "https://api.kimi.com/coding", (&Account{ + Platform: PlatformKimi, Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_protocol": APIProtocolAnthropic, "account_mode": AccountModeCoding}, + }).GetAnthropicProtocolBaseURL()) + require.Equal(t, "https://open.bigmodel.cn/api/anthropic", (&Account{ + Platform: PlatformZhipu, Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_protocol": APIProtocolAnthropic}, + }).GetAnthropicProtocolBaseURL()) + require.Equal(t, "https://api.deepseek.com/anthropic", (&Account{ + Platform: PlatformDeepseek, Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_protocol": APIProtocolAnthropic}, + }).GetAnthropicProtocolBaseURL()) + + // 凭证 base_url 覆盖默认值 + require.Equal(t, "https://custom.example.com/anthropic", (&Account{ + Platform: PlatformZhipu, Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_protocol": APIProtocolAnthropic, "base_url": "https://custom.example.com/anthropic"}, + }).GetAnthropicProtocolBaseURL()) + + // 非 Anthropic 协议返回空串 + require.Empty(t, (&Account{ + Platform: PlatformZhipu, Type: AccountTypeAPIKey, + Credentials: map[string]any{"base_url": "https://open.bigmodel.cn/api/paas/v4"}, + }).GetAnthropicProtocolBaseURL()) +} + +// TestGetOpenAIFormatBaseURL_ProtocolAware anthropic 协议账号的凭证 base_url +// 指向 Anthropic 端点,OpenAI 格式路径(模型同步等)必须回退到 CC 默认 base。 +func TestGetOpenAIFormatBaseURL_ProtocolAware(t *testing.T) { + t.Parallel() + + zhipuAnthropic := &Account{ + Platform: PlatformZhipu, Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_protocol": APIProtocolAnthropic, + "base_url": "https://open.bigmodel.cn/api/anthropic", + }, + } + require.Equal(t, "https://open.bigmodel.cn/api/paas/v4", zhipuAnthropic.GetOpenAIFormatBaseURL()) + + kimiCodingAnthropic := &Account{ + Platform: PlatformKimi, Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_protocol": APIProtocolAnthropic, + "account_mode": AccountModeCoding, + "base_url": "https://api.kimi.com/coding", + }, + } + require.Equal(t, "https://api.kimi.com/coding/v1", kimiCodingAnthropic.GetOpenAIFormatBaseURL()) + + // chat_completions 协议下行为不变(凭证 base_url 原样返回) + ccAccount := &Account{ + Platform: PlatformDeepseek, Type: AccountTypeAPIKey, + Credentials: map[string]any{"base_url": "https://ds-relay.example.com"}, + } + require.Equal(t, "https://ds-relay.example.com", ccAccount.GetOpenAIFormatBaseURL()) +} + +// TestBuildUpstreamModelsRequest_AnthropicProtocol 模型同步使用协议感知 base。 +func TestBuildUpstreamModelsRequest_AnthropicProtocol(t *testing.T) { + t.Parallel() + svc := &AccountTestService{cfg: &config.Config{}} + account := &Account{ + ID: 1, Platform: PlatformZhipu, Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "sk-test", + "api_protocol": APIProtocolAnthropic, + "base_url": "https://open.bigmodel.cn/api/anthropic", + }, + } + req, err := svc.buildUpstreamModelsRequest(context.Background(), account) + require.NoError(t, err) + require.Equal(t, "https://open.bigmodel.cn/api/paas/v4/models", req.URL.String()) +} + +// TestBuildOpenAIResponsesURLForPlatform deepseek 官方端点为 /responses(无 /v1)。 +func TestBuildOpenAIResponsesURLForPlatform(t *testing.T) { + t.Parallel() + require.Equal(t, "https://api.deepseek.com/responses", buildOpenAIResponsesURLForPlatform(PlatformDeepseek, "https://api.deepseek.com")) + require.Equal(t, "https://api.openai.com/v1/responses", buildOpenAIResponsesURLForPlatform(PlatformOpenAI, "https://api.openai.com")) + require.Equal(t, "https://open.bigmodel.cn/api/paas/v4/responses", buildOpenAIResponsesURLForPlatform(PlatformZhipu, "https://open.bigmodel.cn/api/paas/v4")) +} + +// TestNormalizeDeepSeekResponsesRequestBody 无状态适配:强制 store=false、 +// 清除 previous_response_id;非 deepseek responses 协议原样返回。 +func TestNormalizeDeepSeekResponsesRequestBody(t *testing.T) { + t.Parallel() + + deepseekResponses := &Account{ + Platform: PlatformDeepseek, Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_protocol": APIProtocolResponses}, + } + body := []byte(`{"model":"deepseek-v4-pro","store":true,"previous_response_id":"resp_123","input":"hi"}`) + normalized := normalizeDeepSeekResponsesRequestBody(deepseekResponses, body) + require.False(t, gjson.GetBytes(normalized, "store").Bool()) + require.False(t, gjson.GetBytes(normalized, "previous_response_id").Exists()) + require.Equal(t, "deepseek-v4-pro", gjson.GetBytes(normalized, "model").String()) + + // 非 responses 协议(deepseek CC 账号)原样返回 + deepseekCC := &Account{Platform: PlatformDeepseek, Type: AccountTypeAPIKey} + require.Equal(t, string(body), string(normalizeDeepSeekResponsesRequestBody(deepseekCC, body))) + + // openai 账号原样返回 + openai := &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey} + require.Equal(t, string(body), string(normalizeDeepSeekResponsesRequestBody(openai, body))) +} + +// TestGetAnthropicAPIKeyAuthScheme_CNProvider CN 账号可经 extra 覆写鉴权方案, +// 默认保持 x-api-key。 +func TestGetAnthropicAPIKeyAuthScheme_CNProvider(t *testing.T) { + t.Parallel() + + zhipu := &Account{ + Platform: PlatformZhipu, Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_protocol": APIProtocolAnthropic}, + } + require.Equal(t, AnthropicAPIKeyAuthSchemeXAPIKey, zhipu.GetAnthropicAPIKeyAuthScheme()) + + zhipu.Extra = map[string]any{"anthropic_apikey_auth_scheme": "authorization_bearer"} + require.Equal(t, AnthropicAPIKeyAuthSchemeAuthorizationBearer, zhipu.GetAnthropicAPIKeyAuthScheme()) +} diff --git a/backend/internal/service/domain_constants.go b/backend/internal/service/domain_constants.go index 2d5311bb24..c0efa2b527 100644 --- a/backend/internal/service/domain_constants.go +++ b/backend/internal/service/domain_constants.go @@ -43,12 +43,58 @@ const ( PlatformGemini = domain.PlatformGemini PlatformAntigravity = domain.PlatformAntigravity PlatformGrok = domain.PlatformGrok - PlatformComposite = domain.PlatformComposite + // 国产 OpenAI 兼容供应商(与 grok 一样经 OpenAI 网关转发)。 + PlatformKimi = domain.PlatformKimi + PlatformZhipu = domain.PlatformZhipu + PlatformDeepseek = domain.PlatformDeepseek + PlatformComposite = domain.PlatformComposite // PlatformKiro is retained for unsupported-platform threshold tests and legacy // account rows. Scheduling-threshold evaluation never pauses kiro accounts. PlatformKiro = "kiro" ) +// 账号接入模式(国产供应商):按量付费 vs Coding Plan。 +const ( + AccountModePayG = domain.AccountModePayG + AccountModeCoding = domain.AccountModeCoding +) + +// 上游 API 协议(国产供应商):决定转发端点与格式,与接入模式正交。 +const ( + APIProtocolChatCompletions = domain.APIProtocolChatCompletions + APIProtocolAnthropic = domain.APIProtocolAnthropic + APIProtocolResponses = domain.APIProtocolResponses +) + +// 国产 OpenAI 兼容供应商各模式的默认 base_url。 +// 与前端 credentialsBuilder.ts 中的预设保持一致。 +const ( + DefaultKimiPayGBaseURL = "https://api.moonshot.cn/v1" + DefaultKimiCodingBaseURL = "https://api.kimi.com/coding/v1" + DefaultZhipuPayGBaseURL = "https://open.bigmodel.cn/api/paas/v4" + DefaultZhipuCodingBaseURL = "https://open.bigmodel.cn/api/coding/paas/v4" + DefaultDeepseekBaseURL = "https://api.deepseek.com" +) + +// 国产供应商 Anthropic 协议端点的默认 base_url(上游路径为 {base}/v1/messages)。 +// 与前端 credentialsBuilder.ts 中的预设保持一致。 +const ( + DefaultKimiPayGAnthropicBaseURL = "https://api.moonshot.cn/anthropic" + DefaultKimiCodingAnthropicBaseURL = "https://api.kimi.com/coding" + DefaultZhipuAnthropicBaseURL = "https://open.bigmodel.cn/api/anthropic" + DefaultDeepseekAnthropicBaseURL = "https://api.deepseek.com/anthropic" +) + +// IsCNProvider 报告 platform 是否为国产 OpenAI 兼容供应商(kimi/zhipu/deepseek)。 +func IsCNProvider(platform string) bool { + switch platform { + case PlatformKimi, PlatformZhipu, PlatformDeepseek: + return true + default: + return false + } +} + // AllowedQuotaPlatforms 是允许设置 user × platform quota 的平台列表(单一权威来源)。 // ent/schema/user_platform_quota.go 的 Validate 函数独立维护(构建期约束), // 若新增平台需同步修改该 schema。 @@ -58,14 +104,20 @@ var AllowedQuotaPlatforms = []string{ PlatformGemini, PlatformAntigravity, PlatformGrok, + PlatformKimi, + PlatformZhipu, + PlatformDeepseek, } // AllowedSchedulingThresholdPlatforms 是允许设置账号自动停调阈值的平台列表。 -// 仅 openai / anthropic / grok 有原生用量窗口可供评估;其他平台写入阈值无效果。 +// openai/anthropic/grok 有原生用量窗口;kimi/zhipu 的 Coding Plan 同样暴露 5h/weekly +// 滚动窗口,纳入阈值评估。deepseek 为余额型,走余额检测而非阈值。 var AllowedSchedulingThresholdPlatforms = []string{ PlatformOpenAI, PlatformAnthropic, PlatformGrok, + PlatformKimi, + PlatformZhipu, } // IsAllowedQuotaPlatform 报告 s 是否为合法的 quota platform 标识。 diff --git a/backend/internal/service/gateway_anthropic_apikey_passthrough_benchmark_test.go b/backend/internal/service/gateway_anthropic_apikey_passthrough_benchmark_test.go index 37fd709f84..a7b05f175c 100644 --- a/backend/internal/service/gateway_anthropic_apikey_passthrough_benchmark_test.go +++ b/backend/internal/service/gateway_anthropic_apikey_passthrough_benchmark_test.go @@ -14,13 +14,12 @@ func BenchmarkGatewayService_ParseSSEUsage_MessageStart(b *testing.B) { } func BenchmarkGatewayService_ParseSSEUsagePassthrough_MessageStart(b *testing.B) { - svc := &GatewayService{} data := `{"type":"message_start","message":{"usage":{"input_tokens":123,"cache_creation_input_tokens":45,"cache_read_input_tokens":6,"cached_tokens":6,"cache_creation":{"ephemeral_5m_input_tokens":20,"ephemeral_1h_input_tokens":25}}}}` b.ReportAllocs() b.ResetTimer() for i := 0; i < b.N; i++ { usage := &ClaudeUsage{} - svc.parseSSEUsagePassthrough(data, usage) + parseSSEUsagePassthrough(data, usage) } } @@ -36,13 +35,12 @@ func BenchmarkGatewayService_ParseSSEUsage_MessageDelta(b *testing.B) { } func BenchmarkGatewayService_ParseSSEUsagePassthrough_MessageDelta(b *testing.B) { - svc := &GatewayService{} data := `{"type":"message_delta","usage":{"output_tokens":456,"cache_creation_input_tokens":30,"cache_read_input_tokens":7,"cached_tokens":7,"cache_creation":{"ephemeral_5m_input_tokens":10,"ephemeral_1h_input_tokens":20}}}` b.ReportAllocs() b.ResetTimer() for i := 0; i < b.N; i++ { usage := &ClaudeUsage{} - svc.parseSSEUsagePassthrough(data, usage) + parseSSEUsagePassthrough(data, usage) } } diff --git a/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go b/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go index 4fa52cc9ce..eba08a113b 100644 --- a/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go +++ b/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go @@ -1256,11 +1256,10 @@ func TestExtractAnthropicSSEDataLine(t *testing.T) { } func TestGatewayService_ParseSSEUsagePassthrough_MessageStartFallbacks(t *testing.T) { - svc := &GatewayService{} usage := &ClaudeUsage{} data := `{"type":"message_start","message":{"usage":{"input_tokens":12,"cache_creation_input_tokens":0,"cache_read_input_tokens":0,"cached_tokens":9,"cache_creation":{"ephemeral_5m_input_tokens":3,"ephemeral_1h_input_tokens":4}}}}` - svc.parseSSEUsagePassthrough(data, usage) + parseSSEUsagePassthrough(data, usage) require.Equal(t, 12, usage.InputTokens) require.Equal(t, 9, usage.CacheReadInputTokens, "应兼容 cached_tokens 字段") @@ -1270,7 +1269,6 @@ func TestGatewayService_ParseSSEUsagePassthrough_MessageStartFallbacks(t *testin } func TestGatewayService_ParseSSEUsagePassthrough_MessageDeltaSelectiveOverwrite(t *testing.T) { - svc := &GatewayService{} usage := &ClaudeUsage{ InputTokens: 10, CacheCreation5mTokens: 2, @@ -1278,7 +1276,7 @@ func TestGatewayService_ParseSSEUsagePassthrough_MessageDeltaSelectiveOverwrite( } data := `{"type":"message_delta","usage":{"input_tokens":0,"output_tokens":5,"cache_creation_input_tokens":8,"cache_read_input_tokens":0,"cached_tokens":11,"cache_creation":{"ephemeral_5m_input_tokens":1,"ephemeral_1h_input_tokens":0}}}` - svc.parseSSEUsagePassthrough(data, usage) + parseSSEUsagePassthrough(data, usage) require.Equal(t, 10, usage.InputTokens, "message_delta 中 0 值不应覆盖已有 input_tokens") require.Equal(t, 5, usage.OutputTokens) @@ -1289,28 +1287,26 @@ func TestGatewayService_ParseSSEUsagePassthrough_MessageDeltaSelectiveOverwrite( } func TestGatewayService_ParseSSEUsagePassthrough_NoopCases(t *testing.T) { - svc := &GatewayService{} usage := &ClaudeUsage{InputTokens: 3} - svc.parseSSEUsagePassthrough("", usage) + parseSSEUsagePassthrough("", usage) require.Equal(t, 3, usage.InputTokens) - svc.parseSSEUsagePassthrough("[DONE]", usage) + parseSSEUsagePassthrough("[DONE]", usage) require.Equal(t, 3, usage.InputTokens) - svc.parseSSEUsagePassthrough("not-json", usage) + parseSSEUsagePassthrough("not-json", usage) require.Equal(t, 3, usage.InputTokens) // nil usage 不应 panic - svc.parseSSEUsagePassthrough(`{"type":"message_start"}`, nil) + parseSSEUsagePassthrough(`{"type":"message_start"}`, nil) } func TestGatewayService_ParseSSEUsagePassthrough_FallbackFromUsageNode(t *testing.T) { - svc := &GatewayService{} usage := &ClaudeUsage{} data := `{"type":"content_block_delta","usage":{"cached_tokens":6,"cache_creation":{"ephemeral_5m_input_tokens":2,"ephemeral_1h_input_tokens":1}}}` - svc.parseSSEUsagePassthrough(data, usage) + parseSSEUsagePassthrough(data, usage) require.Equal(t, 6, usage.CacheReadInputTokens) require.Equal(t, 3, usage.CacheCreationInputTokens) diff --git a/backend/internal/service/gateway_anthropic_passthrough.go b/backend/internal/service/gateway_anthropic_passthrough.go index c18d5bf047..def4d34648 100644 --- a/backend/internal/service/gateway_anthropic_passthrough.go +++ b/backend/internal/service/gateway_anthropic_passthrough.go @@ -546,7 +546,7 @@ func (s *GatewayService) handleStreamingResponseAnthropicAPIKeyPassthrough( ms := int(time.Since(startTime).Milliseconds()) firstTokenMs = &ms } - s.parseSSEUsagePassthrough(data, usage) + parseSSEUsagePassthrough(data, usage) } else { trimmed := strings.TrimSpace(line) if strings.HasPrefix(trimmed, "event:") && anthropicStreamEventIsTerminal(strings.TrimSpace(strings.TrimPrefix(trimmed, "event:")), "") { @@ -625,7 +625,9 @@ func extractAnthropicSSEDataLine(line string) (string, bool) { return line[start:], true } -func (s *GatewayService) parseSSEUsagePassthrough(data string, usage *ClaudeUsage) { +// parseSSEUsagePassthrough 从 Anthropic SSE data 行提取 usage(包级函数: +// Anthropic 平台 passthrough 与国产供应商原生 Anthropic 直通共用)。 +func parseSSEUsagePassthrough(data string, usage *ClaudeUsage) { if usage == nil || data == "" || data == "[DONE]" { return } @@ -730,8 +732,12 @@ func parseClaudeUsageFromResponseBody(body []byte) *ClaudeUsage { return usage } -func (s *GatewayService) invalidNonStreamingJSONFailoverError( +// invalidNonStreamingJSONFailoverError 把"上游 2xx 返回非 JSON body"归一为 +// failover 错误(包级函数:Anthropic 平台 passthrough 与国产供应商原生 +// Anthropic 直通共用)。 +func invalidNonStreamingJSONFailoverError( ctx context.Context, + rateLimitService *RateLimitService, resp *http.Response, account *Account, body []byte, @@ -759,11 +765,11 @@ func (s *GatewayService) invalidNonStreamingJSONFailoverError( parseErr, ) - if s.rateLimitService != nil && account != nil { + if rateLimitService != nil && account != nil { if len(requestedModel) > 0 { - s.rateLimitService.HandleUpstreamError(ctx, account, statusCode, resp.Header, body, requestedModel[0]) + rateLimitService.HandleUpstreamError(ctx, account, statusCode, resp.Header, body, requestedModel[0]) } else { - s.rateLimitService.HandleUpstreamError(ctx, account, statusCode, resp.Header, body) + rateLimitService.HandleUpstreamError(ctx, account, statusCode, resp.Header, body) } } @@ -798,7 +804,7 @@ func (s *GatewayService) handleNonStreamingResponseAnthropicAPIKeyPassthrough( if resp.StatusCode >= http.StatusOK && resp.StatusCode < http.StatusMultipleChoices { var raw json.RawMessage if err := json.Unmarshal(body, &raw); err != nil { - return nil, s.invalidNonStreamingJSONFailoverError(ctx, resp, account, body, err) + return nil, invalidNonStreamingJSONFailoverError(ctx, s.rateLimitService, resp, account, body, err) } } diff --git a/backend/internal/service/gateway_upstream_response.go b/backend/internal/service/gateway_upstream_response.go index 5a9fedf964..20af32314f 100644 --- a/backend/internal/service/gateway_upstream_response.go +++ b/backend/internal/service/gateway_upstream_response.go @@ -1395,7 +1395,7 @@ func (s *GatewayService) handleNonStreamingResponse(ctx context.Context, resp *h } if err := json.Unmarshal(body, &response); err != nil { if resp.StatusCode >= http.StatusOK && resp.StatusCode < http.StatusMultipleChoices { - return nil, s.invalidNonStreamingJSONFailoverError(ctx, resp, account, body, err, mappedModel) + return nil, invalidNonStreamingJSONFailoverError(ctx, s.rateLimitService, resp, account, body, err, mappedModel) } return nil, fmt.Errorf("parse response: %w", err) } diff --git a/backend/internal/service/openai_account_scheduler.go b/backend/internal/service/openai_account_scheduler.go index 6db0db766d..1daf81988c 100644 --- a/backend/internal/service/openai_account_scheduler.go +++ b/backend/internal/service/openai_account_scheduler.go @@ -381,7 +381,7 @@ func (s *defaultOpenAIAccountScheduler) Select( }() previousResponseID := strings.TrimSpace(req.PreviousResponseID) - if previousResponseID != "" && normalizeOpenAICompatiblePlatform(req.Platform) == PlatformOpenAI && + if previousResponseID != "" && NormalizeOpenAICompatiblePlatform(req.Platform) == PlatformOpenAI && (!req.StickyWeighted || !req.PreviousResponseCanMove) { selection, err := s.service.selectAccountByPreviousResponseIDForCapability( ctx, @@ -486,7 +486,7 @@ func (s *defaultOpenAIAccountScheduler) selectBySessionHash( _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) return nil, false, nil } - if shouldClearStickySession(account, req.RequestedModel) || account.Platform != normalizeOpenAICompatiblePlatform(req.Platform) || !account.IsOpenAICompatible() || !account.IsSchedulable() { + if shouldClearStickySession(account, req.RequestedModel) || account.Platform != NormalizeOpenAICompatiblePlatform(req.Platform) || !account.IsOpenAICompatible() || !account.IsSchedulable() { _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) return nil, false, nil } @@ -1406,7 +1406,7 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance( filterStats.exclude("not_schedulable") continue } - if account.Platform != normalizeOpenAICompatiblePlatform(req.Platform) || !account.IsOpenAICompatible() { + if account.Platform != NormalizeOpenAICompatiblePlatform(req.Platform) || !account.IsOpenAICompatible() { filterStats.exclude("platform_mismatch") continue } @@ -2125,7 +2125,7 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler( return selection, decision, err } // The circuit only ever quarantines PlatformOpenAI accounts. - if normalizeOpenAICompatiblePlatform(platform) != PlatformOpenAI { + if NormalizeOpenAICompatiblePlatform(platform) != PlatformOpenAI { return selection, decision, err } blocked := s.getOpenAIProxyStreamCircuit().activeBlockCount(time.Now()) @@ -2161,7 +2161,7 @@ func (s *OpenAIGatewayService) selectAccountWithSchedulerOnce( if requiredImageCapability == "" { ctx = s.withOpenAIProfitControlGate(ctx, groupID) } - platform = normalizeOpenAICompatiblePlatform(platform) + platform = NormalizeOpenAICompatiblePlatform(platform) decision := OpenAIAccountScheduleDecision{} scheduler := s.getOpenAIAccountScheduler(ctx) if scheduler == nil { diff --git a/backend/internal/service/openai_apikey_responses_probe.go b/backend/internal/service/openai_apikey_responses_probe.go index 79db519e6f..645a581694 100644 --- a/backend/internal/service/openai_apikey_responses_probe.go +++ b/backend/internal/service/openai_apikey_responses_probe.go @@ -119,7 +119,27 @@ func (s *AccountTestService) ProbeOpenAIAPIKeyResponsesSupport(ctx context.Conte logger.LegacyPrintf("service.openai_probe", "probe_load_account_failed: account_id=%d err=%v", accountID, err) return } - if account.Platform != PlatformOpenAI || account.Type != AccountTypeAPIKey { + if account.Type != AccountTypeAPIKey { + return + } + if account.IsCNProvider() { + // 国产 OpenAI 兼容上游(kimi/zhipu/deepseek)普遍仅支持 /v1/chat/completions, + // 不存在 /v1/responses 端点。直接落标 false 走 Chat Completions 直转,跳过网络探测。 + // 例外:deepseek 的 responses 协议账号(api_protocol=responses)使用官方原生 + // /responses 端点,落标 force_responses 强制走 Responses 路径。 + if account.GetAPIProtocol() == APIProtocolResponses { + _ = s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{ + openai_compat.ExtraKeyResponsesMode: string(openai_compat.ResponsesSupportModeForceResponses), + openai_compat.ExtraKeyResponsesSupported: true, + }) + return + } + _ = s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{ + openai_compat.ExtraKeyResponsesSupported: false, + }) + return + } + if account.Platform != PlatformOpenAI { // 仅 OpenAI APIKey 账号需要探测;其他账号类型无能力差异。 return } diff --git a/backend/internal/service/openai_embeddings.go b/backend/internal/service/openai_embeddings.go index c0c5900297..e68017a453 100644 --- a/backend/internal/service/openai_embeddings.go +++ b/backend/internal/service/openai_embeddings.go @@ -46,11 +46,13 @@ func (s *OpenAIGatewayService) ForwardEmbeddings( zap.String("upstream_model", upstreamModel), ) - apiKey := account.GetOpenAIApiKey() + apiKey := strings.TrimSpace(account.GetOpenAIProtocolAPIKey()) if apiKey == "" { return nil, fmt.Errorf("account %d missing api_key", account.ID) } - baseURL := account.GetOpenAIBaseURL() + // 协议感知:Anthropic 协议账号的凭证 base_url 指向 /anthropic 端点, + // embeddings 需使用 OpenAI 格式 base。 + baseURL := account.GetOpenAIFormatBaseURL() if baseURL == "" { baseURL = "https://api.openai.com" } diff --git a/backend/internal/service/openai_gateway_anthropic_native_pump.go b/backend/internal/service/openai_gateway_anthropic_native_pump.go new file mode 100644 index 0000000000..8a19677705 --- /dev/null +++ b/backend/internal/service/openai_gateway_anthropic_native_pump.go @@ -0,0 +1,126 @@ +package service + +// 国产供应商 Anthropic 协议转换路径的上游 SSE 行泵。 +// +// 这两条转换链(CC×anthropic / Responses×anthropic)的上游 ctx 是 +// WithoutCancel(detachStreamUpstreamContext)、http.Client 无整体 Timeout, +// 客户端断开后的排水阶段若上游挂住 SSE(不发数据也不断连), +// scanner.Scan() 将永久阻塞:goroutine 钉死、resp.Body 不归还、连接池位 +// 被占用、usage 永不落库。 +// +// 与 handleAnthropicStreamingResponse / readOpenAICompatBufferedTerminal 的 +// 同类排水一致,本泵用 gateway.stream_data_interval_timeout(默认 180s)作为 +// 逐行读间隔上限,超时即向调用方返回 errAnthropicNativeStreamIdle,由调用方 +// 关闭 resp.Body 解除阻塞的读并结束排水。 + +import ( + "bufio" + "errors" + "io" + "time" +) + +// errAnthropicNativeStreamIdle 表示上游流读间隔超时(见上方文件注释)。 +var errAnthropicNativeStreamIdle = errors.New("stream data interval timeout") + +// anthropicNativeLineEvent 是行泵交付的单次读取结果:line 为一行 SSE 文本, +// err 为 scanner 读错误(流自然结束时 next 返回 io.EOF,不经过本字段)。 +type anthropicNativeLineEvent struct { + line string + err error +} + +// anthropicNativeLinePump 以独立 goroutine 泵送 scanner 的行,并对逐行到达 +// 间隔施加 interval 上限(<=0 表示禁用,保持无界读的旧行为)。 +type anthropicNativeLinePump struct { + events chan anthropicNativeLineEvent + done chan struct{} + timer *time.Timer + interval time.Duration +} + +// newAnthropicNativeLinePump 启动泵 goroutine;调用方 defer pump.stop()。 +func newAnthropicNativeLinePump(scanner *bufio.Scanner, interval time.Duration) *anthropicNativeLinePump { + p := &anthropicNativeLinePump{ + events: make(chan anthropicNativeLineEvent, 16), + done: make(chan struct{}), + interval: interval, + } + if interval > 0 { + p.timer = time.NewTimer(interval) + } + go func() { + defer close(p.events) + for scanner.Scan() { + select { + case p.events <- anthropicNativeLineEvent{line: scanner.Text()}: + case <-p.done: + return + } + } + if err := scanner.Err(); err != nil { + select { + case p.events <- anthropicNativeLineEvent{err: err}: + case <-p.done: + } + } + }() + return p +} + +// next 阻塞返回下一行。返回 io.EOF 表示上游正常收流;errAnthropicNativeStreamIdle +// 表示 interval 内无任何数据到达(计时从收到上一行时起算,事件处理耗时不算入, +// 与 readOpenAICompatBufferedTerminal 的 resetTimeout 语义一致)。 +func (p *anthropicNativeLinePump) next() (string, error) { + var timeoutCh <-chan time.Time + if p.timer != nil { + timeoutCh = p.timer.C + } + select { + case ev, ok := <-p.events: + if !ok { + return "", io.EOF + } + p.resetTimer() + return ev.line, ev.err + case <-timeoutCh: + return "", errAnthropicNativeStreamIdle + } +} + +// resetTimer 在收到一行后重启间隔计时器。 +func (p *anthropicNativeLinePump) resetTimer() { + if p.timer == nil { + return + } + if !p.timer.Stop() { + select { + case <-p.timer.C: + default: + } + } + p.timer.Reset(p.interval) +} + +// stop 终止泵 goroutine。注意:goroutine 若正阻塞在 scanner.Read 上,需由 +// 调用方关闭 resp.Body(间隔超时分支已做)才能真正退出。 +func (p *anthropicNativeLinePump) stop() { + close(p.done) + if p.timer != nil { + if !p.timer.Stop() { + select { + case <-p.timer.C: + default: + } + } + } +} + +// anthropicNativeStreamInterval 返回本组转换路径适用的读间隔上限; +// gateway.stream_data_interval_timeout <= 0 时视为禁用。 +func (s *OpenAIGatewayService) anthropicNativeStreamInterval() time.Duration { + if s.cfg != nil && s.cfg.Gateway.StreamDataIntervalTimeout > 0 { + return time.Duration(s.cfg.Gateway.StreamDataIntervalTimeout) * time.Second + } + return 0 +} diff --git a/backend/internal/service/openai_gateway_anthropic_native_pump_test.go b/backend/internal/service/openai_gateway_anthropic_native_pump_test.go new file mode 100644 index 0000000000..1e9e1b39a6 --- /dev/null +++ b/backend/internal/service/openai_gateway_anthropic_native_pump_test.go @@ -0,0 +1,244 @@ +package service + +// 国产供应商 Anthropic 协议转换路径的上游读间隔超时回归测试(B3): +// 上游挂住 SSE(不发数据也不断连)时,CC×anthropic / Responses×anthropic +// 的读循环必须按 gateway.stream_data_interval_timeout 结束,而不是永久阻塞。 + +import ( + "bufio" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" + "github.com/gin-gonic/gin" +) + +func newNativeAnthropicHangTestService(intervalSec int) *OpenAIGatewayService { + return &OpenAIGatewayService{ + cfg: &config.Config{ + Gateway: config.GatewayConfig{ + StreamDataIntervalTimeout: intervalSec, + MaxLineSize: defaultMaxLineSize, + }, + }, + } +} + +func newHangingUpstreamResponse() (*http.Response, *io.PipeReader, *io.PipeWriter) { + pr, pw := io.Pipe() + return &http.Response{StatusCode: http.StatusOK, Body: pr, Header: http.Header{}}, pr, pw +} + +// miniAnthropicSSEStream 是一段最小可转换的 Anthropic 事件流。 +func miniAnthropicSSEStream() string { + return strings.Join([]string{ + "event: message_start", + `data: {"type":"message_start","message":{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"glm-4.7","usage":{"input_tokens":10,"output_tokens":1}}}`, + "", + "event: content_block_start", + `data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}`, + "", + "event: content_block_delta", + `data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello"}}`, + "", + "event: content_block_stop", + `data: {"type":"content_block_stop","index":0}`, + "", + "event: message_delta", + `data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":5}}`, + "", + "event: message_stop", + `data: {"type":"message_stop"}`, + "", + "", + }, "\n") +} + +func TestAnthropicNativeLinePump_TimesOutWithoutData(t *testing.T) { + pr, _ := io.Pipe() + scanner := bufio.NewScanner(pr) + defer func() { _ = pr.Close() }() + + pump := newAnthropicNativeLinePump(scanner, 50*time.Millisecond) + defer pump.stop() + + start := time.Now() + _, err := pump.next() + if err == nil || !strings.Contains(err.Error(), "stream data interval timeout") { + t.Fatalf("expected interval timeout, got %v", err) + } + if elapsed := time.Since(start); elapsed > 5*time.Second { + t.Fatalf("timeout not respected: %v", elapsed) + } +} + +func TestAnthropicNativeLinePump_DataResetsTimer(t *testing.T) { + pr, pw := io.Pipe() + scanner := bufio.NewScanner(pr) + pump := newAnthropicNativeLinePump(scanner, 1*time.Second) + defer pump.stop() + + go func() { + _, _ = pw.Write([]byte("event: ping\n")) + // 保持流打开且不再发数据:第二次 next 必须超时。 + time.Sleep(3 * time.Second) + _ = pw.Close() + }() + defer func() { _ = pr.Close() }() + + line, err := pump.next() + if err != nil || line != "event: ping" { + t.Fatalf("expected first line, got %q err=%v", line, err) + } + + start := time.Now() + _, err = pump.next() + if err == nil || !strings.Contains(err.Error(), "stream data interval timeout") { + t.Fatalf("expected interval timeout after data stops, got %v", err) + } + if elapsed := time.Since(start); elapsed > 5*time.Second { + t.Fatalf("timeout not respected: %v", elapsed) + } +} + +func TestCCStreamingFromNativeAnthropic_HangTimesOut(t *testing.T) { + gin.SetMode(gin.TestMode) + svc := newNativeAnthropicHangTestService(1) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/", nil) + + resp, pr, pw := newHangingUpstreamResponse() + start := time.Now() + res, err := svc.handleCCStreamingFromNativeAnthropic(resp, c, "glm-4.7", "glm-4.7", "glm-4.7", nil, start, true) + _ = pw.Close() + _ = pr.Close() + + if err == nil || !strings.Contains(err.Error(), "stream data interval timeout") { + t.Fatalf("expected stream timeout error, got %v", err) + } + if res == nil { + t.Fatalf("expected result carrying accumulated usage") + } + if elapsed := time.Since(start); elapsed > 5*time.Second { + t.Fatalf("handler did not respect interval bound: %v", elapsed) + } +} + +func TestCCBufferedFromNativeAnthropic_HangTimesOut(t *testing.T) { + gin.SetMode(gin.TestMode) + svc := newNativeAnthropicHangTestService(1) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/", nil) + + resp, pr, pw := newHangingUpstreamResponse() + start := time.Now() + _, err := svc.handleCCBufferedFromNativeAnthropic(resp, c, "glm-4.7", "glm-4.7", "glm-4.7", nil, start) + _ = pw.Close() + _ = pr.Close() + + if err == nil || !strings.Contains(err.Error(), "stream data interval timeout") { + t.Fatalf("expected stream timeout error, got %v", err) + } + if elapsed := time.Since(start); elapsed > 5*time.Second { + t.Fatalf("handler did not respect interval bound: %v", elapsed) + } + if !strings.Contains(rec.Body.String(), "Upstream stream data interval timeout") { + t.Fatalf("expected 502 error body, got %q", rec.Body.String()) + } +} + +func TestResponsesStreamingFromNativeAnthropic_HangTimesOut(t *testing.T) { + gin.SetMode(gin.TestMode) + svc := newNativeAnthropicHangTestService(1) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/", nil) + + resp, pr, pw := newHangingUpstreamResponse() + start := time.Now() + res, err := svc.handleResponsesStreamingFromNativeAnthropic(resp, c, "glm-4.7", "glm-4.7", "glm-4.7", nil, start, apicompat.ResponsesClientToolMapping{}) + _ = pw.Close() + _ = pr.Close() + + if err == nil || !strings.Contains(err.Error(), "stream data interval timeout") { + t.Fatalf("expected stream timeout error, got %v", err) + } + if res == nil { + t.Fatalf("expected result carrying accumulated usage") + } + if elapsed := time.Since(start); elapsed > 5*time.Second { + t.Fatalf("handler did not respect interval bound: %v", elapsed) + } +} + +func TestCCStreamingFromNativeAnthropic_HappyPathStillConverts(t *testing.T) { + gin.SetMode(gin.TestMode) + svc := newNativeAnthropicHangTestService(5) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/", nil) + + resp, pr, pw := newHangingUpstreamResponse() + go func() { + _, _ = pw.Write([]byte(miniAnthropicSSEStream())) + _ = pw.Close() + }() + defer func() { _ = pr.Close() }() + + res, err := svc.handleCCStreamingFromNativeAnthropic(resp, c, "glm-4.7", "glm-4.7", "glm-4.7", nil, time.Now(), true) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if res == nil { + t.Fatalf("expected result") + } + body := rec.Body.String() + if !strings.Contains(body, "Hello") { + t.Fatalf("expected converted text chunk, got %q", body) + } + if !strings.Contains(body, "data: [DONE]") { + t.Fatalf("expected [DONE] terminator, got %q", body) + } +} + +func TestCCBufferedFromNativeAnthropic_HappyPathStillConverts(t *testing.T) { + gin.SetMode(gin.TestMode) + svc := newNativeAnthropicHangTestService(5) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/", nil) + + resp, pr, pw := newHangingUpstreamResponse() + go func() { + _, _ = pw.Write([]byte(miniAnthropicSSEStream())) + _ = pw.Close() + }() + defer func() { _ = pr.Close() }() + + res, err := svc.handleCCBufferedFromNativeAnthropic(resp, c, "glm-4.7", "glm-4.7", "glm-4.7", nil, time.Now()) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if res == nil { + t.Fatalf("expected result") + } + body := rec.Body.String() + if !strings.Contains(body, "Hello") { + t.Fatalf("expected converted text in buffered response, got %q", body) + } + if res.Usage.InputTokens != 10 || res.Usage.OutputTokens != 5 { + t.Fatalf("expected usage 10/5, got %+v", res.Usage) + } +} diff --git a/backend/internal/service/openai_gateway_cc_pipeline.go b/backend/internal/service/openai_gateway_cc_pipeline.go index 8080d49e5c..ce9dd54d37 100644 --- a/backend/internal/service/openai_gateway_cc_pipeline.go +++ b/backend/internal/service/openai_gateway_cc_pipeline.go @@ -150,7 +150,7 @@ func (s *OpenAIGatewayService) openAIChatCompletionsTargetURL(account *Account) // resolveCCFallbackTarget 解析两条 CC 回退路径共用的账号凭证与上游端点 // (回退路径仅面向 APIKey 账号,凭证恒为 openai api_key)。 func (s *OpenAIGatewayService) resolveCCFallbackTarget(account *Account) (apiKey string, targetURL string, err error) { - apiKey = account.GetOpenAIApiKey() + apiKey = strings.TrimSpace(account.GetOpenAIProtocolAPIKey()) if apiKey == "" { return "", "", fmt.Errorf("account %d missing api_key", account.ID) } diff --git a/backend/internal/service/openai_gateway_chat_completions.go b/backend/internal/service/openai_gateway_chat_completions.go index 87973304c0..fed651bd70 100644 --- a/backend/internal/service/openai_gateway_chat_completions.go +++ b/backend/internal/service/openai_gateway_chat_completions.go @@ -87,6 +87,14 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions( return s.forwardAsRawChatCompletions(ctx, c, account, body, defaultMappedModel) } + // 入口分流(国产供应商 Anthropic 协议):上游为供应商原生 Anthropic 端点, + // CC 入站请求经 CC→Responses→Anthropic 转换链直通该端点。必须先于 + // ShouldUseResponsesAPI 分流:该类账号经 probe 落标 + // openai_responses_supported=false,会先命中下方的 CC 直转分支。 + if account.IsAnthropicProtocol() { + return s.forwardChatCompletionsViaNativeAnthropic(ctx, c, account, body, defaultMappedModel) + } + // 入口分流:APIKey 账号 + 强制或已探测确认上游不支持 Responses,走 CC 直转。 // 自动模式下标记缺失(未探测)按"现状即证据"原则继续走下方原 Responses 转换路径。 if account.Type == AccountTypeAPIKey && !openai_compat.ShouldUseResponsesAPI(account.Extra) { diff --git a/backend/internal/service/openai_gateway_chat_completions_anthropic_native.go b/backend/internal/service/openai_gateway_chat_completions_anthropic_native.go new file mode 100644 index 0000000000..b41bc644c6 --- /dev/null +++ b/backend/internal/service/openai_gateway_chat_completions_anthropic_native.go @@ -0,0 +1,486 @@ +package service + +// 国产供应商 Anthropic 协议账号的 CC 入站反向路径。 +// +// 客户端说 OpenAI Chat Completions、上游是供应商原生 Anthropic 端点 +// (api_protocol=anthropic)时的交叉组合:请求 CC→Responses→Anthropic 转换, +// 响应 Anthropic→Responses→CC 转换。转换链与 Anthropic 平台的 +// gateway_forward_as_chat_completions.go 完全一致(复用同一组 apicompat +// 状态机),仅上游发送/错误处理对齐 OpenAI 网关语义。 + +import ( + "bufio" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" + "github.com/gin-gonic/gin" + "go.uber.org/zap" +) + +// forwardChatCompletionsViaNativeAnthropic serves OpenAI /v1/chat/completions +// clients through a CN provider's native Anthropic endpoint. +// +// Conversion chain: +// +// Request: Chat Completions → Responses → Anthropic (chained) +// Response: Anthropic events → Responses events → CC chunks (chained state machines) +func (s *OpenAIGatewayService) forwardChatCompletionsViaNativeAnthropic( + ctx context.Context, + c *gin.Context, + account *Account, + body []byte, + defaultMappedModel string, +) (*OpenAIForwardResult, error) { + startTime := time.Now() + + // 1. Parse Chat Completions request + var ccReq apicompat.ChatCompletionsRequest + if err := json.Unmarshal(body, &ccReq); err != nil { + writeChatCompletionsError(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") + return nil, fmt.Errorf("parse chat completions request: %w", err) + } + originalModel := ccReq.Model + if strings.TrimSpace(originalModel) == "" { + writeChatCompletionsError(c, http.StatusBadRequest, "invalid_request_error", "model is required") + return nil, fmt.Errorf("missing model in request") + } + clientStream := ccReq.Stream + includeUsage := ccReq.StreamOptions != nil && ccReq.StreamOptions.IncludeUsage + + // 2. Convert CC → Responses → Anthropic (chained conversion) + responsesReq, err := apicompat.ChatCompletionsToResponses(&ccReq) + if err != nil { + writeChatCompletionsError(c, http.StatusBadRequest, "invalid_request_error", "Failed to convert request") + return nil, fmt.Errorf("convert chat completions to responses: %w", err) + } + anthropicReq, err := apicompat.ResponsesToAnthropicRequest(responsesReq) + if err != nil { + writeChatCompletionsError(c, http.StatusBadRequest, "invalid_request_error", "Failed to convert request") + return nil, fmt.Errorf("convert responses to anthropic: %w", err) + } + + // 3. Model mapping(OpenAI 网关统一入口的映射语义) + billingModel := resolveOpenAIForwardModel(account, originalModel, defaultMappedModel) + upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel) + anthropicReq.Model = upstreamModel + + // 4. Force upstream streaming(客户端原始终决定响应格式; + // 上游恒为流式,非流式由缓冲路径组装)。 + anthropicReq.Stream = true + reqStream := true + + logger.L().Debug("openai chat_completions: forwarding via native anthropic endpoint", + zap.Int64("account_id", account.ID), + zap.String("original_model", originalModel), + zap.String("billing_model", billingModel), + zap.String("upstream_model", upstreamModel), + zap.Bool("client_stream", clientStream), + ) + + anthropicBody, err := json.Marshal(anthropicReq) + if err != nil { + return nil, fmt.Errorf("marshal anthropic request: %w", err) + } + + // 与 /v1/messages 直通路径相同的 pre-filter。 + anthropicBody = StripEmptyTextBlocks(anthropicBody) + anthropicBody = FilterWebSearchHistoryBlocks(anthropicBody, upstreamModel) + anthropicBody = enforceCacheControlLimit(anthropicBody) + + apiKey := strings.TrimSpace(account.GetOpenAIProtocolAPIKey()) + if apiKey == "" { + return nil, fmt.Errorf("account %d missing api_key", account.ID) + } + targetURL, err := s.nativeAnthropicTargetURL(account) + if err != nil { + return nil, err + } + + proxyURL := "" + if account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + + upstreamCtx, releaseUpstreamCtx := detachStreamUpstreamContext(ctx, reqStream) + upstreamReq, _, err := s.buildNativeAnthropicUpstreamRequest(upstreamCtx, c, account, anthropicBody, apiKey, targetURL) + releaseUpstreamCtx() + if err != nil { + return nil, fmt.Errorf("build upstream request: %w", err) + } + + resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) + if err != nil { + return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, true) + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode >= 400 { + respBody, upstreamMsg := s.readOpenAIUpstreamError(resp) + if foErr := s.failoverOpenAIUpstreamHTTPError(ctx, c, account, resp, respBody, upstreamMsg, upstreamModel); foErr != nil { + return nil, foErr + } + writeChatCompletionsError(c, mapUpstreamStatusCode(resp.StatusCode), "server_error", upstreamMsg) + return nil, fmt.Errorf("upstream error: %d %s", resp.StatusCode, upstreamMsg) + } + + reasoningEffort := extractCCReasoningEffortFromBody(body) + reasoningEffort = ApplyThinkingEnabledFallback(reasoningEffort, body, billingModel) + + if clientStream { + return s.handleCCStreamingFromNativeAnthropic(resp, c, originalModel, billingModel, upstreamModel, reasoningEffort, startTime, includeUsage) + } + return s.handleCCBufferedFromNativeAnthropic(resp, c, originalModel, billingModel, upstreamModel, reasoningEffort, startTime) +} + +// handleCCBufferedFromNativeAnthropic reads Anthropic SSE events, assembles the +// full response, then converts Anthropic → Responses → Chat Completions. +func (s *OpenAIGatewayService) handleCCBufferedFromNativeAnthropic( + resp *http.Response, + c *gin.Context, + originalModel string, + billingModel string, + upstreamModel string, + reasoningEffort *string, + startTime time.Time, +) (*OpenAIForwardResult, error) { + requestID := resp.Header.Get("x-request-id") + + scanner := bufio.NewScanner(resp.Body) + maxLineSize := defaultMaxLineSize + if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { + maxLineSize = s.cfg.Gateway.MaxLineSize + } + scanner.Buffer(make([]byte, 0, 64*1024), maxLineSize) + + var finalResp *apicompat.AnthropicResponse + var usage ClaudeUsage + + // 读间隔上限:上游挂住 SSE 时中止组装(缓冲路径尚未提交响应头,可回 502)。 + streamInterval := s.anthropicNativeStreamInterval() + pump := newAnthropicNativeLinePump(scanner, streamInterval) + defer pump.stop() + + logReadErr := func(err error) { + if !errors.Is(err, io.EOF) && !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) { + logger.L().Warn("openai cc via native anthropic buffered: read error", + zap.Error(err), + zap.String("request_id", requestID), + ) + } + } + onIdle := func() (*OpenAIForwardResult, error) { + _ = resp.Body.Close() + logger.L().Warn("openai cc via native anthropic buffered: data interval timeout", + zap.String("request_id", requestID), + zap.Duration("interval", streamInterval), + ) + writeChatCompletionsError(c, http.StatusBadGateway, "server_error", "Upstream stream data interval timeout") + return nil, fmt.Errorf("stream data interval timeout") + } + + for { + line, rerr := pump.next() + if rerr != nil { + if errors.Is(rerr, errAnthropicNativeStreamIdle) { + return onIdle() + } + logReadErr(rerr) + break + } + // SSE 规范允许 `event:xxx`(冒号后无空格):Kimi 等 Anthropic 兼容上游 + // 返回紧凑格式,严格匹配 "event: " 会丢弃全部事件(#4653 同根因)。 + if _, ok := extractOpenAISSEEventLine(line); !ok { + continue + } + + dataLine, rerr := pump.next() + if rerr != nil { + if errors.Is(rerr, errAnthropicNativeStreamIdle) { + return onIdle() + } + logReadErr(rerr) + break + } + payload, ok := extractOpenAISSEDataLine(dataLine) + if !ok { + continue + } + + var event apicompat.AnthropicStreamEvent + if err := json.Unmarshal([]byte(payload), &event); err != nil { + continue + } + + if event.Type == "message_start" && event.Message != nil { + finalResp = event.Message + mergeAnthropicUsage(&usage, event.Message.Usage) + } + if event.Type == "message_delta" { + if event.Usage != nil { + mergeAnthropicUsage(&usage, *event.Usage) + } + if event.Delta != nil && event.Delta.StopReason != "" && finalResp != nil { + finalResp.StopReason = apicompat.AnthropicStopReasonPtr(event.Delta.StopReason) + } + } + if event.Type == "content_block_start" && event.ContentBlock != nil && finalResp != nil { + finalResp.Content = append(finalResp.Content, *event.ContentBlock) + } + if event.Type == "content_block_delta" && event.Delta != nil && finalResp != nil && event.Index != nil { + idx := *event.Index + if idx < len(finalResp.Content) { + switch event.Delta.Type { + case "text_delta": + finalResp.Content[idx].Text += event.Delta.Text + case "thinking_delta": + finalResp.Content[idx].Thinking += event.Delta.Thinking + case "input_json_delta": + finalResp.Content[idx].Input = appendRawJSON(finalResp.Content[idx].Input, event.Delta.PartialJSON) + } + } + } + } + + if finalResp == nil { + writeChatCompletionsError(c, http.StatusBadGateway, "server_error", "Upstream stream ended without a response") + return nil, fmt.Errorf("upstream stream ended without response") + } + + if usage.InputTokens > 0 || usage.OutputTokens > 0 { + finalResp.Usage = apicompat.AnthropicUsage{ + InputTokens: usage.InputTokens, + OutputTokens: usage.OutputTokens, + CacheCreationInputTokens: usage.CacheCreationInputTokens, + CacheReadInputTokens: usage.CacheReadInputTokens, + } + } + + responsesResp := apicompat.AnthropicToResponsesResponse(finalResp) + ccResp := apicompat.ResponsesToChatCompletions(responsesResp, originalModel) + + if s.responseHeaderFilter != nil { + responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) + } + // 非流式响应必须是 application/json(上游被强制流式,透传头会污染)。 + c.Writer.Header().Set("Content-Type", "application/json; charset=utf-8") + if respBytes, err := json.Marshal(ccResp); err == nil { + respBytes = reverseToolNamesIfPresent(c, respBytes) + c.Data(http.StatusOK, "application/json; charset=utf-8", respBytes) + } else { + c.JSON(http.StatusOK, ccResp) + } + + return &OpenAIForwardResult{ + RequestID: requestID, + Usage: claudeUsageToOpenAIUsage(&usage), + Model: originalModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, + UpstreamEndpoint: "/v1/messages", + ReasoningEffort: reasoningEffort, + Stream: false, + Duration: time.Since(startTime), + }, nil +} + +// handleCCStreamingFromNativeAnthropic reads Anthropic SSE events, converts each +// to Responses events, then to Chat Completions chunks, and writes them. +func (s *OpenAIGatewayService) handleCCStreamingFromNativeAnthropic( + resp *http.Response, + c *gin.Context, + originalModel string, + billingModel string, + upstreamModel string, + reasoningEffort *string, + startTime time.Time, + includeUsage bool, +) (*OpenAIForwardResult, error) { + requestID := resp.Header.Get("x-request-id") + + if s.responseHeaderFilter != nil { + responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) + } + c.Writer.Header().Set("Content-Type", "text/event-stream") + c.Writer.Header().Set("Cache-Control", "no-cache") + c.Writer.Header().Set("Connection", "keep-alive") + c.Writer.Header().Set("X-Accel-Buffering", "no") + c.Writer.WriteHeader(http.StatusOK) + + anthState := apicompat.NewAnthropicEventToResponsesState() + anthState.Model = originalModel + ccState := apicompat.NewResponsesEventToChatState() + ccState.Model = originalModel + ccState.IncludeUsage = includeUsage + + var usage ClaudeUsage + var firstTokenMs *int + firstChunk := true + clientDisconnected := false + + scanner := bufio.NewScanner(resp.Body) + maxLineSize := defaultMaxLineSize + if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { + maxLineSize = s.cfg.Gateway.MaxLineSize + } + scanner.Buffer(make([]byte, 0, 64*1024), maxLineSize) + + resultWithUsage := func() *OpenAIForwardResult { + return &OpenAIForwardResult{ + RequestID: requestID, + Usage: claudeUsageToOpenAIUsage(&usage), + Model: originalModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, + UpstreamEndpoint: "/v1/messages", + ReasoningEffort: reasoningEffort, + Stream: true, + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, + ClientDisconnect: clientDisconnected, + } + } + + // 读间隔上限:上游挂住 SSE(不发数据也不断连)时结束排水。上游 ctx 为 + // WithoutCancel 且 http.Client 无整体 Timeout,无此界限则客户端断开后 + // scanner.Scan() 永久阻塞(见 anthropic native pump 文件注释)。 + streamInterval := s.anthropicNativeStreamInterval() + pump := newAnthropicNativeLinePump(scanner, streamInterval) + defer pump.stop() + + logReadErr := func(err error) { + if !errors.Is(err, io.EOF) && !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) { + logger.L().Warn("openai cc via native anthropic stream: read error", + zap.Error(err), + zap.String("request_id", requestID), + ) + } + } + // onIdle 关闭上游连接(解除阻塞的读、归还连接池位),并按已累计 usage + // 返回——与 messages 主路径 "stream usage incomplete after timeout" 同语义。 + onIdle := func() (*OpenAIForwardResult, error) { + _ = resp.Body.Close() + if !clientDisconnected { + logger.L().Warn("openai cc via native anthropic stream: data interval timeout", + zap.String("request_id", requestID), + zap.Duration("interval", streamInterval), + ) + } + return resultWithUsage(), fmt.Errorf("stream data interval timeout") + } + + writeChunk := func(chunk apicompat.ChatCompletionsChunk) bool { + if clientDisconnected { + return false // 已断开:不再写客户端,只排水上游累计 usage + } + sse, err := apicompat.ChatChunkToSSE(chunk) + if err != nil { + return false + } + out := string(reverseToolNamesIfPresent(c, []byte(sse))) + if _, err := fmt.Fprint(c.Writer, out); err != nil { + clientDisconnected = true + return false + } + return false + } + + processAnthropicEvent := func(event *apicompat.AnthropicStreamEvent) bool { + if firstChunk { + firstChunk = false + ms := int(time.Since(startTime).Milliseconds()) + firstTokenMs = &ms + } + + // usage 恒累计(含客户端断开后的排水阶段,payg 上游照常计费)。 + if event.Type == "message_delta" && event.Usage != nil { + mergeAnthropicUsage(&usage, *event.Usage) + } + if event.Type == "message_start" && event.Message != nil { + mergeAnthropicUsage(&usage, event.Message.Usage) + } + + // 客户端已断开:跳过转换与写出,继续读上游直到流结束(usage 完整、 + // 连接及时归还),不再提前 return。 + if clientDisconnected { + return false + } + + responsesEvents := apicompat.AnthropicEventToResponsesEvents(event, anthState) + for _, resEvt := range responsesEvents { + ccChunks := apicompat.ResponsesEventToChatChunks(&resEvt, ccState) + for _, chunk := range ccChunks { + writeChunk(chunk) + } + } + if len(responsesEvents) > 0 { + c.Writer.Flush() + } + return false + } + + for { + line, rerr := pump.next() + if rerr != nil { + if errors.Is(rerr, errAnthropicNativeStreamIdle) { + return onIdle() + } + logReadErr(rerr) + break + } + if _, ok := extractOpenAISSEEventLine(line); !ok { + continue + } + + dataLine, rerr := pump.next() + if rerr != nil { + if errors.Is(rerr, errAnthropicNativeStreamIdle) { + return onIdle() + } + // EOF / 读错误:事件行后流终止,进入 finalize。 + logReadErr(rerr) + break + } + payload, ok := extractOpenAISSEDataLine(dataLine) + if !ok { + continue + } + + var event apicompat.AnthropicStreamEvent + if err := json.Unmarshal([]byte(payload), &event); err != nil { + continue + } + + if processAnthropicEvent(&event) { + return resultWithUsage(), nil + } + } + + // Finalize both state machines(客户端已断开时仍执行,保证 usage 汇总完整)。 + finalResEvents := apicompat.FinalizeAnthropicResponsesStream(anthState) + for _, resEvt := range finalResEvents { + ccChunks := apicompat.ResponsesEventToChatChunks(&resEvt, ccState) + for _, chunk := range ccChunks { + writeChunk(chunk) //nolint:errcheck + } + } + finalCCChunks := apicompat.FinalizeResponsesChatStream(ccState) + for _, chunk := range finalCCChunks { + writeChunk(chunk) //nolint:errcheck + } + + if !clientDisconnected { + fmt.Fprint(c.Writer, "data: [DONE]\n\n") //nolint:errcheck + c.Writer.Flush() + } + + return resultWithUsage(), nil +} diff --git a/backend/internal/service/openai_gateway_count_tokens.go b/backend/internal/service/openai_gateway_count_tokens.go index e56b33df45..f20b45b526 100644 --- a/backend/internal/service/openai_gateway_count_tokens.go +++ b/backend/internal/service/openai_gateway_count_tokens.go @@ -43,6 +43,12 @@ type openAIInputTokensCountPrepared struct { // locally. Grok does not expose a compatible token-counting endpoint, so this // path deliberately avoids account selection, credentials, and upstream calls. func EstimateGrokCountTokens(body []byte) (int, error) { + return estimateAnthropicCountTokensLocally(body) +} + +// estimateAnthropicCountTokensLocally 走 Anthropic→Responses→tiktoken 链本地估算 +// count_tokens,不发任何上游请求(上游无兼容端点的平台使用)。 +func estimateAnthropicCountTokensLocally(body []byte) (int, error) { var anthropicReq apicompat.AnthropicRequest if err := json.Unmarshal(body, &anthropicReq); err != nil { return 0, fmt.Errorf("parse anthropic count_tokens request: %w", err) @@ -64,7 +70,7 @@ func EstimateGrokCountTokens(body []byte) (int, error) { ToolChoice: responsesReq.ToolChoice, }) if err != nil { - return 0, fmt.Errorf("estimate grok input tokens: %w", err) + return 0, fmt.Errorf("estimate input tokens: %w", err) } if estimated < openAIInputTokensFallbackMinimum { estimated = openAIInputTokensFallbackMinimum @@ -86,6 +92,31 @@ func (s *OpenAIGatewayService) ForwardCountTokensAsAnthropic( return fmt.Errorf("count_tokens: missing account") } + // 国产供应商 Anthropic 协议:上游有原生 /v1/messages/count_tokens 端点, + // 直接透传(仅模型名映射),不走 /v1/responses/input_tokens 估算。 + if account.IsAnthropicProtocol() { + return s.forwardCountTokensViaNativeAnthropic(ctx, c, account, body, defaultMappedModel) + } + + // 国产供应商其余协议(chat_completions / responses):三家上游均无 + // OpenAI 兼容的 /v1/responses/input_tokens 端点,与 Grok 一样本地估算, + // 不发上游请求(Claude Code 客户端会高频调用 count_tokens)。 + if account.IsCNProvider() { + estimated, err := estimateAnthropicCountTokensLocally(body) + if err != nil { + writeAnthropicCountTokensError(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") + return fmt.Errorf("count_tokens: estimate cn provider input tokens: %w", err) + } + logger.L().Debug("openai count_tokens: cn provider local estimate", + zap.Int64("account_id", account.ID), + zap.Int("estimated_input_tokens", estimated), + ) + c.JSON(http.StatusOK, gin.H{ + "input_tokens": estimated, + }) + return nil + } + prepared, err := prepareOpenAIInputTokensCountRequest(body, account, defaultMappedModel) if err != nil { writeAnthropicCountTokensError(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") diff --git a/backend/internal/service/openai_gateway_forward.go b/backend/internal/service/openai_gateway_forward.go index a3ba6ebab9..a262fb60cc 100644 --- a/backend/internal/service/openai_gateway_forward.go +++ b/backend/internal/service/openai_gateway_forward.go @@ -108,6 +108,13 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco return s.forwardGrokResponses(ctx, c, account, body, originalModel, reqStream, startTime) } + // CN 供应商 anthropic 协议账号:/v1/responses 入站是交叉协议组合 + // (Responses 客户端 × Anthropic 上游),转成 Anthropic 请求走原生端点。 + // 不能落到下面的 raw-CC 分支——其 URL 构造会把 anthropic base 当 CC base 用。 + if account.IsAnthropicProtocol() { + return s.forwardResponsesViaNativeAnthropic(ctx, c, account, body, reqModel) + } + if shouldForwardOpenAIResponsesViaRawChatCompletions(account) { return s.forwardResponsesViaRawChatCompletions(ctx, c, account, body) } @@ -1051,13 +1058,17 @@ func (s *OpenAIGatewayService) buildUpstreamRequest(ctx context.Context, c *gin. if err != nil { return nil, err } - targetURL = buildOpenAIResponsesURL(validatedURL) + targetURL = buildOpenAIResponsesURLForPlatform(account.Platform, validatedURL) } default: targetURL = openaiPlatformAPIURL } targetURL = appendOpenAIResponsesRequestPathSuffix(targetURL, openAIResponsesRequestPathSuffix(c)) + // DeepSeek 原生 Responses 端点为无状态实现:强制 store=false、清除 + // previous_response_id,避免携带状态字段被上游拒绝。 + body = normalizeDeepSeekResponsesRequestBody(account, body) + req, err := http.NewRequestWithContext(ctx, "POST", targetURL, bytes.NewReader(body)) if err != nil { return nil, err diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index 268aec7b5a..43f4977f2d 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -36,6 +36,15 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( ) (*OpenAIForwardResult, error) { beginUpstreamResponseModelObservation(c) + // 入口分流(国产供应商 Anthropic 协议):上游为供应商原生 Anthropic 端点时, + // /v1/messages 请求零转换直通(仅模型名映射 + 少量 body 清洗),完整保留 + // thinking / tool_use / cache 语义,适配 Claude Code 等原生客户端。 + // 必须先于 ShouldUseResponsesAPI 分流:Anthropic 协议账号经 probe 落标 + // openai_responses_supported=false,会先命中下方的 CC 直转分支。 + if account.IsAnthropicProtocol() { + return s.forwardAnthropicViaNativeAnthropicEndpoint(ctx, c, account, body, defaultMappedModel) + } + // 入口分流:APIKey 账号 + 上游不支持 Responses API → 走 CC 直转(与 // ForwardAsChatCompletions 对称)。缺少此分流时,/v1/messages 入站请求 // 会被无条件转为 Responses 格式发往上游 /v1/responses,导致只支持 diff --git a/backend/internal/service/openai_gateway_messages_anthropic_native.go b/backend/internal/service/openai_gateway_messages_anthropic_native.go new file mode 100644 index 0000000000..e785f3755f --- /dev/null +++ b/backend/internal/service/openai_gateway_messages_anthropic_native.go @@ -0,0 +1,619 @@ +package service + +// 国产供应商(kimi/zhipu/deepseek)原生 Anthropic 端点直通路径。 +// +// 当账号 credentials["api_protocol"] = "anthropic" 时,入站 /v1/messages 请求 +// 不再做 Anthropic→CC→Anthropic 双重转换,而是零转换直通供应商的官方 +// Anthropic 兼容端点(如 https://open.bigmodel.cn/api/anthropic/v1/messages), +// 适配 Claude Code 等原生 Anthropic 客户端。转发骨架以 +// gateway_anthropic_passthrough.go 的 APIKey 透传为模板(字节级 SSE 中继 + +// usage 解析),错误/failover 语义对齐 OpenAI 网关其他路径 +// (failoverOpenAIUpstreamHTTPError / handleAnthropicErrorResponse)。 + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strings" + "sync/atomic" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/gin-gonic/gin" + "github.com/tidwall/gjson" + "github.com/tidwall/sjson" +) + +// forwardAnthropicViaNativeAnthropicEndpoint 将 Anthropic Messages 请求零转换 +// 直通到国产供应商的原生 Anthropic 端点。仅做模型名映射与少量 body 清洗 +// (空文本块 / web-search 历史块),协议本身不转换。 +func (s *OpenAIGatewayService) forwardAnthropicViaNativeAnthropicEndpoint( + ctx context.Context, + c *gin.Context, + account *Account, + body []byte, + defaultMappedModel string, +) (*OpenAIForwardResult, error) { + startTime := time.Now() + + originalModel := strings.TrimSpace(gjson.GetBytes(body, "model").String()) + if originalModel == "" { + writeAnthropicError(c, http.StatusBadRequest, "invalid_request_error", "model is required") + return nil, fmt.Errorf("missing model in request") + } + clientStream := gjson.GetBytes(body, "stream").Bool() + + billingModel := resolveOpenAIForwardModel(account, originalModel, defaultMappedModel) + upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel) + if upstreamModel != originalModel { + rewritten, err := sjson.SetBytes(body, "model", upstreamModel) + if err != nil { + return nil, fmt.Errorf("rewrite model: %w", err) + } + body = rewritten + } + + // 与 Anthropic 平台 passthrough 相同的 pre-filter:剥离空文本块与上游 + // 无法接受的 web-search 历史块(GLM/Kimi/DeepSeek 对 server_tool_use 400)。 + body = StripEmptyTextBlocks(body) + body = FilterWebSearchHistoryBlocks(body, upstreamModel) + + logger.LegacyPrintf("service.gateway", "[CN Anthropic 直通] account=%d(%s) platform=%s model=%s upstream=%s stream=%v", + account.ID, account.Name, account.Platform, originalModel, upstreamModel, clientStream) + + apiKey := strings.TrimSpace(account.GetOpenAIProtocolAPIKey()) + if apiKey == "" { + return nil, fmt.Errorf("account %d missing api_key", account.ID) + } + targetURL, err := s.nativeAnthropicTargetURL(account) + if err != nil { + return nil, err + } + + proxyURL := "" + if account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + + upstreamCtx, releaseUpstreamCtx := detachStreamUpstreamContext(ctx, clientStream) + upstreamReq, _, err := s.buildNativeAnthropicUpstreamRequest(upstreamCtx, c, account, body, apiKey, targetURL) + releaseUpstreamCtx() + if err != nil { + return nil, err + } + + resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) + if err != nil { + return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, true) + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode >= 400 { + respBody, upstreamMsg := s.readOpenAIUpstreamError(resp) + if foErr := s.failoverOpenAIUpstreamHTTPError(ctx, c, account, resp, respBody, upstreamMsg, upstreamModel); foErr != nil { + return nil, foErr + } + // 非 failover 错误:经共享 compat handler 以 Anthropic 格式回写 + // (透传规则、ops 记录、cyber_policy 与 CC 回退路径一致)。 + return s.handleAnthropicErrorResponse(resp, c, account, billingModel) + } + + if clientStream { + return s.handleNativeAnthropicStreamingResponse(ctx, resp, c, account, originalModel, billingModel, upstreamModel, startTime) + } + return s.handleNativeAnthropicBufferedResponse(ctx, resp, c, account, originalModel, billingModel, upstreamModel, startTime) +} + +// nativeAnthropicTargetURL 组装国产供应商原生 Anthropic messages 端点。 +// 第三方端点保持朴素路径,不附加 ?beta=true。 +func (s *OpenAIGatewayService) nativeAnthropicTargetURL(account *Account) (string, error) { + baseURL := strings.TrimSpace(account.GetAnthropicProtocolBaseURL()) + if baseURL == "" { + return "", fmt.Errorf("account %d has no anthropic protocol base url", account.ID) + } + validatedURL, err := s.validateUpstreamBaseURL(baseURL) + if err != nil { + return "", fmt.Errorf("invalid base_url: %w", err) + } + return strings.TrimRight(validatedURL, "/") + "/v1/messages", nil +} + +func (s *OpenAIGatewayService) buildNativeAnthropicUpstreamRequest( + ctx context.Context, + c *gin.Context, + account *Account, + body []byte, + apiKey string, + targetURL string, +) (*http.Request, []byte, error) { + // 能力维度 body sanitize:与 Anthropic 平台 passthrough 相同,按 beta + // header 决定是否保留 body 中的 beta 能力字段,避免客户端"body 带字段但 + // header 忘带 token"的 bug 让第三方上游 400。 + clientBeta := "" + if c != nil && c.Request != nil { + clientBeta = getHeaderRaw(c.Request.Header, "anthropic-beta") + } + if beta, ok := account.HeaderOverrideValue("anthropic-beta"); ok { + clientBeta = beta + } + if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta); changed { + body = sanitized + } + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) + if err != nil { + return nil, nil, err + } + + if c != nil && c.Request != nil { + for key, values := range c.Request.Header { + lowerKey := strings.ToLower(strings.TrimSpace(key)) + if !allowedHeaders[lowerKey] { + continue + } + wireKey := resolveWireCasing(key) + for _, v := range values { + addHeaderRaw(req.Header, wireKey, v) + } + } + } + + // 覆盖入站鉴权残留,注入上游认证(默认 x-api-key;可经 extra + // anthropic_apikey_auth_scheme 切换 Authorization: Bearer)。 + req.Header.Del("authorization") + req.Header.Del("x-api-key") + req.Header.Del("x-goog-api-key") + req.Header.Del("cookie") + setAnthropicAPIKeyAuthHeader(req.Header, account, apiKey) + + if getHeaderRaw(req.Header, "content-type") == "" { + setHeaderRaw(req.Header, "content-type", "application/json") + } + if getHeaderRaw(req.Header, "anthropic-version") == "" { + setHeaderRaw(req.Header, "anthropic-version", "2023-06-01") + } + + // 账号级请求头覆写(最终生效,覆盖上面所有来源的同名头) + account.ApplyHeaderOverrides(req.Header) + + return req, body, nil +} + +// handleNativeAnthropicBufferedResponse 处理非流式原生 Anthropic 响应: +// 校验 JSON、解析 usage、透传响应头后原样回写(仅工具名反向还原)。 +func (s *OpenAIGatewayService) handleNativeAnthropicBufferedResponse( + ctx context.Context, + resp *http.Response, + c *gin.Context, + account *Account, + originalModel string, + billingModel string, + upstreamModel string, + startTime time.Time, +) (*OpenAIForwardResult, error) { + if s.rateLimitService != nil { + s.rateLimitService.UpdateSessionWindow(ctx, account, resp.Header) + } + + body, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, anthropicTooLargeError) + if err != nil { + return nil, err + } + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } + observer.ObserveAnthropic(body) + + var raw json.RawMessage + if err := json.Unmarshal(body, &raw); err != nil { + return nil, invalidNonStreamingJSONFailoverError(ctx, s.rateLimitService, resp, account, body, err, billingModel) + } + + usage := parseClaudeUsageFromResponseBody(body) + if IsForceCacheBilling(ctx) && usage.InputTokens > 0 { + body, err = classifyAnthropicResponseInputAsCacheRead(body, usage) + if err != nil { + return nil, err + } + } + + writeAnthropicPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) + contentType := strings.TrimSpace(resp.Header.Get("Content-Type")) + if contentType == "" { + contentType = "application/json" + } + body = reverseToolNamesIfPresent(c, body) + c.Data(resp.StatusCode, contentType, body) + + return &OpenAIForwardResult{ + RequestID: resp.Header.Get("x-request-id"), + Usage: claudeUsageToOpenAIUsage(usage), + Model: originalModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, + UpstreamEndpoint: "/v1/messages", + Stream: false, + Duration: time.Since(startTime), + }, nil +} + +// handleNativeAnthropicStreamingResponse 处理流式原生 Anthropic 响应: +// 字节级 SSE 中继(逐行透传、按事件边界 flush),同时解析 usage。 +// 骨架与 handleStreamingResponseAnthropicAPIKeyPassthrough 一致。 +func (s *OpenAIGatewayService) handleNativeAnthropicStreamingResponse( + ctx context.Context, + resp *http.Response, + c *gin.Context, + account *Account, + originalModel string, + billingModel string, + upstreamModel string, + startTime time.Time, +) (*OpenAIForwardResult, error) { + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } + if s.rateLimitService != nil { + s.rateLimitService.UpdateSessionWindow(ctx, account, resp.Header) + } + + writeAnthropicPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) + + contentType := strings.TrimSpace(resp.Header.Get("Content-Type")) + if contentType == "" { + contentType = "text/event-stream" + } + c.Header("Content-Type", contentType) + if c.Writer.Header().Get("Cache-Control") == "" { + c.Header("Cache-Control", "no-cache") + } + if c.Writer.Header().Get("Connection") == "" { + c.Header("Connection", "keep-alive") + } + c.Header("X-Accel-Buffering", "no") + if v := resp.Header.Get("x-request-id"); v != "" { + c.Header("x-request-id", v) + } + + w := c.Writer + flusher, ok := w.(http.Flusher) + if !ok { + return nil, errors.New("streaming not supported") + } + + usage := &ClaudeUsage{} + var firstTokenMs *int + clientDisconnected := false + sawTerminalEvent := false + + scanner := bufio.NewScanner(resp.Body) + maxLineSize := defaultMaxLineSize + if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { + maxLineSize = s.cfg.Gateway.MaxLineSize + } + scanBuf := getSSEScannerBuf64K() + scanner.Buffer(scanBuf[:0], maxLineSize) + + type scanEvent struct { + line string + err error + } + events := make(chan scanEvent, 16) + done := make(chan struct{}) + sendEvent := func(ev scanEvent) bool { + select { + case events <- ev: + return true + case <-done: + return false + } + } + var lastReadAt int64 + atomic.StoreInt64(&lastReadAt, time.Now().UnixNano()) + go func(scanBuf *sseScannerBuf64K) { + defer putSSEScannerBuf64K(scanBuf) + defer close(events) + for scanner.Scan() { + atomic.StoreInt64(&lastReadAt, time.Now().UnixNano()) + if !sendEvent(scanEvent{line: scanner.Text()}) { + return + } + } + if err := scanner.Err(); err != nil { + _ = sendEvent(scanEvent{err: err}) + } + }(scanBuf) + defer close(done) + + streamInterval := time.Duration(0) + if s.cfg != nil && s.cfg.Gateway.StreamDataIntervalTimeout > 0 { + streamInterval = time.Duration(s.cfg.Gateway.StreamDataIntervalTimeout) * time.Second + } + var intervalTicker *time.Ticker + if streamInterval > 0 { + intervalTicker = time.NewTicker(streamInterval) + defer intervalTicker.Stop() + } + var intervalCh <-chan time.Time + if intervalTicker != nil { + intervalCh = intervalTicker.C + } + + keepaliveInterval := time.Duration(0) + if s.cfg != nil && s.cfg.Gateway.StreamKeepaliveInterval > 0 { + keepaliveInterval = time.Duration(s.cfg.Gateway.StreamKeepaliveInterval) * time.Second + } + var keepaliveTimer *time.Timer + if keepaliveInterval > 0 { + keepaliveTimer = time.NewTimer(keepaliveInterval) + defer keepaliveTimer.Stop() + } + var keepaliveCh <-chan time.Time + if keepaliveTimer != nil { + keepaliveCh = keepaliveTimer.C + } + lastDataAt := time.Now() + resetKeepaliveTimer := func() { + if keepaliveTimer == nil { + return + } + if !keepaliveTimer.Stop() { + select { + case <-keepaliveTimer.C: + default: + } + } + keepaliveTimer.Reset(keepaliveInterval) + } + inPartialEvent := false + + for { + select { + case ev, ok := <-events: + if !ok { + if !clientDisconnected { + flusher.Flush() + } + if !sawTerminalEvent { + return s.nativeAnthropicStreamResult(c, resp, usage, firstTokenMs, clientDisconnected, originalModel, billingModel, upstreamModel, startTime), + fmt.Errorf("stream usage incomplete: missing terminal event") + } + return s.nativeAnthropicStreamResult(c, resp, usage, firstTokenMs, clientDisconnected, originalModel, billingModel, upstreamModel, startTime), nil + } + if ev.err != nil { + if sawTerminalEvent { + return s.nativeAnthropicStreamResult(c, resp, usage, firstTokenMs, clientDisconnected, originalModel, billingModel, upstreamModel, startTime), nil + } + if clientDisconnected { + return s.nativeAnthropicStreamResult(c, resp, usage, firstTokenMs, clientDisconnected, originalModel, billingModel, upstreamModel, startTime), + fmt.Errorf("stream usage incomplete after disconnect: %w", ev.err) + } + if errors.Is(ev.err, context.Canceled) || errors.Is(ev.err, context.DeadlineExceeded) { + return s.nativeAnthropicStreamResult(c, resp, usage, firstTokenMs, clientDisconnected, originalModel, billingModel, upstreamModel, startTime), + fmt.Errorf("stream usage incomplete: %w", ev.err) + } + if errors.Is(ev.err, bufio.ErrTooLong) { + logger.LegacyPrintf("service.gateway", "[CN Anthropic 直通] SSE line too long: account=%d max_size=%d error=%v", account.ID, maxLineSize, ev.err) + return s.nativeAnthropicStreamResult(c, resp, usage, firstTokenMs, clientDisconnected, originalModel, billingModel, upstreamModel, startTime), ev.err + } + return s.nativeAnthropicStreamResult(c, resp, usage, firstTokenMs, clientDisconnected, originalModel, billingModel, upstreamModel, startTime), + fmt.Errorf("stream read error: %w", ev.err) + } + + line := ev.line + if data, ok := extractAnthropicSSEDataLine(line); ok { + trimmed := strings.TrimSpace(data) + observer.ObserveAnthropic([]byte(trimmed)) + if anthropicStreamEventIsTerminal("", trimmed) { + sawTerminalEvent = true + } + if firstTokenMs == nil && trimmed != "" && trimmed != "[DONE]" { + ms := int(time.Since(startTime).Milliseconds()) + firstTokenMs = &ms + } + parseSSEUsagePassthrough(data, usage) + } else { + trimmed := strings.TrimSpace(line) + if strings.HasPrefix(trimmed, "event:") && anthropicStreamEventIsTerminal(strings.TrimSpace(strings.TrimPrefix(trimmed, "event:")), "") { + sawTerminalEvent = true + } + } + + if !clientDisconnected { + restored := string(reverseToolNamesIfPresent(c, []byte(line))) + if _, err := io.WriteString(w, restored); err != nil { + clientDisconnected = true + logger.LegacyPrintf("service.gateway", "[CN Anthropic 直通] Client disconnected during streaming, continue draining upstream for usage: account=%d", account.ID) + } else if _, err := io.WriteString(w, "\n"); err != nil { + clientDisconnected = true + logger.LegacyPrintf("service.gateway", "[CN Anthropic 直通] Client disconnected during streaming, continue draining upstream for usage: account=%d", account.ID) + } else if line == "" { + // 按 SSE 事件边界刷出,减少每行 flush 带来的 syscall 开销。 + flusher.Flush() + lastDataAt = time.Now() + resetKeepaliveTimer() + inPartialEvent = false + } else { + inPartialEvent = true + } + } + + case <-intervalCh: + lastRead := time.Unix(0, atomic.LoadInt64(&lastReadAt)) + if time.Since(lastRead) < streamInterval { + continue + } + if clientDisconnected { + return s.nativeAnthropicStreamResult(c, resp, usage, firstTokenMs, clientDisconnected, originalModel, billingModel, upstreamModel, startTime), + fmt.Errorf("stream usage incomplete after timeout") + } + logger.LegacyPrintf("service.gateway", "[CN Anthropic 直通] Stream data interval timeout: account=%d model=%s interval=%s", account.ID, upstreamModel, streamInterval) + if s.rateLimitService != nil { + s.rateLimitService.HandleStreamTimeout(ctx, account, upstreamModel) + } + return s.nativeAnthropicStreamResult(c, resp, usage, firstTokenMs, clientDisconnected, originalModel, billingModel, upstreamModel, startTime), + fmt.Errorf("stream data interval timeout") + + case <-keepaliveCh: + if clientDisconnected { + continue + } + if inPartialEvent { + resetKeepaliveTimer() + continue + } + if time.Since(lastDataAt) < keepaliveInterval { + resetKeepaliveTimer() + continue + } + if _, err := fmt.Fprint(w, "event: ping\ndata: {\"type\": \"ping\"}\n\n"); err != nil { + clientDisconnected = true + logger.LegacyPrintf("service.gateway", "[CN Anthropic 直通] Client disconnected during keepalive ping, continue draining upstream for usage: account=%d", account.ID) + continue + } + flusher.Flush() + lastDataAt = time.Now() + resetKeepaliveTimer() + } + } +} + +// nativeAnthropicStreamResult 组装流式直通结果;流中断时同样返回已观测到的 +// usage 与错误一起带出,避免上游已计量的请求漏记漏计费(对齐 issue #5148 语义)。 +func (s *OpenAIGatewayService) nativeAnthropicStreamResult( + c *gin.Context, + resp *http.Response, + usage *ClaudeUsage, + firstTokenMs *int, + clientDisconnect bool, + originalModel string, + billingModel string, + upstreamModel string, + startTime time.Time, +) *OpenAIForwardResult { + if usage == nil { + usage = &ClaudeUsage{} + } + return &OpenAIForwardResult{ + RequestID: resp.Header.Get("x-request-id"), + Usage: claudeUsageToOpenAIUsage(usage), + Model: originalModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, + UpstreamEndpoint: "/v1/messages", + Stream: true, + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, + ClientDisconnect: clientDisconnect, + } +} + +// forwardCountTokensViaNativeAnthropic 把 Anthropic count_tokens 请求透传到 +// 国产供应商原生 Anthropic 端点({base}/v1/messages/count_tokens), +// 仅做模型名映射,不做协议转换。 +func (s *OpenAIGatewayService) forwardCountTokensViaNativeAnthropic( + ctx context.Context, + c *gin.Context, + account *Account, + body []byte, + defaultMappedModel string, +) error { + originalModel := strings.TrimSpace(gjson.GetBytes(body, "model").String()) + if originalModel == "" { + writeAnthropicCountTokensError(c, http.StatusBadRequest, "invalid_request_error", "model is required") + return fmt.Errorf("count_tokens: missing model in request") + } + billingModel := resolveOpenAIForwardModel(account, originalModel, strings.TrimSpace(defaultMappedModel)) + upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel) + if upstreamModel != originalModel { + rewritten, err := sjson.SetBytes(body, "model", upstreamModel) + if err != nil { + return fmt.Errorf("count_tokens: rewrite model: %w", err) + } + body = rewritten + } + + apiKey := strings.TrimSpace(account.GetOpenAIProtocolAPIKey()) + if apiKey == "" { + writeAnthropicCountTokensError(c, http.StatusBadGateway, "upstream_error", "Account api_key is missing") + return fmt.Errorf("count_tokens: account %d missing api_key", account.ID) + } + targetURL, err := s.nativeAnthropicTargetURL(account) + if err != nil { + return fmt.Errorf("count_tokens: %w", err) + } + targetURL = strings.TrimSuffix(targetURL, "/v1/messages") + "/v1/messages/count_tokens" + + upstreamReq, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) + if err != nil { + writeAnthropicCountTokensError(c, http.StatusInternalServerError, "api_error", "Failed to build request") + return fmt.Errorf("count_tokens: build request: %w", err) + } + reqHeader := upstreamReq.Header + reqHeader.Del("authorization") + reqHeader.Del("x-api-key") + setAnthropicAPIKeyAuthHeader(reqHeader, account, apiKey) + reqHeader.Set("content-type", "application/json") + reqHeader.Set("accept", "application/json") + account.ApplyHeaderOverrides(reqHeader) + + proxyURL := "" + if account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) + if err != nil { + safeErr := sanitizeUpstreamErrorMessage(err.Error()) + setOpsUpstreamError(c, 0, safeErr, "") + writeAnthropicCountTokensError(c, http.StatusBadGateway, "upstream_error", "Upstream request failed") + return fmt.Errorf("count_tokens: upstream request failed: %s", safeErr) + } + defer func() { _ = resp.Body.Close() }() + + // count_tokens 响应体极小;与其他探测路径一致加 256KB 上限防异常上游放大内存。 + respBody, err := io.ReadAll(io.LimitReader(resp.Body, cnQuotaMaxBodyBytes)) + if err != nil { + writeAnthropicCountTokensError(c, http.StatusBadGateway, "upstream_error", "Failed to read response") + return fmt.Errorf("count_tokens: read response: %w", err) + } + if resp.StatusCode >= 400 { + if s.rateLimitService != nil { + s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody) + } + upstreamMsg := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(respBody))) + setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, "") + writeAnthropicCountTokensError(c, resp.StatusCode, "upstream_error", "Upstream request failed") + return fmt.Errorf("count_tokens: upstream error: %d", resp.StatusCode) + } + + inputTokens := gjson.GetBytes(respBody, "input_tokens") + if !inputTokens.Exists() { + writeAnthropicCountTokensError(c, http.StatusBadGateway, "upstream_error", "Upstream response missing input_tokens") + return fmt.Errorf("count_tokens: response missing input_tokens field") + } + c.JSON(http.StatusOK, gin.H{ + "input_tokens": int(inputTokens.Int()), + }) + return nil +} + +// claudeUsageToOpenAIUsage 把 Anthropic 格式 usage 映射到 OpenAI 网关统一的 +// 用量结构(字段一一对应)。 +func claudeUsageToOpenAIUsage(u *ClaudeUsage) OpenAIUsage { + if u == nil { + return OpenAIUsage{} + } + return OpenAIUsage{ + InputTokens: u.InputTokens, + OutputTokens: u.OutputTokens, + CacheCreationInputTokens: u.CacheCreationInputTokens, + CacheReadInputTokens: u.CacheReadInputTokens, + } +} diff --git a/backend/internal/service/openai_gateway_model_availability.go b/backend/internal/service/openai_gateway_model_availability.go index a665052e60..35adea0d72 100644 --- a/backend/internal/service/openai_gateway_model_availability.go +++ b/backend/internal/service/openai_gateway_model_availability.go @@ -34,7 +34,7 @@ func (s *OpenAIGatewayService) DiagnoseModelAvailabilityForPlatform( return ModelAvailabilityDiagnosis{HasAccountsInPool: true, HasModelSupport: true} } - platform = normalizeOpenAICompatiblePlatform(platform) + platform = NormalizeOpenAICompatiblePlatform(platform) queryGroupID := groupID includeGrouped := false if s.cfg != nil && s.cfg.RunMode == config.RunModeSimple { diff --git a/backend/internal/service/openai_gateway_passthrough.go b/backend/internal/service/openai_gateway_passthrough.go index 1fac0440f6..291f8bcc9e 100644 --- a/backend/internal/service/openai_gateway_passthrough.go +++ b/backend/internal/service/openai_gateway_passthrough.go @@ -380,11 +380,14 @@ func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough( if err != nil { return nil, err } - targetURL = buildOpenAIResponsesURL(validatedURL) + targetURL = buildOpenAIResponsesURLForPlatform(account.Platform, validatedURL) } } targetURL = appendOpenAIResponsesRequestPathSuffix(targetURL, openAIResponsesRequestPathSuffix(c)) + // DeepSeek 原生 Responses 端点为无状态实现(见 normalizeDeepSeekResponsesRequestBody)。 + body = normalizeDeepSeekResponsesRequestBody(account, body) + req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) if err != nil { return nil, err diff --git a/backend/internal/service/openai_gateway_request_body.go b/backend/internal/service/openai_gateway_request_body.go index 8be2d0cf2b..3559fb400d 100644 --- a/backend/internal/service/openai_gateway_request_body.go +++ b/backend/internal/service/openai_gateway_request_body.go @@ -45,6 +45,33 @@ func buildOpenAIResponsesURL(base string) string { return buildOpenAIEndpointURL(base, "/v1/responses") } +// buildOpenAIResponsesURLForPlatform 组装 Responses 端点(平台感知)。 +// DeepSeek 官方 Responses 端点为 /responses(无 /v1 前缀,适配 Codex); +// 其余平台维持 /v1/responses。 +func buildOpenAIResponsesURLForPlatform(platform string, base string) string { + if platform == PlatformDeepseek { + return buildOpenAIEndpointURL(base, "/responses") + } + return buildOpenAIResponsesURL(base) +} + +// normalizeDeepSeekResponsesRequestBody 适配 DeepSeek 无状态 Responses 端点: +// 强制 store=false 并清除 previous_response_id(官方 /responses 不支持服务端 +// 状态存储,携带这些字段会被拒绝)。非 deepseek responses 协议账号原样返回。 +func normalizeDeepSeekResponsesRequestBody(account *Account, body []byte) []byte { + if account == nil || account.Platform != PlatformDeepseek || account.GetAPIProtocol() != APIProtocolResponses { + return body + } + normalized, err := sjson.SetBytes(body, "store", false) + if err != nil { + return body + } + if stripped, err := sjson.DeleteBytes(normalized, "previous_response_id"); err == nil { + normalized = stripped + } + return normalized +} + func trimOpenAIEncryptedReasoningItems(reqBody map[string]any) bool { if len(reqBody) == 0 { return false diff --git a/backend/internal/service/openai_gateway_responses_anthropic_native.go b/backend/internal/service/openai_gateway_responses_anthropic_native.go new file mode 100644 index 0000000000..b72067fc4d --- /dev/null +++ b/backend/internal/service/openai_gateway_responses_anthropic_native.go @@ -0,0 +1,468 @@ +package service + +// 国产供应商 Anthropic 协议账号的 Responses 入站反向路径。 +// +// 客户端说 OpenAI Responses(/v1/responses,Codex 等)、上游是供应商原生 +// Anthropic 端点(api_protocol=anthropic)时的交叉组合:请求 Responses→Anthropic +// 单次转换,响应 Anthropic 事件→Responses 事件转换。转换链与 Anthropic 平台的 +// gateway_forward_as_responses.go 完全一致(复用同一组 apicompat 状态机),仅上游 +// 发送/错误处理对齐 OpenAI 网关语义(模型映射、failover、transport error)。 + +import ( + "bufio" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" + "github.com/gin-gonic/gin" + "github.com/tidwall/gjson" + "go.uber.org/zap" +) + +// forwardResponsesViaNativeAnthropic serves OpenAI /v1/responses clients through +// a CN provider's native Anthropic endpoint. +// +// Conversion chain: +// +// Request: Responses → Anthropic (single conversion) +// Response: Anthropic events → Responses events (stream state machine) +func (s *OpenAIGatewayService) forwardResponsesViaNativeAnthropic( + ctx context.Context, + c *gin.Context, + account *Account, + body []byte, + defaultMappedModel string, +) (*OpenAIForwardResult, error) { + startTime := time.Now() + + // 1. Lower Codex client-side tools to function tools understood by Anthropic. + adaptedBody, clientToolMapping, err := adaptResponsesClientToolsForAnthropic(body) + if err != nil { + writeResponsesError(c, http.StatusBadRequest, "invalid_request_error", "Failed to adapt request tools") + return nil, fmt.Errorf("adapt responses client tools: %w", err) + } + + // 2. Parse Responses request + var responsesReq apicompat.ResponsesRequest + if err := json.Unmarshal(adaptedBody, &responsesReq); err != nil { + writeResponsesError(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") + return nil, fmt.Errorf("parse responses request: %w", err) + } + originalModel := responsesReq.Model + if strings.TrimSpace(originalModel) == "" { + writeResponsesError(c, http.StatusBadRequest, "invalid_request_error", "model is required") + return nil, fmt.Errorf("missing model in request") + } + clientStream := responsesReq.Stream + + // 3. Convert Responses → Anthropic + anthropicReq, err := apicompat.ResponsesToAnthropicRequest(&responsesReq) + if err != nil { + writeResponsesError(c, http.StatusBadRequest, "invalid_request_error", "Failed to convert request") + return nil, fmt.Errorf("convert responses to anthropic: %w", err) + } + + // 4. Model mapping(OpenAI 网关统一入口的映射语义) + billingModel := resolveOpenAIForwardModel(account, originalModel, defaultMappedModel) + upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel) + anthropicReq.Model = upstreamModel + + reasoningEffort := ExtractResponsesReasoningEffortFromBody(body) + reasoningEffort = ApplyThinkingEnabledFallback(reasoningEffort, body, billingModel) + + // 5. Force upstream streaming(客户端原始终决定响应格式; + // 上游恒为流式,非流式由缓冲路径组装)。 + anthropicReq.Stream = true + reqStream := true + + logger.L().Debug("openai responses: forwarding via native anthropic endpoint", + zap.Int64("account_id", account.ID), + zap.String("original_model", originalModel), + zap.String("billing_model", billingModel), + zap.String("upstream_model", upstreamModel), + zap.Bool("client_stream", clientStream), + ) + + anthropicBody, err := json.Marshal(anthropicReq) + if err != nil { + return nil, fmt.Errorf("marshal anthropic request: %w", err) + } + + // 与 /v1/messages 直通路径相同的 pre-filter。 + anthropicBody = StripEmptyTextBlocks(anthropicBody) + anthropicBody = FilterWebSearchHistoryBlocks(anthropicBody, upstreamModel) + anthropicBody = enforceCacheControlLimit(anthropicBody) + + apiKey := strings.TrimSpace(account.GetOpenAIProtocolAPIKey()) + if apiKey == "" { + return nil, fmt.Errorf("account %d missing api_key", account.ID) + } + targetURL, err := s.nativeAnthropicTargetURL(account) + if err != nil { + return nil, err + } + + proxyURL := "" + if account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + + upstreamCtx, releaseUpstreamCtx := detachStreamUpstreamContext(ctx, reqStream) + upstreamReq, _, err := s.buildNativeAnthropicUpstreamRequest(upstreamCtx, c, account, anthropicBody, apiKey, targetURL) + releaseUpstreamCtx() + if err != nil { + return nil, fmt.Errorf("build upstream request: %w", err) + } + + resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) + if err != nil { + return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, true) + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode >= 400 { + respBody, upstreamMsg := s.readOpenAIUpstreamError(resp) + if foErr := s.failoverOpenAIUpstreamHTTPError(ctx, c, account, resp, respBody, upstreamMsg, upstreamModel); foErr != nil { + return nil, foErr + } + writeResponsesError(c, mapUpstreamStatusCode(resp.StatusCode), "server_error", upstreamMsg) + return nil, fmt.Errorf("upstream error: %d %s", resp.StatusCode, upstreamMsg) + } + + if clientStream { + return s.handleResponsesStreamingFromNativeAnthropic(resp, c, originalModel, billingModel, upstreamModel, reasoningEffort, startTime, clientToolMapping) + } + return s.handleResponsesBufferedFromNativeAnthropic(resp, c, originalModel, billingModel, upstreamModel, reasoningEffort, startTime, clientToolMapping) +} + +// handleResponsesBufferedFromNativeAnthropic reads Anthropic SSE events, assembles +// the full response, then converts Anthropic → Responses. +func (s *OpenAIGatewayService) handleResponsesBufferedFromNativeAnthropic( + resp *http.Response, + c *gin.Context, + originalModel string, + billingModel string, + upstreamModel string, + reasoningEffort *string, + startTime time.Time, + clientToolMapping apicompat.ResponsesClientToolMapping, +) (*OpenAIForwardResult, error) { + requestID := resp.Header.Get("x-request-id") + + scanner := bufio.NewScanner(resp.Body) + maxLineSize := defaultMaxLineSize + if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { + maxLineSize = s.cfg.Gateway.MaxLineSize + } + scanner.Buffer(make([]byte, 0, 64*1024), maxLineSize) + + var finalResp *apicompat.AnthropicResponse + var usage ClaudeUsage + + // 读间隔上限:上游挂住 SSE 时中止组装(缓冲路径尚未提交响应头,可回 502)。 + streamInterval := s.anthropicNativeStreamInterval() + pump := newAnthropicNativeLinePump(scanner, streamInterval) + defer pump.stop() + + logReadErr := func(err error) { + if !errors.Is(err, io.EOF) && !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) { + logger.L().Warn("openai responses via native anthropic buffered: read error", + zap.Error(err), + zap.String("request_id", requestID), + ) + } + } + onIdle := func() (*OpenAIForwardResult, error) { + _ = resp.Body.Close() + logger.L().Warn("openai responses via native anthropic buffered: data interval timeout", + zap.String("request_id", requestID), + zap.Duration("interval", streamInterval), + ) + writeResponsesError(c, http.StatusBadGateway, "server_error", "Upstream stream data interval timeout") + return nil, fmt.Errorf("stream data interval timeout") + } + + for { + line, rerr := pump.next() + if rerr != nil { + if errors.Is(rerr, errAnthropicNativeStreamIdle) { + return onIdle() + } + logReadErr(rerr) + break + } + // SSE 规范允许 `event:xxx`(冒号后无空格):Kimi 等上游返回紧凑格式。 + if _, ok := extractOpenAISSEEventLine(line); !ok { + continue + } + + dataLine, rerr := pump.next() + if rerr != nil { + if errors.Is(rerr, errAnthropicNativeStreamIdle) { + return onIdle() + } + logReadErr(rerr) + break + } + payload, ok := extractOpenAISSEDataLine(dataLine) + if !ok { + continue + } + + var event apicompat.AnthropicStreamEvent + if err := json.Unmarshal([]byte(payload), &event); err != nil { + continue + } + + if event.Type == "message_start" && event.Message != nil { + finalResp = event.Message + mergeAnthropicUsage(&usage, event.Message.Usage) + } + if event.Type == "message_delta" { + if event.Usage != nil { + mergeAnthropicUsage(&usage, *event.Usage) + } + if event.Delta != nil && event.Delta.StopReason != "" && finalResp != nil { + finalResp.StopReason = apicompat.AnthropicStopReasonPtr(event.Delta.StopReason) + } + } + if event.Type == "content_block_start" && event.ContentBlock != nil && finalResp != nil { + finalResp.Content = append(finalResp.Content, *event.ContentBlock) + } + if event.Type == "content_block_delta" && event.Delta != nil && finalResp != nil && event.Index != nil { + idx := *event.Index + if idx < len(finalResp.Content) { + switch event.Delta.Type { + case "text_delta": + finalResp.Content[idx].Text += event.Delta.Text + case "thinking_delta": + finalResp.Content[idx].Thinking += event.Delta.Thinking + case "input_json_delta": + finalResp.Content[idx].Input = appendRawJSON(finalResp.Content[idx].Input, event.Delta.PartialJSON) + } + } + } + } + + if finalResp == nil { + writeResponsesError(c, http.StatusBadGateway, "server_error", "Upstream stream ended without a response") + return nil, fmt.Errorf("upstream stream ended without response") + } + + if usage.InputTokens > 0 || usage.OutputTokens > 0 { + finalResp.Usage = apicompat.AnthropicUsage{ + InputTokens: usage.InputTokens, + OutputTokens: usage.OutputTokens, + CacheCreationInputTokens: usage.CacheCreationInputTokens, + CacheReadInputTokens: usage.CacheReadInputTokens, + } + } + + responsesResp := apicompat.AnthropicToResponsesResponse(finalResp) + responsesResp.Model = originalModel + + if s.responseHeaderFilter != nil { + responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) + } + // 非流式响应必须是 application/json(上游被强制流式,透传头会污染)。 + c.Writer.Header().Set("Content-Type", "application/json; charset=utf-8") + if respBytes, err := json.Marshal(responsesResp); err == nil { + respBytes = reverseToolNamesIfPresent(c, respBytes) + respBytes, _, err = apicompat.RestoreResponsesClientToolPayload(respBytes, clientToolMapping) + if err != nil { + return nil, fmt.Errorf("restore responses client tools: %w", err) + } + c.Data(http.StatusOK, "application/json; charset=utf-8", respBytes) + } else { + c.JSON(http.StatusOK, responsesResp) + } + + return &OpenAIForwardResult{ + RequestID: requestID, + Usage: claudeUsageToOpenAIUsage(&usage), + Model: originalModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, + UpstreamEndpoint: "/v1/messages", + ReasoningEffort: reasoningEffort, + Stream: false, + Duration: time.Since(startTime), + }, nil +} + +// handleResponsesStreamingFromNativeAnthropic reads Anthropic SSE events, converts +// each to Responses SSE events, and writes them to the client. +func (s *OpenAIGatewayService) handleResponsesStreamingFromNativeAnthropic( + resp *http.Response, + c *gin.Context, + originalModel string, + billingModel string, + upstreamModel string, + reasoningEffort *string, + startTime time.Time, + clientToolMapping apicompat.ResponsesClientToolMapping, +) (*OpenAIForwardResult, error) { + requestID := resp.Header.Get("x-request-id") + + if s.responseHeaderFilter != nil { + responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) + } + c.Writer.Header().Set("Content-Type", "text/event-stream") + c.Writer.Header().Set("Cache-Control", "no-cache") + c.Writer.Header().Set("Connection", "keep-alive") + c.Writer.Header().Set("X-Accel-Buffering", "no") + c.Writer.WriteHeader(http.StatusOK) + + state := apicompat.NewAnthropicEventToResponsesState() + state.Model = originalModel + clientToolRestorer := apicompat.NewResponsesClientToolStreamRestorer(clientToolMapping) + + var usage ClaudeUsage + var firstTokenMs *int + firstChunk := true + clientDisconnected := false + + scanner := bufio.NewScanner(resp.Body) + maxLineSize := defaultMaxLineSize + if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { + maxLineSize = s.cfg.Gateway.MaxLineSize + } + scanner.Buffer(make([]byte, 0, 64*1024), maxLineSize) + + resultWithUsage := func() *OpenAIForwardResult { + return &OpenAIForwardResult{ + RequestID: requestID, + Usage: claudeUsageToOpenAIUsage(&usage), + Model: originalModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, + UpstreamEndpoint: "/v1/messages", + ReasoningEffort: reasoningEffort, + Stream: true, + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, + ClientDisconnect: clientDisconnected, + } + } + + // 读间隔上限:上游挂住 SSE(不发数据也不断连)时结束转换循环。上游 ctx 为 + // WithoutCancel 且 http.Client 无整体 Timeout,无此界限则 scanner.Scan() + // 永久阻塞(见 anthropic native pump 文件注释)。 + streamInterval := s.anthropicNativeStreamInterval() + pump := newAnthropicNativeLinePump(scanner, streamInterval) + defer pump.stop() + + logReadErr := func(err error) { + if !errors.Is(err, io.EOF) && !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) { + logger.L().Warn("openai responses via native anthropic stream: read error", + zap.Error(err), + zap.String("request_id", requestID), + ) + } + } + onIdle := func() (*OpenAIForwardResult, error) { + _ = resp.Body.Close() + logger.L().Warn("openai responses via native anthropic stream: data interval timeout", + zap.String("request_id", requestID), + zap.Duration("interval", streamInterval), + ) + return resultWithUsage(), fmt.Errorf("stream data interval timeout") + } + + processAnthropicEvent := func(event *apicompat.AnthropicStreamEvent) bool { + if firstChunk { + firstChunk = false + ms := int(time.Since(startTime).Milliseconds()) + firstTokenMs = &ms + } + + if event.Type == "message_delta" && event.Usage != nil { + mergeAnthropicUsage(&usage, *event.Usage) + } + if event.Type == "message_start" && event.Message != nil { + mergeAnthropicUsage(&usage, event.Message.Usage) + } + + events := apicompat.AnthropicEventToResponsesEvents(event, state) + for _, evt := range events { + payload, err := json.Marshal(evt) + if err != nil { + continue + } + payload = reverseToolNamesIfPresent(c, payload) + payloads, _, err := clientToolRestorer.RestoreEvent(payload) + if err != nil { + continue + } + for _, restored := range payloads { + eventType := gjson.GetBytes(restored, "type").String() + if _, err := fmt.Fprintf(c.Writer, "event: %s\ndata: %s\n\n", eventType, restored); err != nil { + clientDisconnected = true + return true + } + } + } + if len(events) > 0 { + c.Writer.Flush() + } + return false + } + + for { + line, rerr := pump.next() + if rerr != nil { + if errors.Is(rerr, errAnthropicNativeStreamIdle) { + return onIdle() + } + logReadErr(rerr) + break + } + if _, ok := extractOpenAISSEEventLine(line); !ok { + continue + } + + dataLine, rerr := pump.next() + if rerr != nil { + if errors.Is(rerr, errAnthropicNativeStreamIdle) { + return onIdle() + } + logReadErr(rerr) + break + } + payload, ok := extractOpenAISSEDataLine(dataLine) + if !ok { + continue + } + + var event apicompat.AnthropicStreamEvent + if err := json.Unmarshal([]byte(payload), &event); err != nil { + continue + } + + if processAnthropicEvent(&event) { + return resultWithUsage(), nil + } + } + + // Finalize state machine(客户端已断开时仍执行,保证 usage 汇总完整)。 + if finalEvents := apicompat.FinalizeAnthropicResponsesStream(state); len(finalEvents) > 0 { + for _, evt := range finalEvents { + sse, err := apicompat.ResponsesEventToSSE(evt) + if err != nil { + continue + } + fmt.Fprint(c.Writer, sse) //nolint:errcheck + } + c.Writer.Flush() + } + + return resultWithUsage(), nil +} diff --git a/backend/internal/service/openai_gateway_scheduling.go b/backend/internal/service/openai_gateway_scheduling.go index f7dff0909f..7dd5414724 100644 --- a/backend/internal/service/openai_gateway_scheduling.go +++ b/backend/internal/service/openai_gateway_scheduling.go @@ -243,15 +243,22 @@ func (s *OpenAIGatewayService) SelectAccountForModelWithExclusions(ctx context.C return s.selectAccountForModelWithExclusions(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, PlatformOpenAI, sessionHash, requestedModel, excludedIDs, false, 0, "", false) } -// noAvailableOpenAISelectionError builds the standard "no account available" error -// while preserving the legacy /responses/compact error when applicable. -func normalizeOpenAICompatiblePlatform(platform string) string { - if platform == PlatformGrok { - return PlatformGrok +// NormalizeOpenAICompatiblePlatform 保留 grok 与国产 OpenAI 兼容供应商(kimi/zhipu/ +// deepseek)的原值,其他值一律归一为 openai。调度器据此对账号与请求做精确平台匹配: +// kimi 分组请求只命中 kimi 账号,语义与 openai/grok 一致。 +// (upstream 曾将本函数改为未导出 normalizeOpenAICompatiblePlatform,本分支的 +// handler 调度入口仍需导出,保持导出名。) +func NormalizeOpenAICompatiblePlatform(platform string) string { + switch platform { + case PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek: + return platform + default: + return PlatformOpenAI } - return PlatformOpenAI } +// noAvailableOpenAISelectionError builds the standard "no account available" error +// while preserving the legacy /responses/compact error when applicable. // details carries an optional machine-parseable exclusion summary (e.g. // "pool=2, filtered: quota_auto_pause_7d=1 runtime_blocked=1") appended in // parentheses. It is for server-side logs / ops diagnostics only: handlers @@ -327,7 +334,7 @@ func isOpenAICompatibleAccountEligibleForRequest(ctx context.Context, account *A // ordinary scheduling gate. Legacy selection uses it before classifying the // profit veto so earlier failures retain their actual reason. func isOpenAICompatibleAccountEligibleForRequestBeforeProfit(ctx context.Context, account *Account, platform string, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) bool { - platform = normalizeOpenAICompatiblePlatform(platform) + platform = NormalizeOpenAICompatiblePlatform(platform) if account == nil || account.Platform != platform || !account.IsOpenAICompatible() || !account.IsSchedulableForModelWithContext(ctx, requestedModel) { return false } @@ -726,7 +733,7 @@ func resolveOpenAIAccountUpstreamModelForRequest(account *Account, requestedMode } func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.Context, groupID *int64, platform string, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64, requiredCapability OpenAIEndpointCapability, preferLowUpstreamRate bool) (*Account, error) { - platform = normalizeOpenAICompatiblePlatform(platform) + platform = NormalizeOpenAICompatiblePlatform(platform) if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) { slog.Warn("channel pricing restriction blocked request", "group_id", derefGroupID(groupID), @@ -779,7 +786,7 @@ func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID if sessionHash == "" { return nil } - platform = normalizeOpenAICompatiblePlatform(platform) + platform = NormalizeOpenAICompatiblePlatform(platform) accountID := stickyAccountID if accountID <= 0 { @@ -846,7 +853,7 @@ func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID // true); the third contains deterministic // exclusion diagnostics for the evaluated snapshot. func (s *OpenAIGatewayService) selectBestAccount(ctx context.Context, groupID *int64, platform string, accounts []Account, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability, preferLowUpstreamRate bool) (*Account, bool, openAISelectionFilterStats) { - platform = normalizeOpenAICompatiblePlatform(platform) + platform = NormalizeOpenAICompatiblePlatform(platform) compactBlocked := false filterStats := openAISelectionFilterStats{pool: len(accounts)} needsUpstreamCheck := s.needsUpstreamChannelRestrictionCheck(ctx, groupID) @@ -958,7 +965,7 @@ func (s *OpenAIGatewayService) SelectAccountWithLoadAwareness(ctx context.Contex } func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Context, groupID *int64, platform string, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability, useUpstreamTokenCost bool) (*AccountSelectionResult, error) { - platform = normalizeOpenAICompatiblePlatform(platform) + platform = NormalizeOpenAICompatiblePlatform(platform) if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) { slog.Warn("channel pricing restriction blocked request", "group_id", derefGroupID(groupID), @@ -1304,7 +1311,7 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex } func (s *OpenAIGatewayService) listSchedulableAccounts(ctx context.Context, groupID *int64, platform string) ([]Account, error) { - platform = normalizeOpenAICompatiblePlatform(platform) + platform = NormalizeOpenAICompatiblePlatform(platform) if s.schedulerSnapshot != nil { accounts, _, err := s.schedulerSnapshot.ListSchedulableAccounts(ctx, groupID, platform, false) if err != nil { @@ -1357,7 +1364,7 @@ func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccountBeforeProfit( if account == nil { return nil } - platform = normalizeOpenAICompatiblePlatform(platform) + platform = NormalizeOpenAICompatiblePlatform(platform) fresh := account if s.schedulerSnapshot != nil { @@ -1415,7 +1422,7 @@ func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDBBeforeProfit(ct if account == nil { return nil } - platform = normalizeOpenAICompatiblePlatform(platform) + platform = NormalizeOpenAICompatiblePlatform(platform) if s.schedulerSnapshot == nil || s.accountRepo == nil { if !isOpenAICompatibleAccountEligibleForRequestBeforeProfit(ctx, account, platform, requestedModel, requireCompact, requiredCapability) { return nil diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index a9200b10b7..25b3333f5a 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -1205,7 +1205,7 @@ func (s *OpenAIGatewayService) GetAccessToken(ctx context.Context, account *Acco } return apiKey, "apikey", nil } - apiKey := account.GetOpenAIApiKey() + apiKey := strings.TrimSpace(account.GetOpenAIProtocolAPIKey()) if apiKey == "" { return "", "", errors.New("api_key not found in credentials") } diff --git a/backend/internal/service/openai_ws_forwarder_payload.go b/backend/internal/service/openai_ws_forwarder_payload.go index 4c02cb30f1..d783c98104 100644 --- a/backend/internal/service/openai_ws_forwarder_payload.go +++ b/backend/internal/service/openai_ws_forwarder_payload.go @@ -42,7 +42,7 @@ func (s *OpenAIGatewayService) buildOpenAIResponsesWSURL(account *Account) (stri if err != nil { return "", err } - targetURL = buildOpenAIResponsesURL(validatedURL) + targetURL = buildOpenAIResponsesURLForPlatform(account.Platform, validatedURL) } default: targetURL = openaiPlatformAPIURL diff --git a/backend/internal/service/ratelimit_cn_providers.go b/backend/internal/service/ratelimit_cn_providers.go new file mode 100644 index 0000000000..4aab5a8751 --- /dev/null +++ b/backend/internal/service/ratelimit_cn_providers.go @@ -0,0 +1,154 @@ +package service + +import ( + "context" + "log/slog" + "net/http" + "strings" + "time" +) + +// 国产供应商(kimi/zhipu/deepseek)的响应式冷却辅助。 +// +// 与 openai/anthropic 不同: +// - 余额不足是「可恢复」状态(充值/检测恢复后自动重新调度),不能走 handleAuthError +// 永久置 status=error。这里改为 SetTempUnschedulable,由 CN 余额检测周期任务 +// (cn_provider_balance_check_service.go)在余额恢复后 ClearTempUnschedulable。 +// - Coding Plan 滚动窗口耗尽(429)的冷却终点应是真实的窗口重置时间(已由 +// CNProviderQuotaService 落入 account.Extra 快照),而非默认的秒级兜底。 + +// cnBalanceExtraSuffixLow 标记账号响应过「余额不足」,供余额检测任务区分 +// 「确属余额不足」与「尚未探测」。 +const cnBalanceExtraSuffixLow = "balance_low" + +// cnBalanceLowReasonPrefix 是余额不足临时停调 reason 的稳定前缀。 +// 周期余额检测任务据此识别「是我们停调的」并在余额恢复后安全清除——不会误清 +// 其他子系统(阈值/限流/401)写入的临时停调。 +const cnBalanceLowReasonPrefix = "cn_balance_low" + +// cnBalanceLowReason 构造余额不足临时停调的 reason(带稳定前缀)。 +func cnBalanceLowReason(upstreamMsg string) string { + if upstreamMsg = strings.TrimSpace(upstreamMsg); upstreamMsg != "" { + return cnBalanceLowReasonPrefix + ": " + upstreamMsg + } + return cnBalanceLowReasonPrefix + ": 余额不足,账号临时停调" +} + +// cnProviderResponseIndicatesInsufficientBalance 通过响应体文案识别余额不足 +// (智谱 payg 无独立余额端点,仅能靠响应文案识别)。 +func cnProviderResponseIndicatesInsufficientBalance(body []byte) bool { + if len(body) == 0 { + return false + } + s := strings.ToLower(string(body)) + return strings.Contains(s, "余额不足") || + strings.Contains(s, "insufficient balance") || + strings.Contains(s, "insufficient_credit") || + strings.Contains(s, "balance is not enough") || + strings.Contains(s, "no enough balance") +} + +// handleCNProviderInsufficientBalance 把余额不足标记为可恢复的临时停调: +// 写入 balance_low 快照 + SetTempUnschedulable 一个余额检测周期, +// 由周期任务在余额恢复后清除。返回前已通知调度阻塞。 +func (s *RateLimitService) handleCNProviderInsufficientBalance( + ctx context.Context, + account *Account, + upstreamMsg string, +) { + msg := cnBalanceLowReason(upstreamMsg) + + if err := s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{ + cnExtraKey(account.Platform, cnBalanceExtraSuffixLow): true, + }); err != nil { + slog.Warn("cn_balance_low_mark_failed", "account_id", account.ID, "error", err) + } + + until := time.Now().Add(s.cnBalanceCooldownDuration()) + s.notifyAccountSchedulingBlocked(account, until, "cn_insufficient_balance") + if err := s.accountRepo.SetTempUnschedulable(ctx, account.ID, until, msg); err != nil { + slog.Warn("cn_balance_set_temp_unschedulable_failed", "account_id", account.ID, "error", err) + return + } + slog.Info("cn_provider_insufficient_balance", + "account_id", account.ID, + "platform", account.Platform, + "until", until.UTC(), + ) +} + +// cnBalanceCooldownDuration 返回余额不足临时停调的持续时长(= 2× 余额检测周期, +// 默认 20 分钟)。周期任务会在余额恢复后提前清除,故此处只需保证冷却覆盖到下一次 +// 周期检测即可。 +func (s *RateLimitService) cnBalanceCooldownDuration() time.Duration { + minutes := 10 + if s != nil && s.cfg != nil { + if cfgMin := s.cfg.Gateway.CNProviders.BalanceCheckIntervalMinutes; cfgMin > 0 { + minutes = cfgMin + } + } + cooldown := time.Duration(minutes) * time.Minute * 2 + if cooldown < time.Minute { + cooldown = 10 * time.Minute + } + return cooldown +} + +// cnProviderQuotaSnapshotReset 读取 Coding Plan 账号快照中最早一个仍在未来的窗口 +// 重置时间(5h / weekly)。429 多数由 5h 滚动窗口触发,取较早的重置点可避免 +// 把账号冷却到 weekly 重置(可达数天)的过度停调;如果确是 weekly 窗口耗尽, +// 周期额度探测刷新快照后阈值评估会再次停调到正确的时间点。 +// 无快照或均已过期返回 nil。 +func cnProviderQuotaSnapshotReset(account *Account, now time.Time) *time.Time { + if account == nil || !account.IsCNProvider() || !account.IsCodingPlan() || len(account.Extra) == 0 { + return nil + } + provider := account.Platform + var earliest *time.Time + for _, suffix := range []string{cnExtraSuffix5hReset, cnExtraSuffixWeeklyReset} { + t := parseSchedulingResetAt(account.Extra[cnExtraKey(provider, suffix)]) + if t == nil || !t.After(now) { + continue + } + if earliest == nil || t.Before(*earliest) { + earliest = t + } + } + return earliest +} + +// applyCNProviderReactive429 处理国产供应商的 429 响应。 +// 返回 true 表示已处理(调用方应 return),false 表示未命中、继续走默认 429 逻辑。 +func (s *RateLimitService) applyCNProviderReactive429( + ctx context.Context, + account *Account, + headers http.Header, + responseBody []byte, +) bool { + if !account.IsCNProvider() { + return false + } + // 1) 余额不足文案:可恢复临时停调(含智谱 payg 这类无余额端点的场景)。 + if cnProviderResponseIndicatesInsufficientBalance(responseBody) { + s.handleCNProviderInsufficientBalance(ctx, account, extractUpstreamErrorMessage(responseBody)) + return true + } + // 2) Coding Plan 窗口耗尽:冷却到快照中最早的窗口重置点(见 + // cnProviderQuotaSnapshotReset:429 多由 5h 窗口触发,取较早点避免过度停调)。 + if account.IsCodingPlan() { + if until := cnProviderQuotaSnapshotReset(account, time.Now()); until != nil { + s.notifyAccountSchedulingBlocked(account, *until, "429") + if err := s.accountRepo.SetRateLimited(ctx, account.ID, *until); err != nil { + slog.Warn("rate_limit_set_failed", "account_id", account.ID, "error", err) + return true + } + slog.Info("cn_coding_plan_rate_limited", + "account_id", account.ID, + "platform", account.Platform, + "reset_at", *until, + ) + return true + } + } + return false +} diff --git a/backend/internal/service/ratelimit_service.go b/backend/internal/service/ratelimit_service.go index ee5b09bbb4..0f14491914 100644 --- a/backend/internal/service/ratelimit_service.go +++ b/backend/internal/service/ratelimit_service.go @@ -436,6 +436,13 @@ func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Acc shouldDisable = true } case 402: + // 国产供应商:余额不足是可恢复状态(充值/检测恢复后由周期任务自动解除), + // 不能走 handleAuthError 永久置 status=error。改为可恢复的临时停调。 + if account.IsCNProvider() { + s.handleCNProviderInsufficientBalance(ctx, account, upstreamMsg) + shouldDisable = true + break + } // OpenAI: deactivated_workspace 表示工作区已停用,直接标记 error if account.Platform == PlatformOpenAI && gjson.GetBytes(responseBody, "detail.code").String() == "deactivated_workspace" { msg := "Workspace deactivated (402): workspace has been deactivated" @@ -1048,6 +1055,13 @@ func (s *RateLimitService) handle429(ctx context.Context, account *Account, head if account.IsShadow() { return } + // 国产供应商(kimi/zhipu/deepseek)的 429 走专用可恢复路径:余额不足 → 临时停调, + // Coding Plan 窗口耗尽 → 冷却到快照重置点。未命中则继续默认 429 逻辑。 + if account.IsCNProvider() { + if s.applyCNProviderReactive429(ctx, account, headers, responseBody) { + return + } + } // 1. OpenAI 平台:优先尝试解析 x-codex-* 响应头(用于 rate_limit_exceeded) if account.Platform == PlatformOpenAI { persistOpenAI429PlanType(ctx, s.accountRepo, account, responseBody) diff --git a/backend/internal/service/upstream_billing_probe.go b/backend/internal/service/upstream_billing_probe.go index d66cd905dc..c6a62bab09 100644 --- a/backend/internal/service/upstream_billing_probe.go +++ b/backend/internal/service/upstream_billing_probe.go @@ -965,7 +965,9 @@ func decodeUpstreamBillingProbeSnapshot(extra map[string]any) *UpstreamBillingPr // IsUpstreamBillingProbeIdentity reports whether an account identity may opt // in to the upstream billing probe. `/v1/sub2api/billing` is a key-scoped -// sub2api convention shared by the five supported API-key platforms. +// sub2api convention shared by the supported API-key platforms (including the +// CN providers, whose official-domain accounts are short-circuited to +// "unsupported" by upstreamBillingProbeTargetIsOfficialAPI). // Non-sub2api upstreams return 404 and the snapshot records "unsupported". // Only AccountTypeAPIKey is in scope. OAuth/Bedrock hold no static API key to // present at all; AccountTypeUpstream (antigravity relay accounts) does carry @@ -978,7 +980,8 @@ func IsUpstreamBillingProbeIdentity(platform, accountType string) bool { return false } switch platform { - case PlatformOpenAI, PlatformAnthropic, PlatformGemini, PlatformAntigravity, PlatformGrok: + case PlatformOpenAI, PlatformAnthropic, PlatformGemini, PlatformAntigravity, PlatformGrok, + PlatformKimi, PlatformZhipu, PlatformDeepseek: return true default: return false @@ -1002,6 +1005,9 @@ func isUpstreamBillingProbeAccount(account *Account) bool { // ollama.com is a first-class configuration here (Ollama Cloud accounts are // platform openai/anthropic with base_url https://ollama.com/v1), and it is // an official provider API just like the rest, so it belongs on this list. +// CN provider domains (moonshot.cn / kimi.com / bigmodel.cn / deepseek.com) +// serve the same role: official APIs that can never host /v1/sub2api/billing, +// so their accounts short-circuit to "unsupported" without a request. var upstreamBillingProbeOfficialAPIDomains = []string{ "anthropic.com", "googleapis.com", @@ -1009,6 +1015,10 @@ var upstreamBillingProbeOfficialAPIDomains = []string{ "grok.com", "openai.com", "ollama.com", + "moonshot.cn", + "kimi.com", + "bigmodel.cn", + "deepseek.com", } func upstreamBillingProbeTargetIsOfficialAPI(baseURL string) bool { diff --git a/backend/internal/service/upstream_billing_probe_multiplatform_test.go b/backend/internal/service/upstream_billing_probe_multiplatform_test.go index 018f2b7198..6690a80c15 100644 --- a/backend/internal/service/upstream_billing_probe_multiplatform_test.go +++ b/backend/internal/service/upstream_billing_probe_multiplatform_test.go @@ -11,11 +11,12 @@ import ( "github.com/stretchr/testify/require" ) -// 探测资格:/v1/sub2api/billing 是 key 级端点,五个 -// 受支持平台的 API-key 账号都可开启探测;OAuth/Bedrock 无静态 Key 仍不合格。 +// 探测资格:/v1/sub2api/billing 是 key 级端点,全部 +// 受支持平台(含国产供应商)的 API-key 账号都可开启探测;OAuth/Bedrock 无静态 Key 仍不合格。 func TestUpstreamBillingProbeIdentityCoversAllAPIKeyPlatforms(t *testing.T) { for _, platform := range []string{ PlatformOpenAI, PlatformGrok, PlatformAnthropic, PlatformGemini, PlatformAntigravity, + PlatformKimi, PlatformZhipu, PlatformDeepseek, } { require.True(t, IsUpstreamBillingProbeIdentity(platform, AccountTypeAPIKey), platform) require.True(t, isUpstreamBillingProbeAccount(&Account{Platform: platform, Type: AccountTypeAPIKey}), platform) @@ -124,6 +125,14 @@ func TestUpstreamBillingProbeOfficialAPIBaseURLIsUnsupportedWithoutRequest(t *te {PlatformAnthropic, "https://ollama.com/v1"}, {PlatformAnthropic, "https://ollama.com"}, {PlatformAnthropic, "https://www.ollama.com/v1"}, + // 国产供应商官方域(含各协议端点)同样是官方 API,创建即开探测也不发请求。 + {PlatformKimi, "https://api.moonshot.cn/v1"}, + {PlatformKimi, "https://api.moonshot.cn/anthropic"}, + {PlatformKimi, "https://api.kimi.com/coding"}, + {PlatformZhipu, "https://open.bigmodel.cn/api/paas/v4"}, + {PlatformZhipu, "https://open.bigmodel.cn/api/anthropic"}, + {PlatformDeepseek, "https://api.deepseek.com"}, + {PlatformDeepseek, "https://api.deepseek.com/anthropic"}, } for i, tc := range cases { account := &Account{ @@ -160,6 +169,11 @@ func TestUpstreamBillingProbeOfficialAPIHostMatchingIsNormalized(t *testing.T) { require.True(t, upstreamBillingProbeTargetIsOfficialAPI("https://ollama.com:443/v1")) require.True(t, upstreamBillingProbeTargetIsOfficialAPI("https://www.ollama.com/v1")) require.True(t, upstreamBillingProbeTargetIsOfficialAPI("HTTPS://OLLAMA.COM./v1")) + // 国产供应商官方域及子域。 + require.True(t, upstreamBillingProbeTargetIsOfficialAPI("https://api.moonshot.cn/v1")) + require.True(t, upstreamBillingProbeTargetIsOfficialAPI("https://api.kimi.com/coding/v1")) + require.True(t, upstreamBillingProbeTargetIsOfficialAPI("https://open.bigmodel.cn/api/anthropic")) + require.True(t, upstreamBillingProbeTargetIsOfficialAPI("https://api.deepseek.com/anthropic")) // 相似但不同的注册域不拦:中转完全可能叫 *-x.ai 之外的任何名字。 require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://relay.example/v1")) require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://notx.ai")) @@ -168,6 +182,11 @@ func TestUpstreamBillingProbeOfficialAPIHostMatchingIsNormalized(t *testing.T) { require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://notollama.com/v1")) require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://ollama.com.evil.example/v1")) require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://ollama.example/v1")) + require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://notmoonshot.cn/v1")) + require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://moonshot.cn.evil.example/v1")) + require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://kimi.example/v1")) + require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://notbigmodel.cn")) + require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://deepseek.example.com")) } // OpenAI 语义保持不变:无自定义 base 时仍探官方域,且沿用 openai 传输画像。 diff --git a/backend/internal/service/upstream_models.go b/backend/internal/service/upstream_models.go index 9a4e4311b0..df926b0297 100644 --- a/backend/internal/service/upstream_models.go +++ b/backend/internal/service/upstream_models.go @@ -137,7 +137,8 @@ func (s *AccountTestService) buildUpstreamModelsRequest(ctx context.Context, acc return s.buildAntigravityAPIKeyModelsRequest(ctx, account) case account.IsGrok(): return s.buildGrokUpstreamModelsRequest(ctx, account) - case account.IsOpenAI(): + case account.IsOpenAI() || account.IsCNProvider(): + // 国产 OpenAI 兼容供应商(kimi/zhipu/deepseek)复用 OpenAI /v1/models 探测。 return s.buildOpenAIUpstreamModelsRequest(ctx, account) case account.IsGemini(): return s.buildGeminiUpstreamModelsRequest(ctx, account) @@ -347,12 +348,14 @@ func (s *AccountTestService) buildOpenAIUpstreamModelsRequest(ctx context.Contex fmt.Sprintf("Unsupported OpenAI account type for upstream model sync: %s", account.Type), nil, ) } - apiKey := strings.TrimSpace(account.GetOpenAIApiKey()) + apiKey := strings.TrimSpace(account.GetOpenAIProtocolAPIKey()) if apiKey == "" { return nil, newUpstreamModelSyncConfigError("No OpenAI API key is available", nil) } - baseURL := account.GetOpenAIBaseURL() + // 协议感知:Anthropic 协议账号的凭证 base_url 指向 /anthropic 端点,模型 + // 列表同步需使用 OpenAI 格式 base(供应商 × 模式默认)。 + baseURL := account.GetOpenAIFormatBaseURL() if strings.TrimSpace(baseURL) == "" { baseURL = "https://api.openai.com" } diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index 2e26c0035c..d6bb56fac4 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -261,6 +261,45 @@ func ProvideGrokQuotaService( return service } +// ProvideCNProviderQuotaService 构造国产供应商 Coding Plan 额度探测服务。 +func ProvideCNProviderQuotaService( + accountRepo AccountRepository, + proxyRepo ProxyRepository, + httpUpstream HTTPUpstream, + cfg *config.Config, +) *CNProviderQuotaService { + return NewCNProviderQuotaService(accountRepo, proxyRepo, httpUpstream, cfg) +} + +// ProvideCNProviderBalanceService 构造国产供应商余额探测服务。 +func ProvideCNProviderBalanceService( + accountRepo AccountRepository, + proxyRepo ProxyRepository, + httpUpstream HTTPUpstream, + cfg *config.Config, +) *CNProviderBalanceService { + return NewCNProviderBalanceService(accountRepo, proxyRepo, httpUpstream, cfg) +} + +// ProvideCNProviderBalanceCheckService 构造并启动周期余额/额度检测任务。 +// payg 账号探余额(低余额停调);coding plan 账号探 5h/weekly 滚动窗口 +// (落 extra 快照供调度阈值评估自动停调)。 +// 间隔取自 gateway.cn_providers.balance_check_interval_minutes;<=0 或关闭时不启动。 +func ProvideCNProviderBalanceCheckService( + accountRepo AccountRepository, + balanceService *CNProviderBalanceService, + quotaService *CNProviderQuotaService, + cfg *config.Config, +) *CNProviderBalanceCheckService { + minutes := 10 + if cfg != nil && cfg.Gateway.CNProviders.BalanceCheckIntervalMinutes > 0 { + minutes = cfg.Gateway.CNProviders.BalanceCheckIntervalMinutes + } + svc := NewCNProviderBalanceCheckService(accountRepo, balanceService, quotaService, cfg, time.Duration(minutes)*time.Minute) + svc.Start() + return svc +} + // ProvideGeminiTokenProvider creates GeminiTokenProvider with OAuthRefreshAPI injection func ProvideGeminiTokenProvider( accountRepo AccountRepository, @@ -796,6 +835,9 @@ var ProviderSet = wire.NewSet( ProvideOpenAITokenProvider, ProvideOpenAIQuotaService, ProvideGrokQuotaService, + ProvideCNProviderQuotaService, + ProvideCNProviderBalanceService, + ProvideCNProviderBalanceCheckService, ProvideClaudeTokenProvider, NewAntigravityGatewayService, ProvideRateLimitService, diff --git a/backend/migrations/224_user_platform_quotas_add_cn_providers.sql b/backend/migrations/224_user_platform_quotas_add_cn_providers.sql new file mode 100644 index 0000000000..011a7762af --- /dev/null +++ b/backend/migrations/224_user_platform_quotas_add_cn_providers.sql @@ -0,0 +1,17 @@ +-- 把 kimi/zhipu/deepseek 平台加入 user_platform_quotas.platform 的 CHECK 约束。 +-- +-- 背景:国产供应商进入 AllowedQuotaPlatforms(internal/service/domain_constants.go), +-- 注册时 GetDefaultPlatformQuotas 会为全部 8 平台预填充默认配额行,但 157 号迁移的 +-- CHECK 仍只允许 5 平台。BulkInsertInitial 是单条多行 INSERT,任一违约行会中止整条 +-- 语句 → 注册路径 fail-open 吞错 → 新用户拿到零条配额记录(含原有 5 平台,缺失配额 +-- 行 = 无限额)。与 157 头注释记载的 grok 同型事故一致。 +-- +-- 修复:把约束与代码平台列表(PlatformKimi/PlatformZhipu/PlatformDeepseek)对齐。 +-- DROP ... IF EXISTS 保证可重入;新约束是旧约束的超集,存量行(仅 5 平台)瞬时校验通过。 +ALTER TABLE user_platform_quotas + DROP CONSTRAINT IF EXISTS user_platform_quotas_platform_check; + +ALTER TABLE user_platform_quotas + ADD CONSTRAINT user_platform_quotas_platform_check + CHECK (platform IN ('anthropic', 'openai', 'gemini', 'antigravity', 'grok', + 'kimi', 'zhipu', 'deepseek')); diff --git a/backend/migrations/user_platform_quota_cn_providers_migration_test.go b/backend/migrations/user_platform_quota_cn_providers_migration_test.go new file mode 100644 index 0000000000..73bfd05767 --- /dev/null +++ b/backend/migrations/user_platform_quota_cn_providers_migration_test.go @@ -0,0 +1,22 @@ +package migrations + +import ( + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +// TestUserPlatformQuotasCNProvidersMigration 校验 224 号迁移把 kimi/zhipu/deepseek +// 加入 user_platform_quotas.platform 的 CHECK 约束(对照 157 号 grok 迁移)。 +// 约束未放宽时,注册预填充 8 平台默认配额会整条 INSERT 中止 → 新用户零配额行 +// (缺失配额行 = 无限额),管理端设置国产平台配额直接 500。 +func TestUserPlatformQuotasCNProvidersMigration(t *testing.T) { + content, err := FS.ReadFile("224_user_platform_quotas_add_cn_providers.sql") + require.NoError(t, err) + + sql := strings.Join(strings.Fields(string(content)), " ") + require.Contains(t, sql, "DROP CONSTRAINT IF EXISTS user_platform_quotas_platform_check") + require.Contains(t, sql, + "CHECK (platform IN ('anthropic', 'openai', 'gemini', 'antigravity', 'grok', 'kimi', 'zhipu', 'deepseek'))") +} diff --git a/frontend/src/api/admin/cnProviders.ts b/frontend/src/api/admin/cnProviders.ts new file mode 100644 index 0000000000..668f67a73e --- /dev/null +++ b/frontend/src/api/admin/cnProviders.ts @@ -0,0 +1,70 @@ +/** + * Admin CN providers (Kimi / Zhipu / DeepSeek) API endpoints. + * Coding-plan rolling-window quota probe + payg balance probe. + */ + +import { apiClient } from '../client' + +/** 滚动用量窗口档(5 小时 / 每周),对齐后端 service.CNQuotaTier。 */ +export interface CNQuotaTier { + window: '5h' | 'weekly' + used_percent: number + reset_at?: string +} + +/** Coding Plan 额度探测结果(kimi / zhipu),对齐后端 CNProviderQuotaProbeResult。 */ +export interface CNProviderQuotaProbeResult { + provider: string + source?: string + success: boolean + credential_valid: boolean + tiers?: CNQuotaTier[] + plan_level?: string + status_code?: number + fetched_at: number + persisted: boolean + error?: string +} + +/** 单币种余额明细(deepseek 双币种账号含 CNY + USD 两条)。 */ +export interface CNProviderBalanceEntry { + currency: string + balance: number +} + +/** payg 余额探测结果(kimi / deepseek),对齐后端 CNProviderBalanceResult。 */ +export interface CNProviderBalanceResult { + provider: string + success: boolean + /** 主币种余额(balances 首条,兼容单币种展示)。 */ + balance: number + currency?: string + /** 多币种明细;缺省时按主币种展示。 */ + balances?: CNProviderBalanceEntry[] + available: boolean + status_code?: number + fetched_at: number + persisted: boolean + error?: string +} + +/** 查询 Coding Plan 滚动窗口用量(5h + weekly)。 */ +export async function queryQuota(id: number): Promise { + const { data } = await apiClient.get( + `/admin/cn-providers/accounts/${id}/quota` + ) + return data +} + +/** 查询 payg 账号余额。 */ +export async function queryBalance(id: number): Promise { + const { data } = await apiClient.get( + `/admin/cn-providers/accounts/${id}/balance` + ) + return data +} + +export default { + queryQuota, + queryBalance +} diff --git a/frontend/src/api/admin/index.ts b/frontend/src/api/admin/index.ts index 80a7073dec..dd976daad5 100644 --- a/frontend/src/api/admin/index.ts +++ b/frontend/src/api/admin/index.ts @@ -18,6 +18,7 @@ import usageAPI from './usage' import geminiAPI from './gemini' import antigravityAPI from './antigravity' import grokAPI from './grok' +import cnProvidersAPI from './cnProviders' import userAttributesAPI from './userAttributes' import opsAPI from './ops' import errorPassthroughAPI from './errorPassthrough' @@ -54,6 +55,7 @@ export const adminAPI = { gemini: geminiAPI, antigravity: antigravityAPI, grok: grokAPI, + cnProviders: cnProvidersAPI, userAttributes: userAttributesAPI, ops: opsAPI, errorPassthrough: errorPassthroughAPI, @@ -88,6 +90,7 @@ export { geminiAPI, antigravityAPI, grokAPI, + cnProvidersAPI, userAttributesAPI, opsAPI, errorPassthroughAPI, diff --git a/frontend/src/api/admin/settings.ts b/frontend/src/api/admin/settings.ts index b176024ae6..c8dbe8fb78 100644 --- a/frontend/src/api/admin/settings.ts +++ b/frontend/src/api/admin/settings.ts @@ -36,13 +36,19 @@ export type SchedulingThresholdPlatformType = | "openai" | "anthropic" | "grok" + | "kimi" + | "zhipu" export type AccountSchedulingThresholdsMap = Record +// 与后端 AllowedSchedulingThresholdPlatforms 保持一致(deepseek 为余额型, +// 走余额检测而非用量阈值)。 export const SCHEDULING_THRESHOLD_PLATFORMS: SchedulingThresholdPlatformType[] = [ "openai", "anthropic", "grok", + "kimi", + "zhipu", ] export function normalizeAccountSchedulingThresholdsMap( diff --git a/frontend/src/components/account/AccountUsageCell.vue b/frontend/src/components/account/AccountUsageCell.vue index 485b4b8701..ae9d55e3c2 100644 --- a/frontend/src/components/account/AccountUsageCell.vue +++ b/frontend/src/components/account/AccountUsageCell.vue @@ -422,6 +422,21 @@ + + +