mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:28:39 +08:00
Merge pull request #5666 from Randark-JMT/feat/cn-providers-kimi-zhipu-deepseek
feat: 国产供应商多协议支持(Kimi/Zhipu/DeepSeek 原生 Anthropic 直通 + DeepSeek Responses)与配额/余额监控
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -66,6 +66,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) {
|
||||
schedulerSnapshotSvc,
|
||||
tokenRefreshSvc,
|
||||
accountExpirySvc,
|
||||
nil, // cnProviderBalanceCheck
|
||||
codexVersionSyncSvc,
|
||||
proxyExpirySvc,
|
||||
subscriptionExpirySvc,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
{
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 ""
|
||||
|
||||
@@ -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,键形如
|
||||
// <provider>_5h_used_percent / <provider>_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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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 ""
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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 标识。
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,导致只支持
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 传输画像。
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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'));
|
||||
@@ -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'))")
|
||||
}
|
||||
@@ -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<CNProviderQuotaProbeResult> {
|
||||
const { data } = await apiClient.get<CNProviderQuotaProbeResult>(
|
||||
`/admin/cn-providers/accounts/${id}/quota`
|
||||
)
|
||||
return data
|
||||
}
|
||||
|
||||
/** 查询 payg 账号余额。 */
|
||||
export async function queryBalance(id: number): Promise<CNProviderBalanceResult> {
|
||||
const { data } = await apiClient.get<CNProviderBalanceResult>(
|
||||
`/admin/cn-providers/accounts/${id}/balance`
|
||||
)
|
||||
return data
|
||||
}
|
||||
|
||||
export default {
|
||||
queryQuota,
|
||||
queryBalance
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -36,13 +36,19 @@ export type SchedulingThresholdPlatformType =
|
||||
| "openai"
|
||||
| "anthropic"
|
||||
| "grok"
|
||||
| "kimi"
|
||||
| "zhipu"
|
||||
|
||||
export type AccountSchedulingThresholdsMap = Record<SchedulingThresholdPlatformType, number>
|
||||
|
||||
// 与后端 AllowedSchedulingThresholdPlatforms 保持一致(deepseek 为余额型,
|
||||
// 走余额检测而非用量阈值)。
|
||||
export const SCHEDULING_THRESHOLD_PLATFORMS: SchedulingThresholdPlatformType[] = [
|
||||
"openai",
|
||||
"anthropic",
|
||||
"grok",
|
||||
"kimi",
|
||||
"zhipu",
|
||||
]
|
||||
|
||||
export function normalizeAccountSchedulingThresholdsMap(
|
||||
|
||||
@@ -422,6 +422,21 @@
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<!-- CN providers (Kimi / Zhipu / DeepSeek): coding-plan quota or payg balance -->
|
||||
<template v-else-if="account.platform === 'kimi' || account.platform === 'zhipu' || account.platform === 'deepseek'">
|
||||
<div class="space-y-1">
|
||||
<!-- 子单元格各自按 模式×平台 判定可见;两者都不可见时(智谱 payg 无公开
|
||||
余额端点、coding 探测也不适用)才回落到占位符。 -->
|
||||
<div
|
||||
v-if="!cnQuotaCellVisible && !cnBalanceCellVisible"
|
||||
class="text-xs text-gray-400"
|
||||
:title="t('admin.accounts.cnProviders.noBalanceEndpoint')"
|
||||
>-</div>
|
||||
<CNProviderQuotaCell :account="account" />
|
||||
<CNProviderBalanceCell :account="account" />
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<!-- Gemini platform: show quota + local usage window -->
|
||||
<template v-else-if="account.platform === 'gemini'">
|
||||
<!-- Auth Type + Tier Badge (first line) -->
|
||||
@@ -625,7 +640,10 @@ import UsageProgressBar from './UsageProgressBar.vue'
|
||||
import AccountQuotaInfo from './AccountQuotaInfo.vue'
|
||||
import OpenAIQuotaResetCell from './OpenAIQuotaResetCell.vue'
|
||||
import GrokQuotaProbeCell from './GrokQuotaProbeCell.vue'
|
||||
import CNProviderQuotaCell from './CNProviderQuotaCell.vue'
|
||||
import CNProviderBalanceCell from './CNProviderBalanceCell.vue'
|
||||
import OllamaCloudUsageCell from './OllamaCloudUsageCell.vue'
|
||||
import { cnQuotaCellVisible as cnQuotaCellVisibleFn, cnBalanceCellVisible as cnBalanceCellVisibleFn } from './credentialsBuilder'
|
||||
|
||||
// Module-level cache shared across all AccountUsageCell instances
|
||||
const _usageCache = new Map<number, { data: AccountUsageInfo; ts: number }>()
|
||||
@@ -690,6 +708,15 @@ let visibilityObserver: IntersectionObserver | null = null
|
||||
const showUsageWindows = computed(() => {
|
||||
// Gemini: we can always compute local usage windows from DB logs (simulated quotas).
|
||||
if (props.account.platform === 'gemini') return true
|
||||
// CN providers: apikey 账号也有滚动用量窗口(coding plan)或余额(payg),
|
||||
// 由 CNProviderQuotaCell / CNProviderBalanceCell 自行探测与展示。
|
||||
if (
|
||||
props.account.platform === 'kimi' ||
|
||||
props.account.platform === 'zhipu' ||
|
||||
props.account.platform === 'deepseek'
|
||||
) {
|
||||
return true
|
||||
}
|
||||
return props.account.type === 'oauth' || props.account.type === 'setup-token'
|
||||
})
|
||||
|
||||
@@ -712,6 +739,15 @@ const shouldFetchUsage = computed(() => {
|
||||
return false
|
||||
})
|
||||
|
||||
// CN 供应商子单元格可见性(与 CNProviderQuotaCell / CNProviderBalanceCell 共用
|
||||
// credentialsBuilder 的单一实现):都不可见时显示 `-` 占位符。
|
||||
const cnAccountMode = computed(() => {
|
||||
const mode = props.account.credentials?.account_mode
|
||||
return typeof mode === 'string' ? mode : ''
|
||||
})
|
||||
const cnQuotaCellVisible = computed(() => cnQuotaCellVisibleFn(props.account.platform, cnAccountMode.value))
|
||||
const cnBalanceCellVisible = computed(() => cnBalanceCellVisibleFn(props.account.platform, cnAccountMode.value))
|
||||
|
||||
const isBatchManaged = computed(() => typeof props.requestBatchedUsage === 'function')
|
||||
|
||||
const showGeminiTodayStats = computed(() => {
|
||||
|
||||
@@ -0,0 +1,168 @@
|
||||
<template>
|
||||
<div v-if="visible" class="space-y-1">
|
||||
<div class="flex flex-wrap items-center gap-1.5">
|
||||
<button
|
||||
type="button"
|
||||
:class="[
|
||||
'inline-flex items-center gap-0.5 rounded px-1.5 py-0.5 text-[10px] font-medium transition-colors hover:bg-gray-100 disabled:cursor-not-allowed disabled:opacity-50 dark:hover:bg-dark-600',
|
||||
platformTextClass(account.platform)
|
||||
]"
|
||||
:disabled="loading"
|
||||
@click="handleProbe"
|
||||
>
|
||||
<svg
|
||||
class="h-2.5 w-2.5"
|
||||
:class="{ 'animate-spin': loading }"
|
||||
fill="none"
|
||||
stroke="currentColor"
|
||||
viewBox="0 0 24 24"
|
||||
>
|
||||
<path
|
||||
stroke-linecap="round"
|
||||
stroke-linejoin="round"
|
||||
stroke-width="2"
|
||||
d="M4 4v5h.582m15.356 2A8.001 8.001 0 004.582 9m0 0H9m11 11v-5h-.581m0 0a8.003 8.003 0 01-15.357-2m15.357 2H15"
|
||||
/>
|
||||
</svg>
|
||||
{{ balanceLabel }}
|
||||
</button>
|
||||
|
||||
<!-- Low balance badge (reactive 402/429 marker or probe-detected) -->
|
||||
<span
|
||||
v-if="balanceLow"
|
||||
class="inline-flex items-center rounded bg-red-100 px-1 py-0.5 text-[10px] font-medium text-red-700 dark:bg-red-900/30 dark:text-red-300"
|
||||
>
|
||||
{{ t('admin.accounts.cnProviders.balanceLow') }}
|
||||
</span>
|
||||
</div>
|
||||
|
||||
<div v-if="error" class="truncate text-[10px] text-red-600 dark:text-red-400" :title="error">
|
||||
{{ truncatedError }}
|
||||
</div>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { computed, ref, watch } from 'vue'
|
||||
import { useI18n } from 'vue-i18n'
|
||||
import { adminAPI } from '@/api/admin'
|
||||
import type { CNProviderBalanceEntry, CNProviderBalanceResult } from '@/api/admin/cnProviders'
|
||||
import type { Account } from '@/types'
|
||||
import { platformTextClass } from '@/utils/platformColors'
|
||||
import { cnBalanceCellVisible } from './credentialsBuilder'
|
||||
|
||||
const props = defineProps<{
|
||||
account: Account
|
||||
}>()
|
||||
|
||||
const { t } = useI18n()
|
||||
|
||||
const readMode = (): string => {
|
||||
const mode = props.account.credentials?.account_mode
|
||||
return typeof mode === 'string' ? mode : ''
|
||||
}
|
||||
|
||||
// 仅 kimi / deepseek payg 账号有公开余额端点(智谱 payg 无)。
|
||||
const visible = computed(() => cnBalanceCellVisible(props.account.platform, readMode()))
|
||||
|
||||
const loading = ref(false)
|
||||
const error = ref<string | null>(null)
|
||||
const data = ref<CNProviderBalanceResult | null>(null)
|
||||
|
||||
const extraKey = (suffix: string) => `${props.account.platform}_${suffix}`
|
||||
|
||||
// 落库快照(后端周期探测/响应式写入 account.Extra)。
|
||||
const snapshotBalance = computed(() => {
|
||||
const v = props.account.extra?.[extraKey('balance')]
|
||||
return typeof v === 'number' ? v : null
|
||||
})
|
||||
const snapshotCurrency = computed(() => {
|
||||
const v = props.account.extra?.[extraKey('balance_currency')]
|
||||
return typeof v === 'string' ? v : ''
|
||||
})
|
||||
// 多币种快照(后端写 <platform>_balances:[{currency, balance}],deepseek CNY+USD)。
|
||||
const snapshotBalances = computed<CNProviderBalanceEntry[]>(() => {
|
||||
const v = props.account.extra?.[extraKey('balances')]
|
||||
if (!Array.isArray(v)) return []
|
||||
return v.flatMap((item): CNProviderBalanceEntry[] => {
|
||||
if (!item || typeof item !== 'object') return []
|
||||
const { currency, balance } = item as Record<string, unknown>
|
||||
if (typeof currency !== 'string' || typeof balance !== 'number') return []
|
||||
return [{ currency, balance }]
|
||||
})
|
||||
})
|
||||
const balanceLow = computed(() => props.account.extra?.[extraKey('balance_low')] === true)
|
||||
|
||||
// 优先用探测结果,其次落库快照。多币种返回全部明细,否则主币种单条。
|
||||
const currentEntries = computed<CNProviderBalanceEntry[]>(() => {
|
||||
if (data.value && data.value.success) {
|
||||
if (data.value.balances && data.value.balances.length > 0) return data.value.balances
|
||||
return [{ currency: data.value.currency || '', balance: data.value.balance }]
|
||||
}
|
||||
if (snapshotBalances.value.length > 0) return snapshotBalances.value
|
||||
if (snapshotBalance.value != null) {
|
||||
return [{ currency: snapshotCurrency.value, balance: snapshotBalance.value }]
|
||||
}
|
||||
return []
|
||||
})
|
||||
|
||||
const formatEntry = (entry: CNProviderBalanceEntry): string => {
|
||||
const fixed = entry.balance >= 100 ? entry.balance.toFixed(0) : entry.balance.toFixed(2)
|
||||
return `${entry.currency || '¥'} ${fixed}`
|
||||
}
|
||||
|
||||
const balanceLabel = computed(() => {
|
||||
if (currentEntries.value.length === 0) {
|
||||
return t('admin.accounts.grokBalance')?.trim() || 'Balance'
|
||||
}
|
||||
return currentEntries.value.map(formatEntry).join(' · ')
|
||||
})
|
||||
|
||||
const extractErrorMessage = (e: unknown): string => {
|
||||
const err = e as {
|
||||
message?: string
|
||||
reason?: string
|
||||
response?: { data?: { message?: string; error?: string } }
|
||||
}
|
||||
return (
|
||||
err?.message ||
|
||||
err?.reason ||
|
||||
err?.response?.data?.message ||
|
||||
err?.response?.data?.error ||
|
||||
t('common.error')
|
||||
)
|
||||
}
|
||||
|
||||
const truncatedError = computed(() => {
|
||||
if (!error.value) return ''
|
||||
return error.value.length > 80 ? `${error.value.slice(0, 80)}...` : error.value
|
||||
})
|
||||
|
||||
const handleProbe = async () => {
|
||||
if (loading.value) return
|
||||
loading.value = true
|
||||
error.value = null
|
||||
try {
|
||||
const result = await adminAPI.cnProviders.queryBalance(props.account.id)
|
||||
// 失败时保留快照展示(仅显示错误行),成功才覆盖。
|
||||
if (result.success) {
|
||||
data.value = result
|
||||
} else {
|
||||
error.value = result.error || t('common.error')
|
||||
}
|
||||
} catch (e) {
|
||||
error.value = extractErrorMessage(e)
|
||||
} finally {
|
||||
loading.value = false
|
||||
}
|
||||
}
|
||||
|
||||
watch(
|
||||
() => props.account.id,
|
||||
() => {
|
||||
data.value = null
|
||||
error.value = null
|
||||
loading.value = false
|
||||
}
|
||||
)
|
||||
</script>
|
||||
@@ -0,0 +1,219 @@
|
||||
<template>
|
||||
<div v-if="visible" class="space-y-1">
|
||||
<div class="flex flex-wrap items-center gap-1.5">
|
||||
<button
|
||||
type="button"
|
||||
:class="[
|
||||
'inline-flex items-center gap-0.5 rounded px-1.5 py-0.5 text-[10px] font-medium transition-colors hover:bg-gray-100 disabled:cursor-not-allowed disabled:opacity-50 dark:hover:bg-dark-600',
|
||||
platformTextClass(account.platform)
|
||||
]"
|
||||
:disabled="loading"
|
||||
:title="t('admin.accounts.cnProviders.probeTooltip')"
|
||||
@click="handleProbe()"
|
||||
>
|
||||
<svg
|
||||
class="h-2.5 w-2.5"
|
||||
:class="{ 'animate-spin': loading }"
|
||||
fill="none"
|
||||
stroke="currentColor"
|
||||
viewBox="0 0 24 24"
|
||||
>
|
||||
<path
|
||||
stroke-linecap="round"
|
||||
stroke-linejoin="round"
|
||||
stroke-width="2"
|
||||
d="M4 4v5h.582m15.356 2A8.001 8.001 0 004.582 9m0 0H9m11 11v-5h-.581m0 0a8.003 8.003 0 01-15.357-2m15.357 2H15"
|
||||
/>
|
||||
</svg>
|
||||
{{ t('admin.accounts.cnProviders.window5h') }}/{{ t('admin.accounts.cnProviders.windowWeekly') }}
|
||||
</button>
|
||||
</div>
|
||||
|
||||
<!-- Tier rows: 5h + weekly utilization bars -->
|
||||
<div v-if="data?.success && data.tiers?.length" class="space-y-1">
|
||||
<div v-for="tier in data.tiers" :key="tier.window" class="flex items-center gap-1.5 text-[10px]">
|
||||
<span class="w-10 shrink-0 text-gray-500 dark:text-gray-400">{{ windowLabel(tier.window) }}</span>
|
||||
<div class="h-1.5 w-16 shrink-0 overflow-hidden rounded-full bg-gray-200 dark:bg-dark-600">
|
||||
<div
|
||||
class="h-full rounded-full transition-all"
|
||||
:class="utilizationColor(tier.used_percent)"
|
||||
:style="{ width: `${Math.min(100, Math.max(0, tier.used_percent))}%` }"
|
||||
/>
|
||||
</div>
|
||||
<span :class="['shrink-0 font-medium', utilizationTextColor(tier.used_percent)]">
|
||||
{{ Math.round(tier.used_percent) }}%
|
||||
</span>
|
||||
<span v-if="tier.reset_at" class="truncate text-gray-400 dark:text-gray-500" :title="tier.reset_at">
|
||||
· {{ formatReset(tier.reset_at) }}
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div v-if="error" class="truncate text-[10px] text-red-600 dark:text-red-400" :title="error">
|
||||
{{ truncatedError }}
|
||||
</div>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { computed, onMounted, ref, watch } from 'vue'
|
||||
import { useI18n } from 'vue-i18n'
|
||||
import { adminAPI } from '@/api/admin'
|
||||
import type { CNProviderQuotaProbeResult } from '@/api/admin/cnProviders'
|
||||
import type { Account } from '@/types'
|
||||
import { platformTextClass } from '@/utils/platformColors'
|
||||
import { cnQuotaCellVisible } from './credentialsBuilder'
|
||||
|
||||
const props = defineProps<{
|
||||
account: Account
|
||||
}>()
|
||||
|
||||
const { t } = useI18n()
|
||||
|
||||
const readMode = (): string => {
|
||||
const mode = props.account.credentials?.account_mode
|
||||
return typeof mode === 'string' ? mode : ''
|
||||
}
|
||||
|
||||
const visible = computed(() => cnQuotaCellVisible(props.account.platform, readMode()))
|
||||
|
||||
const loading = ref(false)
|
||||
const error = ref<string | null>(null)
|
||||
const data = ref<CNProviderQuotaProbeResult | null>(null)
|
||||
|
||||
// 后端周期任务/手动探测写入的 extra 快照键(<provider>_ 前缀,与后端
|
||||
// cnQuotaExtraUpdates 对齐)。页面加载即有数据,无需等待探测。
|
||||
const SNAPSHOT_STALE_MS = 15 * 60 * 1000
|
||||
|
||||
// 自动探测去抖窗口与最近一次自动探测时间(模块级,跨实例共享)。
|
||||
const AUTO_PROBE_DEBOUNCE_MS = 5 * 60 * 1000
|
||||
const lastAutoProbeAt = new Map<number, number>()
|
||||
|
||||
const readExtraNumber = (key: string): number | null => {
|
||||
const v = (props.account.extra as Record<string, unknown> | undefined)?.[key]
|
||||
return typeof v === 'number' && Number.isFinite(v) ? v : null
|
||||
}
|
||||
|
||||
const readExtraString = (key: string): string => {
|
||||
const v = (props.account.extra as Record<string, unknown> | undefined)?.[key]
|
||||
return typeof v === 'string' ? v : ''
|
||||
}
|
||||
|
||||
// 从持久化快照构造展示数据(缺少 5h/weekly 两档键时返回 null)。
|
||||
const snapshotData = computed<CNProviderQuotaProbeResult | null>(() => {
|
||||
const platform = props.account.platform
|
||||
const used5h = readExtraNumber(`${platform}_5h_used_percent`)
|
||||
const usedWeekly = readExtraNumber(`${platform}_weekly_used_percent`)
|
||||
if (used5h == null && usedWeekly == null) return null
|
||||
const tiers: CNProviderQuotaProbeResult['tiers'] = []
|
||||
if (used5h != null) {
|
||||
tiers.push({ window: '5h', used_percent: used5h, reset_at: readExtraString(`${platform}_5h_reset_at`) || undefined })
|
||||
}
|
||||
if (usedWeekly != null) {
|
||||
tiers.push({ window: 'weekly', used_percent: usedWeekly, reset_at: readExtraString(`${platform}_weekly_reset_at`) || undefined })
|
||||
}
|
||||
return { success: true, tiers } as CNProviderQuotaProbeResult
|
||||
})
|
||||
|
||||
// 快照是否过期(无更新时间或超过 staleness 窗口)→ 挂载时需要自动探测。
|
||||
const snapshotIsStale = computed(() => {
|
||||
const updatedAt = readExtraString(`${props.account.platform}_usage_updated_at`)
|
||||
if (!updatedAt) return true
|
||||
const ts = new Date(updatedAt).getTime()
|
||||
return Number.isNaN(ts) || Date.now() - ts > SNAPSHOT_STALE_MS
|
||||
})
|
||||
|
||||
// 挂载时:先用持久化快照渲染;快照缺失或过期再自动探测一次(失败显示错误,
|
||||
// 避免静默失败导致单元格空白无提示)。
|
||||
onMounted(() => {
|
||||
if (!visible.value) return
|
||||
data.value = snapshotData.value
|
||||
if (!snapshotIsStale.value) return
|
||||
// 模块级去抖:列表页每行一个实例,翻页/筛选/刷新会重复挂载;同一账号
|
||||
// 短时间内已自动探测过则跳过,避免对上游形成探测风暴。
|
||||
const last = lastAutoProbeAt.get(props.account.id) ?? 0
|
||||
if (Date.now() - last < AUTO_PROBE_DEBOUNCE_MS) return
|
||||
lastAutoProbeAt.set(props.account.id, Date.now())
|
||||
handleProbe()
|
||||
})
|
||||
|
||||
const extractErrorMessage = (e: unknown): string => {
|
||||
const err = e as {
|
||||
message?: string
|
||||
reason?: string
|
||||
response?: { data?: { message?: string; error?: string } }
|
||||
}
|
||||
return (
|
||||
err?.message ||
|
||||
err?.reason ||
|
||||
err?.response?.data?.message ||
|
||||
err?.response?.data?.error ||
|
||||
t('common.error')
|
||||
)
|
||||
}
|
||||
|
||||
const truncatedError = computed(() => {
|
||||
if (!error.value) return ''
|
||||
return error.value.length > 80 ? `${error.value.slice(0, 80)}...` : error.value
|
||||
})
|
||||
|
||||
const windowLabel = (window: string) =>
|
||||
window === 'weekly'
|
||||
? t('admin.accounts.cnProviders.windowWeekly')
|
||||
: t('admin.accounts.cnProviders.window5h')
|
||||
|
||||
const utilizationColor = (pct: number) => {
|
||||
if (pct >= 90) return 'bg-red-500'
|
||||
if (pct >= 75) return 'bg-amber-500'
|
||||
return 'bg-emerald-500'
|
||||
}
|
||||
|
||||
const utilizationTextColor = (pct: number) => {
|
||||
if (pct >= 90) return 'text-red-600 dark:text-red-400'
|
||||
if (pct >= 75) return 'text-amber-600 dark:text-amber-400'
|
||||
return 'text-emerald-600 dark:text-emerald-400'
|
||||
}
|
||||
|
||||
// 重置时间相对/绝对简短显示。
|
||||
const formatReset = (iso: string) => {
|
||||
const d = new Date(iso)
|
||||
if (isNaN(d.getTime())) return iso
|
||||
const now = Date.now()
|
||||
const diffMs = d.getTime() - now
|
||||
if (diffMs <= 0) return t('admin.accounts.cnProviders.resetSoon')
|
||||
if (diffMs < 3_600_000) return `${Math.max(1, Math.round(diffMs / 60_000))}m`
|
||||
const hours = Math.round(diffMs / 3_600_000)
|
||||
if (hours < 48) return `${hours}h`
|
||||
const mm = String(d.getMonth() + 1).padStart(2, '0')
|
||||
const dd = String(d.getDate()).padStart(2, '0')
|
||||
return `${mm}-${dd}`
|
||||
}
|
||||
|
||||
const handleProbe = async () => {
|
||||
if (loading.value) return
|
||||
loading.value = true
|
||||
error.value = null
|
||||
try {
|
||||
const result = await adminAPI.cnProviders.queryQuota(props.account.id)
|
||||
// 失败时保留已渲染的快照条形图(仅显示错误行),成功才覆盖。
|
||||
if (result.success) {
|
||||
data.value = result
|
||||
} else {
|
||||
error.value = result.error || t('common.error')
|
||||
}
|
||||
} catch (e) {
|
||||
error.value = extractErrorMessage(e)
|
||||
} finally {
|
||||
loading.value = false
|
||||
}
|
||||
}
|
||||
|
||||
watch(
|
||||
() => props.account.id,
|
||||
() => {
|
||||
data.value = null
|
||||
error.value = null
|
||||
loading.value = false
|
||||
}
|
||||
)
|
||||
</script>
|
||||
@@ -0,0 +1,53 @@
|
||||
<template>
|
||||
<div class="flex flex-wrap gap-2">
|
||||
<button
|
||||
v-for="preset in presets"
|
||||
:key="preset.mode + ':' + preset.protocol + ':' + preset.url"
|
||||
type="button"
|
||||
data-testid="cn-base-url-preset"
|
||||
:class="[
|
||||
'rounded-lg px-3 py-1 text-xs transition-colors',
|
||||
isActive(preset)
|
||||
? 'bg-primary-100 text-primary-700 dark:bg-primary-900/30 dark:text-primary-300'
|
||||
: 'bg-gray-100 text-gray-700 hover:bg-primary-50 hover:text-primary-700 dark:bg-dark-600 dark:text-gray-300 dark:hover:bg-primary-900/30 dark:hover:text-primary-400'
|
||||
]"
|
||||
@click="emit('select', preset)"
|
||||
>
|
||||
{{ preset.label }} ({{ displayUrl(preset.url) }})
|
||||
</button>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { computed } from 'vue'
|
||||
import { CN_BASE_URL_PRESETS, type CnBaseUrlPreset } from './credentialsBuilder'
|
||||
|
||||
// 国产供应商快捷端点:点击把预设地址(及对应账号类型/协议)回填到调用方。
|
||||
// 与 Grok 预设一致,仅作快速填充,输入框仍接受任意第三方转发地址。
|
||||
// 传入 protocol 时只显示该协议档的预设(协议 × 账号类型正交分档)。
|
||||
const props = defineProps<{
|
||||
platform: 'kimi' | 'zhipu' | 'deepseek'
|
||||
/** 当前已选账号类型,用于过滤和高亮匹配的预设 */
|
||||
mode?: 'payg' | 'coding'
|
||||
/** 当前已选 API 协议,用于过滤和高亮匹配的预设 */
|
||||
protocol?: 'chat_completions' | 'anthropic' | 'responses'
|
||||
/** 当前输入框中的 base url,用于高亮完全匹配项 */
|
||||
currentUrl?: string
|
||||
}>()
|
||||
|
||||
const emit = defineEmits<{
|
||||
(e: 'select', preset: CnBaseUrlPreset): void
|
||||
}>()
|
||||
|
||||
const presets = computed(() => {
|
||||
const all = CN_BASE_URL_PRESETS[props.platform] ?? []
|
||||
// 只按协议过滤:同协议下 payg/coding 两档都展示,点击即同时切换账号类型。
|
||||
if (props.protocol == null) return all
|
||||
return all.filter(p => p.protocol === props.protocol)
|
||||
})
|
||||
|
||||
const isActive = (preset: CnBaseUrlPreset) =>
|
||||
(props.mode != null && preset.mode === props.mode) || preset.url === props.currentUrl
|
||||
|
||||
const displayUrl = (url: string) => url.replace(/^https?:\/\//i, '')
|
||||
</script>
|
||||
@@ -161,6 +161,48 @@
|
||||
Grok
|
||||
</button>
|
||||
</div>
|
||||
<!-- CN providers row: Kimi / Zhipu GLM / DeepSeek -->
|
||||
<div class="mt-2 flex flex-wrap rounded-lg bg-gray-100 p-1 dark:bg-dark-700">
|
||||
<button
|
||||
type="button"
|
||||
@click="selectCNPlatform('kimi')"
|
||||
:class="[
|
||||
'flex flex-1 items-center justify-center gap-2 rounded-md px-4 py-2.5 text-sm font-medium transition-all',
|
||||
form.platform === 'kimi'
|
||||
? 'bg-white text-pink-600 shadow-sm dark:bg-dark-600 dark:text-pink-400'
|
||||
: 'text-gray-600 hover:text-gray-900 dark:text-gray-400 dark:hover:text-gray-200'
|
||||
]"
|
||||
>
|
||||
<PlatformIcon platform="kimi" size="sm" />
|
||||
Kimi
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
@click="selectCNPlatform('zhipu')"
|
||||
:class="[
|
||||
'flex flex-1 items-center justify-center gap-2 rounded-md px-4 py-2.5 text-sm font-medium transition-all',
|
||||
form.platform === 'zhipu'
|
||||
? 'bg-white text-indigo-600 shadow-sm dark:bg-dark-600 dark:text-indigo-400'
|
||||
: 'text-gray-600 hover:text-gray-900 dark:text-gray-400 dark:hover:text-gray-200'
|
||||
]"
|
||||
>
|
||||
<PlatformIcon platform="zhipu" size="sm" />
|
||||
Zhipu GLM
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
@click="selectCNPlatform('deepseek')"
|
||||
:class="[
|
||||
'flex flex-1 items-center justify-center gap-2 rounded-md px-4 py-2.5 text-sm font-medium transition-all',
|
||||
form.platform === 'deepseek'
|
||||
? 'bg-white text-teal-600 shadow-sm dark:bg-dark-600 dark:text-teal-400'
|
||||
: 'text-gray-600 hover:text-gray-900 dark:text-gray-400 dark:hover:text-gray-200'
|
||||
]"
|
||||
>
|
||||
<PlatformIcon platform="deepseek" size="sm" />
|
||||
DeepSeek
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Account Type Selection (Anthropic) -->
|
||||
@@ -411,6 +453,100 @@
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Account Mode Selection (Kimi / Zhipu / DeepSeek) -->
|
||||
<div v-if="isCNPlatform">
|
||||
<label class="input-label">{{ t('admin.accounts.cnProviders.accountMode.title') }}</label>
|
||||
<div class="mt-2 grid grid-cols-1 gap-3 sm:grid-cols-2" data-tour="account-form-mode">
|
||||
<!-- Pay-as-you-go (token balance) -->
|
||||
<button
|
||||
type="button"
|
||||
@click="accountMode = 'payg'"
|
||||
:class="[
|
||||
'flex items-center gap-3 rounded-lg border-2 p-3 text-left transition-all',
|
||||
accountMode === 'payg'
|
||||
? cnAccentActiveClass
|
||||
: 'border-gray-200 hover:border-gray-400 dark:border-dark-600 dark:hover:border-gray-600'
|
||||
]"
|
||||
>
|
||||
<div
|
||||
:class="[
|
||||
'flex h-8 w-8 shrink-0 items-center justify-center rounded-lg',
|
||||
accountMode === 'payg'
|
||||
? cnAccentIconClass
|
||||
: 'bg-gray-100 text-gray-500 dark:bg-dark-600 dark:text-gray-400'
|
||||
]"
|
||||
>
|
||||
<Icon name="creditCard" size="sm" />
|
||||
</div>
|
||||
<div>
|
||||
<span class="block text-sm font-medium text-gray-900 dark:text-white">{{ t('admin.accounts.cnProviders.accountMode.payg') }}</span>
|
||||
<span class="text-xs text-gray-500 dark:text-gray-400">{{ t('admin.accounts.cnProviders.accountMode.paygDesc') }}</span>
|
||||
</div>
|
||||
</button>
|
||||
<!-- Coding Plan (kimi / zhipu only — DeepSeek has no coding plan) -->
|
||||
<button
|
||||
v-if="form.platform !== 'deepseek'"
|
||||
type="button"
|
||||
@click="accountMode = 'coding'"
|
||||
:class="[
|
||||
'flex items-center gap-3 rounded-lg border-2 p-3 text-left transition-all',
|
||||
accountMode === 'coding'
|
||||
? cnAccentActiveClass
|
||||
: 'border-gray-200 hover:border-gray-400 dark:border-dark-600 dark:hover:border-gray-600'
|
||||
]"
|
||||
>
|
||||
<div
|
||||
:class="[
|
||||
'flex h-8 w-8 shrink-0 items-center justify-center rounded-lg',
|
||||
accountMode === 'coding'
|
||||
? cnAccentIconClass
|
||||
: 'bg-gray-100 text-gray-500 dark:bg-dark-600 dark:text-gray-400'
|
||||
]"
|
||||
>
|
||||
<Icon name="bolt" size="sm" />
|
||||
</div>
|
||||
<div>
|
||||
<span class="block text-sm font-medium text-gray-900 dark:text-white">{{ t('admin.accounts.cnProviders.accountMode.coding') }}</span>
|
||||
<span class="text-xs text-gray-500 dark:text-gray-400">{{ t('admin.accounts.cnProviders.accountMode.codingDesc') }}</span>
|
||||
</div>
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- API Protocol Selection (Kimi / Zhipu / DeepSeek) -->
|
||||
<div v-if="isCNPlatform" class="mt-4">
|
||||
<label class="input-label">{{ t('admin.accounts.cnProviders.apiProtocol.title') }}</label>
|
||||
<div class="mt-2 grid grid-cols-1 gap-3 sm:grid-cols-3">
|
||||
<button
|
||||
v-for="opt in cnProtocolOptions"
|
||||
:key="opt.value"
|
||||
type="button"
|
||||
@click="apiProtocol = opt.value"
|
||||
:class="[
|
||||
'flex items-center gap-3 rounded-lg border-2 p-3 text-left transition-all',
|
||||
apiProtocol === opt.value
|
||||
? cnAccentActiveClass
|
||||
: 'border-gray-200 hover:border-gray-400 dark:border-dark-600 dark:hover:border-gray-600'
|
||||
]"
|
||||
>
|
||||
<div
|
||||
:class="[
|
||||
'flex h-8 w-8 shrink-0 items-center justify-center rounded-lg',
|
||||
apiProtocol === opt.value
|
||||
? cnAccentIconClass
|
||||
: 'bg-gray-100 text-gray-500 dark:bg-dark-600 dark:text-gray-400'
|
||||
]"
|
||||
>
|
||||
<Icon :name="opt.value === 'anthropic' ? 'sparkles' : opt.value === 'responses' ? 'terminal' : 'chat'" size="sm" />
|
||||
</div>
|
||||
<div>
|
||||
<span class="block text-sm font-medium text-gray-900 dark:text-white">{{ t(`admin.accounts.cnProviders.apiProtocol.${opt.labelKey}`) }}</span>
|
||||
<span class="text-xs text-gray-500 dark:text-gray-400">{{ t(`admin.accounts.cnProviders.apiProtocol.${opt.labelKey}Desc`) }}</span>
|
||||
</div>
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Account Type Selection (Gemini) -->
|
||||
<div v-if="form.platform === 'gemini'">
|
||||
<div class="flex items-center justify-between">
|
||||
@@ -1120,15 +1256,7 @@
|
||||
v-model="apiKeyBaseUrl"
|
||||
type="text"
|
||||
class="input"
|
||||
:placeholder="
|
||||
form.platform === 'openai'
|
||||
? 'https://api.openai.com'
|
||||
: form.platform === 'gemini'
|
||||
? 'https://generativelanguage.googleapis.com'
|
||||
: form.platform === 'grok'
|
||||
? 'https://api.x.ai/v1'
|
||||
: 'https://api.anthropic.com'
|
||||
"
|
||||
:placeholder="apiKeyBaseUrlPlaceholder"
|
||||
/>
|
||||
<p v-if="baseUrlHint" class="input-hint">{{ baseUrlHint }}</p>
|
||||
<GrokBaseUrlPresets
|
||||
@@ -1136,6 +1264,15 @@
|
||||
class="mt-2"
|
||||
@select="apiKeyBaseUrl = $event"
|
||||
/>
|
||||
<CnBaseUrlPresets
|
||||
v-if="isCNPlatform"
|
||||
class="mt-2"
|
||||
:platform="cnPresetPlatform"
|
||||
:mode="accountMode"
|
||||
:protocol="apiProtocol"
|
||||
:current-url="apiKeyBaseUrl"
|
||||
@select="onCnPresetSelect"
|
||||
/>
|
||||
</div>
|
||||
<div>
|
||||
<label class="input-label">{{ t('admin.accounts.apiKeyRequired') }}</label>
|
||||
@@ -1144,15 +1281,7 @@
|
||||
type="password"
|
||||
required
|
||||
class="input font-mono"
|
||||
:placeholder="
|
||||
form.platform === 'openai'
|
||||
? 'sk-proj-...'
|
||||
: form.platform === 'gemini'
|
||||
? 'AIza...'
|
||||
: form.platform === 'grok'
|
||||
? 'xai-...'
|
||||
: 'sk-ant-...'
|
||||
"
|
||||
:placeholder="apiKeyValuePlaceholder"
|
||||
/>
|
||||
<p v-if="apiKeyHint" class="input-hint">{{ apiKeyHint }}</p>
|
||||
</div>
|
||||
@@ -3609,13 +3738,17 @@ import ModelWhitelistSelector from '@/components/account/ModelWhitelistSelector.
|
||||
import QuotaLimitCard from '@/components/account/QuotaLimitCard.vue'
|
||||
import Toggle from '@/components/common/Toggle.vue'
|
||||
import GrokBaseUrlPresets from '@/components/account/GrokBaseUrlPresets.vue'
|
||||
import CnBaseUrlPresets from '@/components/account/CnBaseUrlPresets.vue'
|
||||
import HeaderOverrideEditor from '@/components/account/HeaderOverrideEditor.vue'
|
||||
import {
|
||||
applyAntigravityProjectID,
|
||||
applyHeaderOverride,
|
||||
applyInterceptWarmup,
|
||||
defaultCNBaseUrl,
|
||||
isHeaderOverrideCapable,
|
||||
validateHeaderOverrideRows,
|
||||
type CnAccountMode,
|
||||
type CnApiProtocol,
|
||||
type HeaderOverrideRow
|
||||
} from '@/components/account/credentialsBuilder'
|
||||
import { formatDateTimeLocalInput, parseDateTimeLocalInput } from '@/utils/format'
|
||||
@@ -3674,6 +3807,42 @@ const apiKeyHint = computed(() => {
|
||||
return t('admin.accounts.apiKeyHint')
|
||||
})
|
||||
|
||||
// Base URL / API Key 占位符:国产供应商随账号类型变化。
|
||||
const apiKeyBaseUrlPlaceholder = computed(() => {
|
||||
if (isCNPlatform.value) {
|
||||
return defaultCNBaseUrl(form.platform, accountMode.value, apiProtocol.value) || 'https://api.example.com'
|
||||
}
|
||||
switch (form.platform) {
|
||||
case 'openai':
|
||||
return 'https://api.openai.com'
|
||||
case 'gemini':
|
||||
return 'https://generativelanguage.googleapis.com'
|
||||
case 'grok':
|
||||
return 'https://api.x.ai/v1'
|
||||
default:
|
||||
return 'https://api.anthropic.com'
|
||||
}
|
||||
})
|
||||
|
||||
const apiKeyValuePlaceholder = computed(() => {
|
||||
switch (form.platform) {
|
||||
case 'openai':
|
||||
return 'sk-proj-...'
|
||||
case 'gemini':
|
||||
return 'AIza...'
|
||||
case 'grok':
|
||||
return 'xai-...'
|
||||
case 'kimi':
|
||||
return 'sk-...'
|
||||
case 'zhipu':
|
||||
return '<api-key>.<secret>'
|
||||
case 'deepseek':
|
||||
return 'sk-...'
|
||||
default:
|
||||
return 'sk-ant-...'
|
||||
}
|
||||
})
|
||||
|
||||
interface Props {
|
||||
show: boolean
|
||||
proxies: Proxy[]
|
||||
@@ -3753,6 +3922,86 @@ const apiKeyBaseUrl = ref('https://api.anthropic.com')
|
||||
const apiKeyValue = ref('')
|
||||
const upstreamBillingAutoProbeEnabled = ref(true)
|
||||
|
||||
// ── 国产供应商(Kimi / Zhipu / DeepSeek)账号类型、API 协议与端点 ──
|
||||
const accountMode = ref<CnAccountMode>('payg')
|
||||
// API 协议决定转发端点与格式:cc=现有转换链,anthropic=原生直通(Claude Code),
|
||||
// responses=deepseek 原生 Responses 端点(Codex)。与账号类型正交。
|
||||
const apiProtocol = ref<CnApiProtocol>('chat_completions')
|
||||
const isCNPlatform = computed(
|
||||
() => form.platform === 'kimi' || form.platform === 'zhipu' || form.platform === 'deepseek'
|
||||
)
|
||||
// CnBaseUrlPresets 的 platform prop 是平台字面量联合类型,模板里不能写
|
||||
// `as` 断言(其中的 `|` 会被 eslint 误判为 Vue2 filter 语法),经此 computed 传递。
|
||||
const cnPresetPlatform = computed<'kimi' | 'zhipu' | 'deepseek'>(() => {
|
||||
if (form.platform === 'kimi' || form.platform === 'zhipu' || form.platform === 'deepseek') {
|
||||
return form.platform
|
||||
}
|
||||
return 'kimi'
|
||||
})
|
||||
// 当前平台可选的协议档(responses 仅 deepseek)。
|
||||
const cnProtocolOptions = computed<Array<{ value: CnApiProtocol; labelKey: string }>>(() => {
|
||||
const opts: Array<{ value: CnApiProtocol; labelKey: string }> = [
|
||||
{ value: 'chat_completions', labelKey: 'chatCompletions' },
|
||||
{ value: 'anthropic', labelKey: 'anthropic' }
|
||||
]
|
||||
if (form.platform === 'deepseek') {
|
||||
opts.push({ value: 'responses', labelKey: 'responses' })
|
||||
}
|
||||
return opts
|
||||
})
|
||||
// 当前选中平台的品牌色(选中卡片描边 / 图标底色),与 platformColors 取色一致。
|
||||
const cnAccentActiveClass = computed(() => {
|
||||
switch (form.platform) {
|
||||
case 'kimi':
|
||||
return 'border-pink-500 bg-pink-50 dark:bg-pink-900/20'
|
||||
case 'zhipu':
|
||||
return 'border-indigo-500 bg-indigo-50 dark:bg-indigo-900/20'
|
||||
case 'deepseek':
|
||||
return 'border-teal-500 bg-teal-50 dark:bg-teal-900/20'
|
||||
default:
|
||||
return 'border-primary-500 bg-primary-50 dark:bg-primary-900/20'
|
||||
}
|
||||
})
|
||||
const cnAccentIconClass = computed(() => {
|
||||
switch (form.platform) {
|
||||
case 'kimi':
|
||||
return 'bg-pink-500 text-white'
|
||||
case 'zhipu':
|
||||
return 'bg-indigo-500 text-white'
|
||||
case 'deepseek':
|
||||
return 'bg-teal-500 text-white'
|
||||
default:
|
||||
return 'bg-primary-500 text-white'
|
||||
}
|
||||
})
|
||||
// 切换国产供应商平台:强制 apikey 类型,deepseek 无 coding 套餐故锁定 payg,
|
||||
// 协议回落 chat_completions,并把 base url 重置为该平台默认端点。
|
||||
function selectCNPlatform(platform: 'kimi' | 'zhipu' | 'deepseek') {
|
||||
form.platform = platform
|
||||
form.type = 'apikey'
|
||||
accountCategory.value = 'apikey'
|
||||
apiProtocol.value = 'chat_completions'
|
||||
if (platform === 'deepseek') {
|
||||
accountMode.value = 'payg'
|
||||
}
|
||||
apiKeyBaseUrl.value = defaultCNBaseUrl(platform, accountMode.value, apiProtocol.value)
|
||||
}
|
||||
// 账号类型 / 协议变更时同步默认 base url。
|
||||
watch(accountMode, (mode) => {
|
||||
if (!isCNPlatform.value) return
|
||||
apiKeyBaseUrl.value = defaultCNBaseUrl(form.platform, mode, apiProtocol.value)
|
||||
})
|
||||
watch(apiProtocol, (protocol) => {
|
||||
if (!isCNPlatform.value) return
|
||||
apiKeyBaseUrl.value = defaultCNBaseUrl(form.platform, accountMode.value, protocol)
|
||||
})
|
||||
// 点击预设端点:同时回填 base url、账号类型与协议。
|
||||
function onCnPresetSelect(preset: { mode: CnAccountMode; protocol: CnApiProtocol; url: string }) {
|
||||
accountMode.value = preset.mode
|
||||
apiProtocol.value = preset.protocol
|
||||
apiKeyBaseUrl.value = preset.url
|
||||
}
|
||||
|
||||
const syncPreviewCredentials = computed(() => {
|
||||
if (!apiKeyValue.value) return undefined
|
||||
return {
|
||||
@@ -4248,14 +4497,18 @@ watch(
|
||||
() => form.platform,
|
||||
(newPlatform) => {
|
||||
// Reset base URL based on platform
|
||||
apiKeyBaseUrl.value =
|
||||
(newPlatform === 'openai')
|
||||
? 'https://api.openai.com'
|
||||
: newPlatform === 'gemini'
|
||||
? 'https://generativelanguage.googleapis.com'
|
||||
: newPlatform === 'grok'
|
||||
? 'https://api.x.ai/v1'
|
||||
: 'https://api.anthropic.com'
|
||||
if (newPlatform === 'kimi' || newPlatform === 'zhipu' || newPlatform === 'deepseek') {
|
||||
apiKeyBaseUrl.value = defaultCNBaseUrl(newPlatform, accountMode.value, apiProtocol.value)
|
||||
} else {
|
||||
apiKeyBaseUrl.value =
|
||||
(newPlatform === 'openai')
|
||||
? 'https://api.openai.com'
|
||||
: newPlatform === 'gemini'
|
||||
? 'https://generativelanguage.googleapis.com'
|
||||
: newPlatform === 'grok'
|
||||
? 'https://api.x.ai/v1'
|
||||
: 'https://api.anthropic.com'
|
||||
}
|
||||
// Clear model-related settings
|
||||
allowedModels.value = []
|
||||
modelMappings.value = []
|
||||
@@ -4696,6 +4949,8 @@ const resetForm = () => {
|
||||
form.expires_at = null
|
||||
accountCategory.value = 'oauth-based'
|
||||
addMethod.value = 'oauth'
|
||||
accountMode.value = 'payg'
|
||||
apiProtocol.value = 'chat_completions'
|
||||
apiKeyBaseUrl.value = 'https://api.anthropic.com'
|
||||
apiKeyValue.value = ''
|
||||
upstreamBillingAutoProbeEnabled.value = true
|
||||
@@ -5158,6 +5413,20 @@ const handleSubmit = async () => {
|
||||
credentials.tier_id = geminiTierAIStudio.value
|
||||
}
|
||||
|
||||
// 国产供应商:账号模式 + 协议 + 对应端点写入凭据;后端按 account_mode 路由
|
||||
// 额度/余额探测,按 api_protocol 路由转发端点与格式。注意 CN apikey 走本函数
|
||||
// 的通用路径(直接 doCreateAccount),不经过 createAccountAndFinish。
|
||||
if (form.platform === 'kimi' || form.platform === 'zhipu' || form.platform === 'deepseek') {
|
||||
credentials.account_mode = accountMode.value
|
||||
credentials.api_protocol = apiProtocol.value
|
||||
const resolvedCNBase = (
|
||||
apiKeyBaseUrl.value.trim() || defaultCNBaseUrl(form.platform, accountMode.value, apiProtocol.value)
|
||||
).trim()
|
||||
if (resolvedCNBase) {
|
||||
credentials.base_url = resolvedCNBase
|
||||
}
|
||||
}
|
||||
|
||||
// Add model mapping if configured(OpenAI 开启自动透传时不应用)
|
||||
if (!isOpenAIModelRestrictionDisabled.value) {
|
||||
const modelMapping = buildModelMappingObject(modelRestrictionMode.value, allowedModels.value, modelMappings.value)
|
||||
|
||||
@@ -52,6 +52,57 @@
|
||||
class="mt-2"
|
||||
@select="editBaseUrl = $event"
|
||||
/>
|
||||
<CnBaseUrlPresets
|
||||
v-if="isCNApiKeyAccount"
|
||||
class="mt-2"
|
||||
:platform="cnPresetPlatform"
|
||||
:mode="editAccountMode"
|
||||
:protocol="editApiProtocol"
|
||||
:current-url="editBaseUrl"
|
||||
@select="onCnPresetSelect"
|
||||
/>
|
||||
</div>
|
||||
<!-- Account Mode Selection (CN providers) -->
|
||||
<div v-if="isCNApiKeyAccount">
|
||||
<label class="input-label">{{ t('admin.accounts.cnProviders.accountMode.title') }}</label>
|
||||
<div class="mt-2 flex flex-wrap gap-2">
|
||||
<button
|
||||
v-for="opt in cnAccountModeOptions"
|
||||
:key="opt.value"
|
||||
type="button"
|
||||
:class="[
|
||||
'rounded-lg border-2 px-3 py-1.5 text-xs transition-all',
|
||||
editAccountMode === opt.value
|
||||
? 'border-primary-500 bg-primary-50 font-medium text-primary-700 dark:bg-primary-900/30 dark:text-primary-300'
|
||||
: 'border-gray-200 text-gray-700 hover:border-gray-400 dark:border-dark-600 dark:text-gray-300 dark:hover:border-gray-600'
|
||||
]"
|
||||
@click="editAccountMode = opt.value"
|
||||
>
|
||||
{{ t(`admin.accounts.cnProviders.accountMode.${opt.labelKey}`) }}
|
||||
</button>
|
||||
</div>
|
||||
<p class="input-hint">{{ t(`admin.accounts.cnProviders.accountMode.${editAccountMode}Desc`) }}</p>
|
||||
</div>
|
||||
<!-- API Protocol Selection (CN providers) -->
|
||||
<div v-if="isCNApiKeyAccount">
|
||||
<label class="input-label">{{ t('admin.accounts.cnProviders.apiProtocol.title') }}</label>
|
||||
<div class="mt-2 flex flex-wrap gap-2">
|
||||
<button
|
||||
v-for="opt in cnProtocolOptions"
|
||||
:key="opt.value"
|
||||
type="button"
|
||||
:class="[
|
||||
'rounded-lg border-2 px-3 py-1.5 text-xs transition-all',
|
||||
editApiProtocol === opt.value
|
||||
? 'border-primary-500 bg-primary-50 font-medium text-primary-700 dark:bg-primary-900/30 dark:text-primary-300'
|
||||
: 'border-gray-200 text-gray-700 hover:border-gray-400 dark:border-dark-600 dark:text-gray-300 dark:hover:border-gray-600'
|
||||
]"
|
||||
@click="editApiProtocol = opt.value"
|
||||
>
|
||||
{{ t(`admin.accounts.cnProviders.apiProtocol.${opt.labelKey}`) }}
|
||||
</button>
|
||||
</div>
|
||||
<p class="input-hint">{{ t(`admin.accounts.cnProviders.apiProtocol.${cnProtocolDescKey}Desc`) }}</p>
|
||||
</div>
|
||||
<div>
|
||||
<label class="input-label">{{ t('admin.accounts.apiKey') }}</label>
|
||||
@@ -2710,7 +2761,7 @@
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { ref, reactive, computed, watch } from 'vue'
|
||||
import { ref, reactive, computed, watch, nextTick } from 'vue'
|
||||
import { useI18n } from 'vue-i18n'
|
||||
import { useAppStore } from '@/stores/app'
|
||||
import { useAuthStore } from '@/stores/auth'
|
||||
@@ -2737,6 +2788,7 @@ import GroupSelector from '@/components/common/GroupSelector.vue'
|
||||
import ModelWhitelistSelector from '@/components/account/ModelWhitelistSelector.vue'
|
||||
import QuotaLimitCard from '@/components/account/QuotaLimitCard.vue'
|
||||
import GrokBaseUrlPresets from '@/components/account/GrokBaseUrlPresets.vue'
|
||||
import CnBaseUrlPresets from '@/components/account/CnBaseUrlPresets.vue'
|
||||
import HeaderOverrideEditor from '@/components/account/HeaderOverrideEditor.vue'
|
||||
import OllamaCloudUsageSettings from '@/components/account/OllamaCloudUsageSettings.vue'
|
||||
import {
|
||||
@@ -2750,8 +2802,11 @@ import {
|
||||
isHeaderOverrideCapable,
|
||||
splitHeaderOverridesObject,
|
||||
validateHeaderOverrideRows,
|
||||
defaultCNBaseUrl,
|
||||
HEADER_OVERRIDE_ENABLED_CREDENTIAL_KEY,
|
||||
HEADER_OVERRIDES_CREDENTIAL_KEY,
|
||||
type CnAccountMode,
|
||||
type CnApiProtocol,
|
||||
type HeaderOverrideRow
|
||||
} from '@/components/account/credentialsBuilder'
|
||||
import { formatDateTime, formatDateTimeLocalInput, parseDateTimeLocalInput } from '@/utils/format'
|
||||
@@ -2829,6 +2884,78 @@ interface TempUnschedRuleForm {
|
||||
const submitting = ref(false)
|
||||
const editBaseUrl = ref('https://api.anthropic.com')
|
||||
const editApiKey = ref('')
|
||||
|
||||
// ── 国产供应商(Kimi / Zhipu / DeepSeek)account_mode / api_protocol 编辑 ──
|
||||
// account_mode 决定额度/余额监控路径,api_protocol 决定转发端点与格式;
|
||||
// 二者均可修正(早期创建的账号可能存错默认值),切换时重置 base_url 预置。
|
||||
const isCNApiKeyAccount = computed(
|
||||
() =>
|
||||
props.account?.type === 'apikey' &&
|
||||
(props.account.platform === 'kimi' ||
|
||||
props.account.platform === 'zhipu' ||
|
||||
props.account.platform === 'deepseek')
|
||||
)
|
||||
// CnBaseUrlPresets 的 platform prop 是平台字面量联合类型,模板里不能写
|
||||
// `as` 断言(其中的 `|` 会被 eslint 误判为 Vue2 filter 语法),经此 computed 传递。
|
||||
const cnPresetPlatform = computed<'kimi' | 'zhipu' | 'deepseek'>(() => {
|
||||
const platform = props.account?.platform
|
||||
if (platform === 'kimi' || platform === 'zhipu' || platform === 'deepseek') {
|
||||
return platform
|
||||
}
|
||||
return 'kimi'
|
||||
})
|
||||
const editApiProtocol = ref<CnApiProtocol>('chat_completions')
|
||||
const editAccountMode = ref<CnAccountMode>('payg')
|
||||
// 回填窗口标志:syncFormFromAccount 会同步改写 editAccountMode / editApiProtocol,
|
||||
// 而 watcher(pre-flush)在同步代码执行完之后才触发——若不抑制,会把刚恢复的
|
||||
// 存储版 base_url(可能是用户自定义/中转地址)覆盖为官方预设并在下次保存时持久化。
|
||||
// nextTick 后解除,此后用户主动切换模式/协议仍正常联动重置。
|
||||
const syncingForm = ref(false)
|
||||
const cnAccountModeOptions = computed<Array<{ value: CnAccountMode; labelKey: 'payg' | 'coding' }>>(
|
||||
() => {
|
||||
// DeepSeek 无 coding 套餐(与创建弹窗一致),仅保留按量付费。
|
||||
if (props.account?.platform === 'deepseek') {
|
||||
return [{ value: 'payg', labelKey: 'payg' }]
|
||||
}
|
||||
return [
|
||||
{ value: 'payg', labelKey: 'payg' },
|
||||
{ value: 'coding', labelKey: 'coding' }
|
||||
]
|
||||
}
|
||||
)
|
||||
const cnProtocolOptions = computed<Array<{ value: CnApiProtocol; labelKey: string }>>(() => {
|
||||
const opts: Array<{ value: CnApiProtocol; labelKey: string }> = [
|
||||
{ value: 'chat_completions', labelKey: 'chatCompletions' },
|
||||
{ value: 'anthropic', labelKey: 'anthropic' }
|
||||
]
|
||||
if (props.account?.platform === 'deepseek') {
|
||||
opts.push({ value: 'responses', labelKey: 'responses' })
|
||||
}
|
||||
return opts
|
||||
})
|
||||
watch(editApiProtocol, (protocol) => {
|
||||
if (!isCNApiKeyAccount.value || syncingForm.value) return
|
||||
editBaseUrl.value = defaultCNBaseUrl(props.account!.platform, editAccountMode.value, protocol)
|
||||
})
|
||||
watch(editAccountMode, (mode) => {
|
||||
if (!isCNApiKeyAccount.value || syncingForm.value) return
|
||||
// deepseek 无 coding 套餐:防御性回退(UI 已隐藏该选项)。
|
||||
const effectiveMode = props.account!.platform === 'deepseek' && mode === 'coding' ? 'payg' : mode
|
||||
if (effectiveMode !== mode) {
|
||||
editAccountMode.value = effectiveMode
|
||||
return
|
||||
}
|
||||
editBaseUrl.value = defaultCNBaseUrl(props.account!.platform, mode, editApiProtocol.value)
|
||||
})
|
||||
const cnProtocolDescKey = computed(
|
||||
() => cnProtocolOptions.value.find(o => o.value === editApiProtocol.value)?.labelKey ?? 'chatCompletions'
|
||||
)
|
||||
// 点击预设端点:回填 base url 与对应模式/协议。
|
||||
function onCnPresetSelect(preset: { mode: CnAccountMode; protocol: CnApiProtocol; url: string }) {
|
||||
editAccountMode.value = preset.mode
|
||||
editApiProtocol.value = preset.protocol
|
||||
editBaseUrl.value = preset.url
|
||||
}
|
||||
// Bedrock credentials
|
||||
const editBedrockAccessKeyId = ref('')
|
||||
const editBedrockSecretAccessKey = ref('')
|
||||
@@ -3272,6 +3399,15 @@ const defaultBaseUrl = computed(() => {
|
||||
if (props.account?.platform === 'openai') return 'https://api.openai.com'
|
||||
if (props.account?.platform === 'gemini') return 'https://generativelanguage.googleapis.com'
|
||||
if (props.account?.platform === 'grok') return 'https://api.x.ai/v1'
|
||||
// CN 供应商:按当前模式/协议回落到官方预设(清空输入框提交时使用),
|
||||
// 不能落到 anthropic 默认值(会被当 CC base 拼出错误端点)。
|
||||
if (
|
||||
props.account?.platform === 'kimi' ||
|
||||
props.account?.platform === 'zhipu' ||
|
||||
props.account?.platform === 'deepseek'
|
||||
) {
|
||||
return defaultCNBaseUrl(props.account.platform, editAccountMode.value, editApiProtocol.value)
|
||||
}
|
||||
return 'https://api.anthropic.com'
|
||||
})
|
||||
|
||||
@@ -3381,6 +3517,11 @@ const syncFormFromAccount = (newAccount: Account | null) => {
|
||||
if (!newAccount) {
|
||||
return
|
||||
}
|
||||
// 进入回填窗口:抑制 CN 模式/协议 watcher 联动重置 base_url(见 syncingForm 注释)。
|
||||
syncingForm.value = true
|
||||
void nextTick(() => {
|
||||
syncingForm.value = false
|
||||
})
|
||||
antigravityMixedChannelConfirmed.value = false
|
||||
showMixedChannelWarning.value = false
|
||||
mixedChannelWarningDetails.value = null
|
||||
@@ -3622,6 +3763,17 @@ const syncFormFromAccount = (newAccount: Account | null) => {
|
||||
// Initialize API Key fields for apikey type
|
||||
if (newAccount.type === 'apikey' && newAccount.credentials) {
|
||||
const credentials = newAccount.credentials as Record<string, unknown>
|
||||
// 国产供应商:读取 account_mode 与 api_protocol 作为可编辑初始值
|
||||
// (编辑弹窗允许修正两者,用于修复早期存错默认值的账号)。
|
||||
if (newAccount.platform === 'kimi' || newAccount.platform === 'zhipu' || newAccount.platform === 'deepseek') {
|
||||
editAccountMode.value = credentials.account_mode === 'coding' ? 'coding' : 'payg'
|
||||
const storedProtocol = credentials.api_protocol
|
||||
editApiProtocol.value =
|
||||
storedProtocol === 'anthropic' || storedProtocol === 'responses' ? storedProtocol : 'chat_completions'
|
||||
if (newAccount.platform !== 'deepseek' && editApiProtocol.value === 'responses') {
|
||||
editApiProtocol.value = 'chat_completions'
|
||||
}
|
||||
}
|
||||
const platformDefaultUrl =
|
||||
newAccount.platform === 'openai'
|
||||
? 'https://api.openai.com'
|
||||
@@ -3629,7 +3781,11 @@ const syncFormFromAccount = (newAccount: Account | null) => {
|
||||
? 'https://generativelanguage.googleapis.com'
|
||||
: newAccount.platform === 'grok'
|
||||
? 'https://api.x.ai/v1'
|
||||
: 'https://api.anthropic.com'
|
||||
: newAccount.platform === 'kimi' ||
|
||||
newAccount.platform === 'zhipu' ||
|
||||
newAccount.platform === 'deepseek'
|
||||
? defaultCNBaseUrl(newAccount.platform, editAccountMode.value, editApiProtocol.value)
|
||||
: 'https://api.anthropic.com'
|
||||
editBaseUrl.value = (credentials.base_url as string) || platformDefaultUrl
|
||||
|
||||
// Load model mappings and detect mode
|
||||
@@ -4298,6 +4454,12 @@ const handleSubmit = async () => {
|
||||
base_url: newBaseUrl
|
||||
}
|
||||
|
||||
// 国产供应商:模式与协议写入凭据(决定额度/余额探测与转发端点/格式)。
|
||||
if (isCNApiKeyAccount.value) {
|
||||
newCredentials.account_mode = editAccountMode.value
|
||||
newCredentials.api_protocol = editApiProtocol.value
|
||||
}
|
||||
|
||||
// Handle API key
|
||||
// 后端响应已脱敏:currentCredentials 不会再包含 api_key 原文。
|
||||
// 用户填入新值则覆盖;留空时优先看 credentials_status.has_api_key;
|
||||
|
||||
@@ -199,7 +199,16 @@ const normalizedPlatforms = computed(() => {
|
||||
)
|
||||
})
|
||||
|
||||
const upstreamSyncPlatforms = new Set(['anthropic', 'openai', 'gemini', 'antigravity', 'grok'])
|
||||
const upstreamSyncPlatforms = new Set([
|
||||
'anthropic',
|
||||
'openai',
|
||||
'gemini',
|
||||
'antigravity',
|
||||
'grok',
|
||||
'kimi',
|
||||
'zhipu',
|
||||
'deepseek'
|
||||
])
|
||||
const canSyncUpstream = computed(() => {
|
||||
if (props.accountId) {
|
||||
if (normalizedPlatforms.value.length === 0) return true
|
||||
|
||||
@@ -242,6 +242,91 @@ export const GROK_BASE_URL_PRESETS: GrokBaseUrlPreset[] = [
|
||||
{ label: 'eu-west-1', url: 'https://eu-west-1.api.x.ai/v1' }
|
||||
]
|
||||
|
||||
// ========== 国产供应商(Kimi / Zhipu / DeepSeek)base_url 预设 ==========
|
||||
// 与后端 service/domain_constants.go 的默认 base url 保持一致。
|
||||
// 账号类型(payg 按量付费 / coding 编程套餐)决定额度监控方式;
|
||||
// API 协议(chat_completions / anthropic / responses)决定转发端点与格式,
|
||||
// 两者正交。同协议请求零转换直通,跨协议组合才走转换链。
|
||||
|
||||
export type CnAccountMode = 'payg' | 'coding'
|
||||
|
||||
/** 仅 deepseek 支持 responses 协议(官方原生 /responses 端点,适配 Codex)。 */
|
||||
export type CnApiProtocol = 'chat_completions' | 'anthropic' | 'responses'
|
||||
|
||||
export interface CnBaseUrlPreset {
|
||||
mode: CnAccountMode
|
||||
protocol: CnApiProtocol
|
||||
/** 专有名词,不参与 i18n */
|
||||
label: string
|
||||
url: string
|
||||
}
|
||||
|
||||
/** 各供应商按账号类型 × API 协议分档的快捷端点(点击快速填充,输入框仍可自由填写)。 */
|
||||
export const CN_BASE_URL_PRESETS: Record<'kimi' | 'zhipu' | 'deepseek', CnBaseUrlPreset[]> = {
|
||||
kimi: [
|
||||
{ mode: 'payg', protocol: 'chat_completions', label: 'Moonshot', url: 'https://api.moonshot.cn/v1' },
|
||||
{ mode: 'payg', protocol: 'anthropic', label: 'Moonshot Anthropic', url: 'https://api.moonshot.cn/anthropic' },
|
||||
{ mode: 'coding', protocol: 'chat_completions', label: 'Kimi For Coding', url: 'https://api.kimi.com/coding/v1' },
|
||||
{ mode: 'coding', protocol: 'anthropic', label: 'Kimi Coding Anthropic', url: 'https://api.kimi.com/coding' }
|
||||
],
|
||||
zhipu: [
|
||||
{ mode: 'payg', protocol: 'chat_completions', label: 'GLM PaaS', url: 'https://open.bigmodel.cn/api/paas/v4' },
|
||||
{ mode: 'payg', protocol: 'anthropic', label: 'GLM Anthropic', url: 'https://open.bigmodel.cn/api/anthropic' },
|
||||
{ mode: 'coding', protocol: 'chat_completions', label: 'GLM Coding', url: 'https://open.bigmodel.cn/api/coding/paas/v4' },
|
||||
{ mode: 'coding', protocol: 'anthropic', label: 'GLM Coding Anthropic', url: 'https://open.bigmodel.cn/api/anthropic' }
|
||||
],
|
||||
deepseek: [
|
||||
{ mode: 'payg', protocol: 'chat_completions', label: 'DeepSeek', url: 'https://api.deepseek.com' },
|
||||
{ mode: 'payg', protocol: 'anthropic', label: 'DeepSeek Anthropic', url: 'https://api.deepseek.com/anthropic' },
|
||||
{ mode: 'payg', protocol: 'responses', label: 'DeepSeek Responses', url: 'https://api.deepseek.com' }
|
||||
]
|
||||
}
|
||||
|
||||
/** 返回指定供应商 + 账号类型 + API 协议的默认 base url。 */
|
||||
export function defaultCNBaseUrl(
|
||||
platform: string,
|
||||
mode: CnAccountMode,
|
||||
protocol: CnApiProtocol = 'chat_completions'
|
||||
): string {
|
||||
if (protocol === 'anthropic') {
|
||||
switch (platform) {
|
||||
case 'kimi':
|
||||
return mode === 'coding' ? 'https://api.kimi.com/coding' : 'https://api.moonshot.cn/anthropic'
|
||||
case 'zhipu':
|
||||
return 'https://open.bigmodel.cn/api/anthropic'
|
||||
case 'deepseek':
|
||||
return 'https://api.deepseek.com/anthropic'
|
||||
default:
|
||||
return ''
|
||||
}
|
||||
}
|
||||
// responses 仅 deepseek:base 与 chat_completions 相同(端点路径差异由后端处理)。
|
||||
switch (platform) {
|
||||
case 'kimi':
|
||||
return mode === 'coding' ? 'https://api.kimi.com/coding/v1' : 'https://api.moonshot.cn/v1'
|
||||
case 'zhipu':
|
||||
return mode === 'coding'
|
||||
? 'https://open.bigmodel.cn/api/coding/paas/v4'
|
||||
: 'https://open.bigmodel.cn/api/paas/v4'
|
||||
case 'deepseek':
|
||||
return 'https://api.deepseek.com'
|
||||
default:
|
||||
return ''
|
||||
}
|
||||
}
|
||||
|
||||
// ===== 国产供应商用量单元格可见性(单一事实源) =====
|
||||
// CNProviderQuotaCell / CNProviderBalanceCell 与 AccountUsageCell 的占位符判定
|
||||
// 共用,避免多处复制条件后一处改另一处漏改。
|
||||
|
||||
export function cnQuotaCellVisible(platform: string, accountMode: string): boolean {
|
||||
return (platform === 'kimi' || platform === 'zhipu') && accountMode === 'coding'
|
||||
}
|
||||
|
||||
export function cnBalanceCellVisible(platform: string, accountMode: string): boolean {
|
||||
return (platform === 'kimi' || platform === 'deepseek') && accountMode !== 'coding'
|
||||
}
|
||||
|
||||
/**
|
||||
* 将请求头覆写写入 credentials。
|
||||
* create 模式:关闭时不写入任何字段;edit 模式:关闭时删除字段(全量替换语义)。
|
||||
|
||||
@@ -25,6 +25,25 @@
|
||||
d="M9.27 15.29l7.978-5.897c.391-.29.95-.177 1.137.272.98 2.369.542 5.215-1.41 7.169-1.951 1.954-4.667 2.382-7.149 1.406l-2.711 1.257c3.889 2.661 8.611 2.003 11.562-.953 2.341-2.344 3.066-5.539 2.388-8.42l.006.007c-.983-4.232.242-5.924 2.75-9.383.06-.082.12-.164.179-.248l-3.301 3.305v-.01L9.267 15.292M7.623 16.723c-2.792-2.67-2.31-6.801.071-9.184 1.761-1.763 4.647-2.483 7.166-1.425l2.705-1.25a7.808 7.808 0 00-1.829-1A8.975 8.975 0 005.984 5.83c-2.533 2.536-3.33 6.436-1.962 9.764 1.022 2.487-.653 4.246-2.34 6.022-.599.63-1.199 1.259-1.682 1.925l7.62-6.815"
|
||||
/>
|
||||
</svg>
|
||||
<!-- Kimi / Moonshot official logo mark (stylized K) -->
|
||||
<svg v-else-if="platform === 'kimi'" :class="sizeClass" viewBox="0 0 24 24" fill="currentColor">
|
||||
<path
|
||||
d="M21.765.351C22.998.351 24 1.353 24 2.586S22.998 4.82 21.765 4.82h-1.974c-.15 0-.26-.12-.26-.26V2.586A2.237 2.237 0 0 1 21.765.35M9.41 13.388l8.447-8.377c.16-.16.07-.471-.14-.471h-4.55s-.1.02-.14.06l-9.099 9.029c-.14.14-.35.02-.35-.21V4.81c0-.15-.1-.27-.221-.27H.22c-.12 0-.22.12-.22.27v18.57c0 .15.1.27.22.27h3.137c.12 0 .22-.12.22-.27v-3.79c0-.08.03-.16.08-.21l2.826-2.796c.07-.07.16-.08.241-.03l7.546 5.551a8.9 8.9 0 0 0 4.018 1.493c.12.01.23-.11.23-.27V19.76c0-.14-.08-.25-.19-.26a5.8 5.8 0 0 1-2.355-.942l-6.533-4.73c-.14-.09-.15-.32-.03-.441"
|
||||
/>
|
||||
</svg>
|
||||
<!-- Zhipu AI official logo mark (particle constellation) -->
|
||||
<svg v-else-if="platform === 'zhipu'" :class="sizeClass" viewBox="0 0 24 24" fill="currentColor">
|
||||
<path
|
||||
fill-rule="evenodd"
|
||||
d="M11.991 23.503a.24.24 0 0 0-.244.248a.24.24 0 0 0 .244.249a.24.24 0 0 0 .245-.249a.24.24 0 0 0-.22-.247zM9.671 5.365a1.697 1.697 0 0 1 1.099 2.132l-.071.172l-.016.04l-.018.054c-.07.16-.104.32-.104.498c-.035.71.47 1.279 1.186 1.314h.366c1.309.053 2.338 1.173 2.286 2.523c-.052 1.332-1.152 2.38-2.478 2.327h-.174c-.715.018-1.274.64-1.239 1.368c0 .124.018.23.053.337c.209.373.54.658.96.8c.75.23 1.517-.125 1.9-.782l.018-.035c.402-.64 1.17-.96 1.92-.711c.854.284 1.378 1.226 1.099 2.167a1.66 1.66 0 0 1-2.077 1.102a1.7 1.7 0 0 1-.907-.711l-.017-.035c-.2-.323-.463-.58-.851-.711l-.056-.018a1.646 1.646 0 0 0-1.954.746a1.66 1.66 0 0 1-1.065.764a1.677 1.677 0 0 1-1.989-1.279c-.209-.906.332-1.83 1.257-2.043a1.5 1.5 0 0 1 .296-.035h.018c.68-.071 1.151-.622 1.116-1.333a1.3 1.3 0 0 0-.227-.693a2.5 2.5 0 0 1-.366-1.403a2.4 2.4 0 0 1 .366-1.208c.14-.195.21-.444.227-.693c.018-.71-.506-1.261-1.186-1.332l-.07-.018a1.4 1.4 0 0 1-.299-.07l-.05-.019a1.7 1.7 0 0 1-1.047-2.114a1.68 1.68 0 0 1 2.094-1.101m-5.575 10.11c.26-.264.639-.367.994-.27s.633.379.728.74c.095.362-.007.748-.267 1.013c-.402.41-1.053.41-1.455 0a1.06 1.06 0 0 1 0-1.482zm14.845-.294c.359-.09.738.024.992.297c.254.274.344.665.237 1.025s-.396.634-.756.718c-.551.128-1.1-.22-1.23-.781a1.05 1.05 0 0 1 .757-1.26zm-.064-4.39c.314.32.49.753.49 1.206s-.176.886-.49 1.206c-.315.32-.74.5-1.185.5c-.444 0-.87-.18-1.184-.5a1.727 1.727 0 0 1 0-2.412a1.654 1.654 0 0 1 2.369 0m-11.243.163c.364.484.447 1.128.218 1.691a1.665 1.665 0 0 1-2.188.923c-.855-.36-1.26-1.358-.907-2.228a1.68 1.68 0 0 1 1.33-1.038a1.66 1.66 0 0 1 1.547.652m11.545-4.221c.368 0 .708.2.892.524s.184.724 0 1.048a1.03 1.03 0 0 1-.892.524a1.04 1.04 0 0 1-1.03-1.048a1.04 1.04 0 0 1 1.03-1.048m-14.358 0c.368 0 .707.2.891.524s.184.724 0 1.048a1.03 1.03 0 0 1-.891.524a1.04 1.04 0 0 1-1.03-1.048c0-.579.461-1.048 1.03-1.048m10.031-1.475c.925 0 1.675.764 1.675 1.706s-.75 1.705-1.675 1.705s-1.674-.763-1.674-1.705s.75-1.706 1.674-1.706m-2.626-.684c.362-.082.653-.356.761-.718a1.06 1.06 0 0 0-.238-1.028a1.02 1.02 0 0 0-.996-.294c-.547.14-.881.7-.752 1.257c.13.558.675.907 1.225.783m0 16.876c.359-.087.644-.36.75-.72a1.06 1.06 0 0 0-.237-1.019a1.02 1.02 0 0 0-.985-.301a1.04 1.04 0 0 0-.762.717c-.108.361-.017.754.239 1.028c.245.263.606.377.953.305zM17.19 3.5a.63.63 0 0 0 .628-.64a.63.63 0 0 0-.628-.64a.63.63 0 0 0-.628.64c0 .355.28.64.628.64m-10.38 0a.63.63 0 0 0 .628-.64c0-.355-.28-.64-.628-.64a.63.63 0 0 0-.628.64c0 .355.279.64.628.64m-5.182 7.852a.63.63 0 0 0-.628.64c0 .354.28.639.628.639a.63.63 0 0 0 .627-.606l.001-.034a.62.62 0 0 0-.628-.64zm5.182 9.13a.63.63 0 0 0-.628.64c0 .355.279.64.628.64a.63.63 0 0 0 .628-.64c0-.355-.28-.64-.628-.64m10.38.018a.63.63 0 0 0-.628.64c0 .355.28.64.628.64a.63.63 0 0 0 .628-.64a.63.63 0 0 0-.628-.64m5.182-9.148a.63.63 0 0 0-.628.64c0 .354.279.639.628.639a.63.63 0 0 0 .628-.64c0-.355-.28-.64-.628-.64zm-.384-4.992a.24.24 0 0 0 .244-.249a.24.24 0 0 0-.244-.249a.24.24 0 0 0-.244.249c0 .142.122.249.244.249M11.991.497a.24.24 0 0 0 .245-.248A.24.24 0 0 0 11.99 0a.24.24 0 0 0-.244.249c0 .133.108.236.223.247zM2.011 6.36a.24.24 0 0 0 .245-.249a.24.24 0 0 0-.244-.249a.24.24 0 0 0-.244.249a.24.24 0 0 0 .244.249zm0 11.263a.24.24 0 0 0-.243.248a.24.24 0 0 0 .244.249a.24.24 0 0 0 .244-.249a.25.25 0 0 0-.244-.248zm19.995-.018a.24.24 0 0 0-.245.248a.24.24 0 0 0 .245.25a.24.24 0 0 0 .244-.25a.25.25 0 0 0-.244-.248z"
|
||||
/>
|
||||
</svg>
|
||||
<!-- DeepSeek official logo mark (whale) -->
|
||||
<svg v-else-if="platform === 'deepseek'" :class="sizeClass" viewBox="0 0 24 24" fill="currentColor">
|
||||
<path
|
||||
d="M23.748 4.651c-.254-.124-.364.113-.512.233-.051.04-.094.09-.137.137-.372.397-.806.657-1.373.626-.829-.046-1.537.214-2.163.848-.133-.782-.575-1.248-1.247-1.548-.352-.155-.708-.311-.955-.65-.172-.24-.219-.509-.305-.774-.055-.16-.11-.323-.293-.35-.2-.031-.278.136-.356.276-.313.572-.434 1.202-.422 1.84.027 1.436.633 2.58 1.838 3.393.137.094.172.187.129.323-.082.28-.18.553-.266.833-.055.179-.137.218-.328.14a5.5 5.5 0 0 1-1.737-1.179c-.857-.828-1.631-1.743-2.597-2.46a12 12 0 0 0-.689-.47c-.985-.957.13-1.743.387-1.836.27-.098.094-.433-.778-.428-.872.003-1.67.295-2.687.685a3 3 0 0 1-.465.136a9.6 9.6 0 0 0-2.883-.101c-1.885.21-3.39 1.1-4.497 2.622C.082 8.776-.231 10.854.152 13.02c.403 2.284 1.568 4.175 3.36 5.653 1.857 1.533 3.997 2.284 6.438 2.14 1.482-.085 3.132-.284 4.994-1.86.47.234.962.328 1.78.398.629.058 1.235-.031 1.705-.129.735-.155.684-.836.418-.961-2.155-1.004-1.682-.595-2.112-.926 1.095-1.295 2.768-3.598 3.284-6.733.05-.346.115-.834.108-1.114-.004-.171.035-.238.23-.257a4.2 4.2 0 0 0 1.545-.475c1.397-.763 1.96-2.016 2.093-3.517.02-.23-.004-.467-.247-.588M11.58 18.168c-2.088-1.642-3.101-2.183-3.52-2.16-.39.024-.32.472-.234.763.09.288.207.487.371.74.114.167.192.416-.113.603-.673.416-1.842-.14-1.897-.168-1.361-.801-2.5-1.86-3.301-3.306-.775-1.393-1.225-2.888-1.299-4.482-.02-.385.094-.522.477-.592a4.7 4.7 0 0 1 1.53-.038c2.131.311 3.946 1.264 5.467 2.774.868.86 1.525 1.887 2.202 2.89.72 1.066 1.494 2.082 2.48 2.915.348.291.626.513.892.677-.802.09-2.14.109-3.055-.615zm1.001-6.44a.306.306 0 0 1 .415-.287a.3.3 0 0 1 .113.074a.3.3 0 0 1 .086.214c0 .17-.136.307-.308.307a.303.303 0 0 1-.306-.307m3.11 1.596c-.2.081-.4.151-.591.16a1.25 1.25 0 0 1-.798-.254c-.274-.23-.47-.358-.551-.758a1.7 1.7 0 0 1 .015-.588c.07-.327-.007-.537-.238-.727-.188-.156-.426-.199-.689-.199a.6.6 0 0 1-.254-.078a.253.253 0 0 1-.114-.358a1 1 0 0 1 .192-.21c.356-.202.767-.136 1.146.016.352.144.618.408 1.001.782.392.451.462.576.685.915.176.264.336.536.446.848.066.194-.02.353-.25.45"
|
||||
/>
|
||||
</svg>
|
||||
<!-- Composite group icon -->
|
||||
<svg v-else-if="platform === 'composite'" :class="sizeClass" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2">
|
||||
<circle cx="6" cy="12" r="3" />
|
||||
|
||||
@@ -90,6 +90,9 @@ const platformLabel = computed(() => {
|
||||
if (props.platform === 'openai') return 'OpenAI'
|
||||
if (props.platform === 'antigravity') return 'Antigravity'
|
||||
if (props.platform === 'grok') return 'Grok'
|
||||
if (props.platform === 'kimi') return 'Kimi'
|
||||
if (props.platform === 'zhipu') return 'Zhipu GLM'
|
||||
if (props.platform === 'deepseek') return 'DeepSeek'
|
||||
return 'Gemini'
|
||||
})
|
||||
|
||||
@@ -188,6 +191,15 @@ const platformClass = computed(() => {
|
||||
if (props.platform === 'grok') {
|
||||
return 'bg-zinc-100 text-zinc-700 dark:bg-zinc-800 dark:text-zinc-300'
|
||||
}
|
||||
if (props.platform === 'kimi') {
|
||||
return 'bg-pink-100 text-pink-700 dark:bg-pink-900/30 dark:text-pink-400'
|
||||
}
|
||||
if (props.platform === 'zhipu') {
|
||||
return 'bg-indigo-100 text-indigo-700 dark:bg-indigo-900/30 dark:text-indigo-400'
|
||||
}
|
||||
if (props.platform === 'deepseek') {
|
||||
return 'bg-teal-100 text-teal-700 dark:bg-teal-900/30 dark:text-teal-400'
|
||||
}
|
||||
return 'bg-blue-100 text-blue-700 dark:bg-blue-900/30 dark:text-blue-400'
|
||||
})
|
||||
|
||||
@@ -204,6 +216,15 @@ const typeClass = computed(() => {
|
||||
if (props.platform === 'grok') {
|
||||
return 'bg-zinc-100 text-zinc-600 dark:bg-zinc-800 dark:text-zinc-300'
|
||||
}
|
||||
if (props.platform === 'kimi') {
|
||||
return 'bg-pink-100 text-pink-600 dark:bg-pink-900/30 dark:text-pink-400'
|
||||
}
|
||||
if (props.platform === 'zhipu') {
|
||||
return 'bg-indigo-100 text-indigo-600 dark:bg-indigo-900/30 dark:text-indigo-400'
|
||||
}
|
||||
if (props.platform === 'deepseek') {
|
||||
return 'bg-teal-100 text-teal-600 dark:bg-teal-900/30 dark:text-teal-400'
|
||||
}
|
||||
return 'bg-blue-100 text-blue-600 dark:bg-blue-900/30 dark:text-blue-400'
|
||||
})
|
||||
|
||||
|
||||
@@ -182,7 +182,9 @@ const yiModels = [
|
||||
// Moonshot/Kimi
|
||||
const moonshotModels = [
|
||||
'moonshot-v1-8k', 'moonshot-v1-32k', 'moonshot-v1-128k',
|
||||
'kimi-latest'
|
||||
'kimi-latest',
|
||||
'kimi-for-coding',
|
||||
'kimi-k2'
|
||||
]
|
||||
|
||||
// 字节跳动 豆包
|
||||
@@ -428,7 +430,8 @@ export function getModelsByPlatform(platform: string): string[] {
|
||||
case 'grok': return xaiModels
|
||||
case 'cohere': return cohereModels
|
||||
case 'yi': return yiModels
|
||||
case 'moonshot': return moonshotModels
|
||||
case 'moonshot':
|
||||
case 'kimi': return moonshotModels
|
||||
case 'doubao': return doubaoModels
|
||||
case 'minimax': return minimaxModels
|
||||
case 'baidu': return baiduModels
|
||||
|
||||
@@ -104,6 +104,33 @@ export default {
|
||||
gemini: 'Gemini',
|
||||
antigravity: 'Antigravity',
|
||||
grok: 'Grok',
|
||||
kimi: 'Kimi',
|
||||
zhipu: 'Zhipu GLM',
|
||||
deepseek: 'DeepSeek',
|
||||
},
|
||||
cnProviders: {
|
||||
accountMode: {
|
||||
title: 'Account Type',
|
||||
payg: 'Pay-as-you-go',
|
||||
paygDesc: 'Consumes account balance, billed per token. Auto-cools down on low balance and recovers after top-up.',
|
||||
coding: 'Coding Plan',
|
||||
codingDesc: 'Subscription coding package, rate-limited by 5-hour / weekly rolling usage windows.',
|
||||
},
|
||||
apiProtocol: {
|
||||
title: 'API Protocol',
|
||||
chatCompletions: 'Chat Completions',
|
||||
chatCompletionsDesc: 'Standard OpenAI-compatible endpoint; requests in other formats are converted.',
|
||||
anthropic: 'Anthropic',
|
||||
anthropicDesc: 'Native passthrough to the provider’s Anthropic endpoint — ideal for Claude Code.',
|
||||
responses: 'Responses',
|
||||
responsesDesc: 'Provider’s native Responses endpoint — ideal for Codex.',
|
||||
},
|
||||
window5h: '5-hour window',
|
||||
windowWeekly: 'Weekly window',
|
||||
probeTooltip: 'Query the provider quota endpoint for 5-hour / weekly rolling window usage',
|
||||
balanceLow: 'Insufficient balance',
|
||||
noBalanceEndpoint: 'This platform has no balance query endpoint',
|
||||
resetSoon: 'reset soon',
|
||||
},
|
||||
types: {
|
||||
oauth: 'OAuth',
|
||||
|
||||
@@ -307,6 +307,33 @@ export default {
|
||||
gemini: 'Gemini',
|
||||
antigravity: 'Antigravity',
|
||||
grok: 'Grok',
|
||||
kimi: 'Kimi',
|
||||
zhipu: 'Zhipu GLM',
|
||||
deepseek: 'DeepSeek',
|
||||
},
|
||||
cnProviders: {
|
||||
accountMode: {
|
||||
title: '账号类型',
|
||||
payg: '按量付费',
|
||||
paygDesc: '消耗账户余额,按 Token 计费。余额不足自动冷却,充值后恢复。',
|
||||
coding: 'Coding Plan',
|
||||
codingDesc: '订阅制编程套餐,按 5 小时 / 每周滚动用量窗口限流。',
|
||||
},
|
||||
apiProtocol: {
|
||||
title: 'API 协议',
|
||||
chatCompletions: 'Chat Completions',
|
||||
chatCompletionsDesc: '标准 OpenAI 兼容端点,其他格式请求将被转换。',
|
||||
anthropic: 'Anthropic',
|
||||
anthropicDesc: '直通供应商原生 Anthropic 端点,零转换,适配 Claude Code。',
|
||||
responses: 'Responses',
|
||||
responsesDesc: '供应商原生 Responses 端点,适配 Codex。',
|
||||
},
|
||||
window5h: '5 小时窗口',
|
||||
windowWeekly: '每周窗口',
|
||||
probeTooltip: '请求供应商额度端点,查询 5 小时 / 每周滚动窗口用量',
|
||||
balanceLow: '余额不足',
|
||||
noBalanceEndpoint: '该平台暂无余额查询接口',
|
||||
resetSoon: '即将重置',
|
||||
},
|
||||
types: {
|
||||
oauth: 'OAuth',
|
||||
|
||||
@@ -525,7 +525,7 @@ export interface PaginationConfig {
|
||||
|
||||
// ==================== API Key & Group Types ====================
|
||||
|
||||
export type GroupPlatform = 'anthropic' | 'openai' | 'gemini' | 'antigravity' | 'grok' | 'composite'
|
||||
export type GroupPlatform = 'anthropic' | 'openai' | 'gemini' | 'antigravity' | 'grok' | 'kimi' | 'zhipu' | 'deepseek' | 'composite'
|
||||
|
||||
export type VideoModelPrices = Record<string, Record<string, number>>
|
||||
|
||||
@@ -882,7 +882,7 @@ export interface UpdateGroupRequest {
|
||||
|
||||
// ==================== Account & Proxy Types ====================
|
||||
|
||||
export type AccountPlatform = 'anthropic' | 'openai' | 'gemini' | 'antigravity' | 'grok'
|
||||
export type AccountPlatform = 'anthropic' | 'openai' | 'gemini' | 'antigravity' | 'grok' | 'kimi' | 'zhipu' | 'deepseek'
|
||||
export type AccountType = 'oauth' | 'setup-token' | 'apikey' | 'upstream' | 'bedrock' | 'service_account'
|
||||
export type OAuthAddMethod = 'oauth' | 'setup-token'
|
||||
export type ProxyProtocol = 'http' | 'https' | 'socks5' | 'socks5h'
|
||||
|
||||
@@ -5,7 +5,16 @@
|
||||
* instead of defining their own color mappings.
|
||||
*/
|
||||
|
||||
export type Platform = 'anthropic' | 'openai' | 'antigravity' | 'gemini' | 'grok' | 'composite'
|
||||
export type Platform =
|
||||
| 'anthropic'
|
||||
| 'openai'
|
||||
| 'antigravity'
|
||||
| 'gemini'
|
||||
| 'grok'
|
||||
| 'kimi'
|
||||
| 'zhipu'
|
||||
| 'deepseek'
|
||||
| 'composite'
|
||||
|
||||
// ── Badge (bg + text + border, for inline badges with border) ───────
|
||||
const BADGE: Record<Platform, string> = {
|
||||
@@ -14,6 +23,9 @@ const BADGE: Record<Platform, string> = {
|
||||
antigravity: 'bg-purple-500/10 text-purple-600 border-purple-500/30 dark:text-purple-400',
|
||||
gemini: 'bg-blue-500/10 text-blue-600 border-blue-500/30 dark:text-blue-400',
|
||||
grok: 'bg-zinc-800/10 text-zinc-800 border-zinc-800/30 dark:bg-zinc-500/10 dark:text-zinc-200 dark:border-zinc-500/30',
|
||||
kimi: 'bg-pink-500/10 text-pink-600 border-pink-500/30 dark:text-pink-400',
|
||||
zhipu: 'bg-indigo-500/10 text-indigo-600 border-indigo-500/30 dark:text-indigo-400',
|
||||
deepseek: 'bg-teal-500/10 text-teal-600 border-teal-500/30 dark:text-teal-400',
|
||||
composite: 'bg-cyan-500/10 text-cyan-700 border-cyan-500/30 dark:text-cyan-300',
|
||||
}
|
||||
const BADGE_DEFAULT = 'bg-slate-500/10 text-slate-600 border-slate-500/30 dark:text-slate-400'
|
||||
@@ -25,6 +37,9 @@ const BADGE_LIGHT: Record<Platform, string> = {
|
||||
antigravity: 'bg-purple-500/10 text-purple-600 dark:bg-purple-500/10 dark:text-purple-300',
|
||||
gemini: 'bg-blue-500/10 text-blue-600 dark:bg-blue-500/10 dark:text-blue-300',
|
||||
grok: 'bg-zinc-800/10 text-zinc-800 dark:bg-zinc-500/10 dark:text-zinc-200',
|
||||
kimi: 'bg-pink-500/10 text-pink-600 dark:bg-pink-500/10 dark:text-pink-300',
|
||||
zhipu: 'bg-indigo-500/10 text-indigo-600 dark:bg-indigo-500/10 dark:text-indigo-300',
|
||||
deepseek: 'bg-teal-500/10 text-teal-600 dark:bg-teal-500/10 dark:text-teal-300',
|
||||
composite: 'bg-cyan-500/10 text-cyan-700 dark:bg-cyan-500/10 dark:text-cyan-300',
|
||||
}
|
||||
|
||||
@@ -35,6 +50,9 @@ const BORDER: Record<Platform, string> = {
|
||||
antigravity: 'border-purple-500/20 dark:border-purple-500/20',
|
||||
gemini: 'border-blue-500/20 dark:border-blue-500/20',
|
||||
grok: 'border-zinc-800/20 dark:border-zinc-500/20',
|
||||
kimi: 'border-pink-500/20 dark:border-pink-500/20',
|
||||
zhipu: 'border-indigo-500/20 dark:border-indigo-500/20',
|
||||
deepseek: 'border-teal-500/20 dark:border-teal-500/20',
|
||||
composite: 'border-cyan-500/20 dark:border-cyan-500/20',
|
||||
}
|
||||
const BORDER_DEFAULT = 'border-gray-200 dark:border-dark-700'
|
||||
@@ -46,6 +64,9 @@ const BORDER_STRONG: Record<Platform, string> = {
|
||||
antigravity: 'border-purple-500/35 dark:border-purple-500/30',
|
||||
gemini: 'border-blue-500/35 dark:border-blue-500/30',
|
||||
grok: 'border-zinc-800/35 dark:border-zinc-500/35',
|
||||
kimi: 'border-pink-500/35 dark:border-pink-500/30',
|
||||
zhipu: 'border-indigo-500/35 dark:border-indigo-500/30',
|
||||
deepseek: 'border-teal-500/35 dark:border-teal-500/30',
|
||||
composite: 'border-cyan-500/35 dark:border-cyan-500/30',
|
||||
}
|
||||
const BORDER_STRONG_DEFAULT = 'border-gray-300 dark:border-dark-600'
|
||||
@@ -58,6 +79,9 @@ const ACCENT: Record<Platform, string> = {
|
||||
antigravity: '#a855f7', // purple-500
|
||||
gemini: '#3b82f6', // blue-500
|
||||
grok: '#71717a', // zinc-500
|
||||
kimi: '#ec4899', // pink-500
|
||||
zhipu: '#6366f1', // indigo-500
|
||||
deepseek: '#14b8a6', // teal-500
|
||||
composite: '#06b6d4', // cyan-500
|
||||
}
|
||||
const ACCENT_DEFAULT = '#14b8a6' // primary-500 (teal)
|
||||
@@ -69,6 +93,9 @@ const ACCENT_BAR: Record<Platform, string> = {
|
||||
antigravity: 'bg-gradient-to-r from-purple-400 to-purple-500',
|
||||
gemini: 'bg-gradient-to-r from-blue-400 to-blue-500',
|
||||
grok: 'bg-gradient-to-r from-zinc-700 to-zinc-900',
|
||||
kimi: 'bg-gradient-to-r from-pink-400 to-pink-500',
|
||||
zhipu: 'bg-gradient-to-r from-indigo-400 to-indigo-500',
|
||||
deepseek: 'bg-gradient-to-r from-teal-400 to-teal-500',
|
||||
composite: 'bg-gradient-to-r from-slate-500 to-cyan-500',
|
||||
}
|
||||
const ACCENT_BAR_DEFAULT = 'bg-gradient-to-r from-primary-400 to-primary-500'
|
||||
@@ -80,6 +107,9 @@ const TEXT: Record<Platform, string> = {
|
||||
antigravity: 'text-purple-600 dark:text-purple-400',
|
||||
gemini: 'text-blue-600 dark:text-blue-400',
|
||||
grok: 'text-zinc-800 dark:text-zinc-200',
|
||||
kimi: 'text-pink-600 dark:text-pink-400',
|
||||
zhipu: 'text-indigo-600 dark:text-indigo-400',
|
||||
deepseek: 'text-teal-600 dark:text-teal-400',
|
||||
composite: 'text-cyan-700 dark:text-cyan-300',
|
||||
}
|
||||
const TEXT_DEFAULT = 'text-primary-600 dark:text-primary-400'
|
||||
@@ -91,6 +121,9 @@ const ICON: Record<Platform, string> = {
|
||||
antigravity: 'text-purple-500 dark:text-purple-400',
|
||||
gemini: 'text-blue-500 dark:text-blue-400',
|
||||
grok: 'text-zinc-800 dark:text-zinc-200',
|
||||
kimi: 'text-pink-500 dark:text-pink-400',
|
||||
zhipu: 'text-indigo-500 dark:text-indigo-400',
|
||||
deepseek: 'text-teal-500 dark:text-teal-400',
|
||||
composite: 'text-cyan-600 dark:text-cyan-300',
|
||||
}
|
||||
const ICON_DEFAULT = 'text-primary-500 dark:text-primary-400'
|
||||
@@ -102,6 +135,9 @@ const BUTTON: Record<Platform, string> = {
|
||||
antigravity: 'bg-purple-500 text-white hover:bg-purple-600 active:bg-purple-700 dark:bg-purple-500/80 dark:hover:bg-purple-500',
|
||||
gemini: 'bg-blue-500 text-white hover:bg-blue-600 active:bg-blue-700 dark:bg-blue-500/80 dark:hover:bg-blue-500',
|
||||
grok: 'bg-zinc-800 text-white hover:bg-zinc-900 active:bg-black dark:bg-zinc-700 dark:hover:bg-zinc-600',
|
||||
kimi: 'bg-pink-500 text-white hover:bg-pink-600 active:bg-pink-700 dark:bg-pink-500/80 dark:hover:bg-pink-500',
|
||||
zhipu: 'bg-indigo-500 text-white hover:bg-indigo-600 active:bg-indigo-700 dark:bg-indigo-500/80 dark:hover:bg-indigo-500',
|
||||
deepseek: 'bg-teal-500 text-white hover:bg-teal-600 active:bg-teal-700 dark:bg-teal-500/80 dark:hover:bg-teal-500',
|
||||
composite: 'bg-cyan-700 text-white hover:bg-cyan-800 active:bg-cyan-900 dark:bg-cyan-600 dark:hover:bg-cyan-500',
|
||||
}
|
||||
const BUTTON_DEFAULT = 'bg-primary-500 text-white hover:bg-primary-600 dark:bg-primary-600 dark:hover:bg-primary-500'
|
||||
@@ -113,6 +149,9 @@ const DISCOUNT: Record<Platform, string> = {
|
||||
antigravity: 'bg-purple-100 text-purple-700 dark:bg-purple-900/40 dark:text-purple-300',
|
||||
gemini: 'bg-blue-100 text-blue-700 dark:bg-blue-900/40 dark:text-blue-300',
|
||||
grok: 'bg-zinc-100 text-zinc-800 dark:bg-zinc-800 dark:text-zinc-200',
|
||||
kimi: 'bg-pink-100 text-pink-700 dark:bg-pink-900/40 dark:text-pink-300',
|
||||
zhipu: 'bg-indigo-100 text-indigo-700 dark:bg-indigo-900/40 dark:text-indigo-300',
|
||||
deepseek: 'bg-teal-100 text-teal-700 dark:bg-teal-900/40 dark:text-teal-300',
|
||||
composite: 'bg-cyan-100 text-cyan-800 dark:bg-cyan-900/40 dark:text-cyan-300',
|
||||
}
|
||||
const DISCOUNT_DEFAULT = 'bg-red-100 text-red-700 dark:bg-red-900/40 dark:text-red-300'
|
||||
@@ -124,6 +163,9 @@ const GRADIENT: Record<Platform, string> = {
|
||||
antigravity: 'from-purple-500 to-purple-600',
|
||||
gemini: 'from-blue-500 to-blue-600',
|
||||
grok: 'from-zinc-700 to-zinc-900',
|
||||
kimi: 'from-pink-500 to-pink-600',
|
||||
zhipu: 'from-indigo-500 to-indigo-600',
|
||||
deepseek: 'from-teal-500 to-teal-600',
|
||||
composite: 'from-slate-600 to-cyan-600',
|
||||
}
|
||||
const GRADIENT_DEFAULT = 'from-primary-500 to-primary-600'
|
||||
@@ -135,6 +177,9 @@ const GRADIENT_TEXT: Record<Platform, string> = {
|
||||
antigravity: 'text-purple-100',
|
||||
gemini: 'text-blue-100',
|
||||
grok: 'text-zinc-100',
|
||||
kimi: 'text-pink-100',
|
||||
zhipu: 'text-indigo-100',
|
||||
deepseek: 'text-teal-100',
|
||||
composite: 'text-cyan-100',
|
||||
}
|
||||
const GRADIENT_TEXT_DEFAULT = 'text-primary-100'
|
||||
@@ -145,6 +190,9 @@ const GRADIENT_SUBTEXT: Record<Platform, string> = {
|
||||
antigravity: 'text-purple-200',
|
||||
gemini: 'text-blue-200',
|
||||
grok: 'text-zinc-300',
|
||||
kimi: 'text-pink-200',
|
||||
zhipu: 'text-indigo-200',
|
||||
deepseek: 'text-teal-200',
|
||||
composite: 'text-cyan-200',
|
||||
}
|
||||
const GRADIENT_SUBTEXT_DEFAULT = 'text-primary-200'
|
||||
@@ -152,7 +200,17 @@ const GRADIENT_SUBTEXT_DEFAULT = 'text-primary-200'
|
||||
// ── Public API ──────────────────────────────────────────────────────
|
||||
|
||||
function isPlatform(p: string): p is Platform {
|
||||
return p === 'anthropic' || p === 'openai' || p === 'antigravity' || p === 'gemini' || p === 'grok' || p === 'composite'
|
||||
return (
|
||||
p === 'anthropic' ||
|
||||
p === 'openai' ||
|
||||
p === 'antigravity' ||
|
||||
p === 'gemini' ||
|
||||
p === 'grok' ||
|
||||
p === 'kimi' ||
|
||||
p === 'zhipu' ||
|
||||
p === 'deepseek' ||
|
||||
p === 'composite'
|
||||
)
|
||||
}
|
||||
|
||||
export function platformBadgeClass(p: string): string {
|
||||
@@ -214,6 +272,9 @@ export function platformLabel(p: string): string {
|
||||
case 'antigravity': return 'Antigravity'
|
||||
case 'gemini': return 'Gemini'
|
||||
case 'grok': return 'Grok'
|
||||
case 'kimi': return 'Kimi'
|
||||
case 'zhipu': return 'Zhipu GLM'
|
||||
case 'deepseek': return 'DeepSeek'
|
||||
case 'composite': return 'Composite'
|
||||
default: return p || 'API'
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user