diff --git a/backend/cmd/server/VERSION b/backend/cmd/server/VERSION index 3ef481c326..a0e8ec1d4e 100644 --- a/backend/cmd/server/VERSION +++ b/backend/cmd/server/VERSION @@ -1 +1 @@ -0.1.144 +0.1.145 diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index aae4c405b9..e148c5c363 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -67,7 +67,11 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { serviceUserPlatformQuotaRepository := repository.NewUserPlatformQuotaServiceAdapter(userPlatformQuotaRepository) billingCacheService := service.ProvideBillingCacheService(billingCache, userRepository, userSubscriptionRepository, apiKeyRepository, userRPMCache, userGroupRateRepository, configConfig, serviceUserPlatformQuotaRepository) apiKeyCache := repository.NewAPIKeyCache(redisClient) - apiKeyService := service.ProvideAPIKeyService(apiKeyRepository, userRepository, groupRepository, userSubscriptionRepository, userGroupRateRepository, apiKeyCache, configConfig, billingCacheService) + concurrencyCache := repository.ProvideConcurrencyCache(redisClient, configConfig) + schedulerCache := repository.ProvideSchedulerCache(redisClient, configConfig) + accountRepository := repository.NewAccountRepository(client, db, schedulerCache) + concurrencyService := service.ProvideConcurrencyService(concurrencyCache, accountRepository, configConfig) + apiKeyService := service.ProvideAPIKeyService(apiKeyRepository, userRepository, groupRepository, userSubscriptionRepository, userGroupRateRepository, apiKeyCache, configConfig, billingCacheService, concurrencyService) apiKeyAuthCacheInvalidator := service.ProvideAPIKeyAuthCacheInvalidator(apiKeyService) promoService := service.NewPromoService(promoCodeRepository, userRepository, billingCacheService, client, apiKeyAuthCacheInvalidator) subscriptionService := service.NewSubscriptionService(groupRepository, userSubscriptionRepository, billingCacheService, client, configConfig) @@ -92,13 +96,9 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { usageLogRepository := repository.NewUsageLogRepository(client, db) usageService := service.NewUsageService(usageLogRepository, userRepository, client, apiKeyAuthCacheInvalidator) opsRepository := repository.NewOpsRepository(db) - schedulerCache := repository.ProvideSchedulerCache(redisClient, configConfig) - accountRepository := repository.NewAccountRepository(client, db, schedulerCache) batchImageRepository := repository.NewBatchImageRepository(db) batchImageQueue := repository.NewBatchImageQueue(redisClient, configConfig) batchImageDownloadLimiter := repository.NewBatchImageDownloadLimiter(redisClient, configConfig) - concurrencyCache := repository.ProvideConcurrencyCache(redisClient, configConfig) - concurrencyService := service.ProvideConcurrencyService(concurrencyCache, accountRepository, configConfig) usageBillingRepository := repository.NewUsageBillingRepository(client, db) gatewayCache := repository.NewGatewayCache(redisClient) schedulerOutboxRepository := repository.NewSchedulerOutboxRepository(db) diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index efe95801c3..8c91245fbf 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -196,14 +196,6 @@ type AccountSchedulerGroupScore struct { const accountListGroupUngroupedQueryValue = "ungrouped" -type openAIAccountSchedulerScorePoolLister interface { - ListOpenAISchedulableAccountsForSchedulerScore(ctx context.Context, groupID *int64) ([]service.Account, error) -} - -type accountSchedulerScoreFilterPoolLister interface { - ListAccountsForSchedulerScoreFilter(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]service.Account, error) -} - func (h *AccountHandler) buildAccountResponseWithRuntime(ctx context.Context, account *service.Account) AccountWithConcurrency { item := AccountWithConcurrency{ Account: dto.AccountFromService(account), @@ -250,33 +242,27 @@ func (h *AccountHandler) buildAccountResponseWithRuntime(ctx context.Context, ac return item } -func (h *AccountHandler) scoreOpenAIAccountSchedulerPool(ctx context.Context, accounts []service.Account) map[int64]AccountSchedulerScore { +// scoreOpenAIAccountSchedulerPool 对池内 OpenAI 账号计算调度分数快照。 +// loadMap 为共享的账号负载数据(含池内全部账号即可,多余条目无害);传 nil 时自行批查。 +func (h *AccountHandler) scoreOpenAIAccountSchedulerPool(ctx context.Context, accounts []service.Account, loadMap map[int64]*service.AccountLoadInfo) map[int64]AccountSchedulerScore { if len(accounts) == 0 { return nil } openAIAccounts := make([]*service.Account, 0, len(accounts)) - loadReq := make([]service.AccountWithConcurrency, 0, len(accounts)) for i := range accounts { account := &accounts[i] if account.Platform != service.PlatformOpenAI { continue } openAIAccounts = append(openAIAccounts, account) - loadReq = append(loadReq, service.AccountWithConcurrency{ - ID: account.ID, - MaxConcurrency: account.EffectiveLoadFactor(), - }) } if len(openAIAccounts) == 0 { return nil } - loadMap := map[int64]*service.AccountLoadInfo{} - if h.concurrencyService != nil { - if batchLoad, err := h.concurrencyService.GetAccountsLoadBatch(ctx, loadReq); err == nil && batchLoad != nil { - loadMap = batchLoad - } + if loadMap == nil { + loadMap = h.fetchOpenAIAccountLoadMap(ctx, openAIAccounts) } var scores map[int64]service.OpenAIAccountSchedulerScoreSnapshot @@ -297,6 +283,36 @@ func (h *AccountHandler) scoreOpenAIAccountSchedulerPool(ctx context.Context, ac return result } +// fetchOpenAIAccountLoadMap 一次性批查给定 OpenAI 账号的负载数据; +// 失败时记录日志并返回空表(分数按零负载计算,属可接受降级)。 +func (h *AccountHandler) fetchOpenAIAccountLoadMap(ctx context.Context, openAIAccounts []*service.Account) map[int64]*service.AccountLoadInfo { + loadMap := map[int64]*service.AccountLoadInfo{} + if h.concurrencyService == nil || len(openAIAccounts) == 0 { + return loadMap + } + seen := make(map[int64]struct{}, len(openAIAccounts)) + loadReq := make([]service.AccountWithConcurrency, 0, len(openAIAccounts)) + for _, account := range openAIAccounts { + if account == nil { + continue + } + if _, ok := seen[account.ID]; ok { + continue + } + seen[account.ID] = struct{}{} + loadReq = append(loadReq, service.AccountWithConcurrency{ + ID: account.ID, + MaxConcurrency: account.EffectiveLoadFactor(), + }) + } + if batchLoad, err := h.concurrencyService.GetAccountsLoadBatch(ctx, loadReq); err != nil { + slog.Warn("openai_scheduler_score_load_batch_failed", "error", err) + } else if batchLoad != nil { + loadMap = batchLoad + } + return loadMap +} + func (h *AccountHandler) buildOpenAIAccountSchedulerScores( ctx context.Context, accounts []service.Account, @@ -309,12 +325,6 @@ func (h *AccountHandler) buildOpenAIAccountSchedulerScores( filterPool = accounts } - baseScores := make(map[int64]*AccountSchedulerScore) - for accountID, score := range h.scoreOpenAIAccountSchedulerPool(ctx, filterPool) { - copiedScore := score - baseScores[accountID] = &copiedScore - } - pageOpenAIAccountIDs := make(map[int64]struct{}) groupIDs := make(map[int64]struct{}) for i := range accounts { @@ -338,7 +348,48 @@ func (h *AccountHandler) buildOpenAIAccountSchedulerScores( } } if len(pageOpenAIAccountIDs) == 0 { - return baseScores, nil + return nil, nil + } + + // 先取各分组池,再对"过滤池 ∪ 分组池"的账号并集做一次负载批查, + // 避免每个池各查一次 Redis 的 N+1。 + groupIDList := make([]int64, 0, len(groupIDs)) + for groupID := range groupIDs { + groupIDList = append(groupIDList, groupID) + } + sort.Slice(groupIDList, func(i, j int) bool { return groupIDList[i] < groupIDList[j] }) + + groupPools := make(map[int64][]service.Account, len(groupIDList)) + if h.adminService != nil { + for _, groupID := range groupIDList { + gid := groupID + pool, err := h.adminService.ListOpenAISchedulableAccountsForSchedulerScore(ctx, &gid) + if err != nil { + slog.Warn("openai_scheduler_group_score_pool_failed", "group_id", gid, "error", err) + continue + } + groupPools[gid] = pool + } + } + + loadUnion := make([]*service.Account, 0, len(filterPool)) + collectOpenAIAccounts := func(pool []service.Account) { + for i := range pool { + if pool[i].Platform == service.PlatformOpenAI { + loadUnion = append(loadUnion, &pool[i]) + } + } + } + collectOpenAIAccounts(filterPool) + for _, pool := range groupPools { + collectOpenAIAccounts(pool) + } + loadMap := h.fetchOpenAIAccountLoadMap(ctx, loadUnion) + + baseScores := make(map[int64]*AccountSchedulerScore) + for accountID, score := range h.scoreOpenAIAccountSchedulerPool(ctx, filterPool, loadMap) { + copiedScore := score + baseScores[accountID] = &copiedScore } groupScoresByAccount := make(map[int64][]AccountSchedulerGroupScore) @@ -346,7 +397,7 @@ func (h *AccountHandler) buildOpenAIAccountSchedulerScores( if len(pool) == 0 { return } - scores := h.scoreOpenAIAccountSchedulerPool(ctx, pool) + scores := h.scoreOpenAIAccountSchedulerPool(ctx, pool, loadMap) for accountID, schedulerScore := range scores { if _, ok := pageOpenAIAccountIDs[accountID]; !ok { continue @@ -365,37 +416,27 @@ func (h *AccountHandler) buildOpenAIAccountSchedulerScores( } } - if lister, ok := h.adminService.(openAIAccountSchedulerScorePoolLister); ok { - groupIDList := make([]int64, 0, len(groupIDs)) - for groupID := range groupIDs { - groupIDList = append(groupIDList, groupID) + for _, groupID := range groupIDList { + gid := groupID + pool, ok := groupPools[gid] + if !ok { + continue } - sort.Slice(groupIDList, func(i, j int) bool { return groupIDList[i] < groupIDList[j] }) - - for _, groupID := range groupIDList { - gid := groupID - pool, err := lister.ListOpenAISchedulableAccountsForSchedulerScore(ctx, &gid) - if err != nil { - slog.Warn("openai_scheduler_group_score_pool_failed", "group_id", gid, "error", err) - continue - } - groupNameByID := make(map[int64]string) - groupPriorityByAccount := make(map[int64]int) - for i := range pool { - account := &pool[i] - for _, accountGroup := range account.AccountGroups { - if accountGroup.GroupID != gid { - continue - } - groupPriorityByAccount[account.ID] = accountGroup.Priority - if accountGroup.Group != nil { - groupNameByID[gid] = accountGroup.Group.Name - } + groupNameByID := make(map[int64]string) + groupPriorityByAccount := make(map[int64]int) + for i := range pool { + account := &pool[i] + for _, accountGroup := range account.AccountGroups { + if accountGroup.GroupID != gid { + continue + } + groupPriorityByAccount[account.ID] = accountGroup.Priority + if accountGroup.Group != nil { + groupNameByID[gid] = accountGroup.Group.Name } } - scoreGroupPool(&gid, groupNameByID, groupPriorityByAccount, pool) } - + scoreGroupPool(&gid, groupNameByID, groupPriorityByAccount, pool) } for accountID := range groupScoresByAccount { @@ -417,11 +458,9 @@ func (h *AccountHandler) listAccountSchedulerScoreFilterPool( if h.adminService == nil || (platform != "" && platform != service.PlatformOpenAI) { return nil } - lister, ok := h.adminService.(accountSchedulerScoreFilterPoolLister) - if !ok { - return nil - } - accounts, err := lister.ListAccountsForSchedulerScoreFilter(ctx, platform, accountType, status, search, groupID, privacyMode) + // 池只用于 OpenAI 分数计算(非 OpenAI 账号会在打分时被丢弃), + // 无论列表页平台过滤为何,查询一律限定 openai,避免无过滤时全表扫描。 + accounts, err := h.adminService.ListAccountsForSchedulerScoreFilter(ctx, service.PlatformOpenAI, accountType, status, search, groupID, privacyMode) if err != nil { slog.Warn("openai_scheduler_filter_score_pool_failed", "error", err) return nil @@ -481,8 +520,20 @@ func (h *AccountHandler) List(c *gin.Context) { var windowCosts map[int64]float64 var activeSessions map[int64]int var rpmCounts map[int64]int - schedulerFilterPool := h.listAccountSchedulerScoreFilterPool(c.Request.Context(), platform, accountType, status, search, groupID, privacyMode) - schedulerScores, schedulerGroupScores := h.buildOpenAIAccountSchedulerScores(c.Request.Context(), accounts, schedulerFilterPool) + // 仅当前页存在 OpenAI 账号时才计算调度分数,避免为空结果付出池查询开销。 + var schedulerScores map[int64]*AccountSchedulerScore + var schedulerGroupScores map[int64][]AccountSchedulerGroupScore + pageHasOpenAIAccounts := false + for i := range accounts { + if accounts[i].Platform == service.PlatformOpenAI { + pageHasOpenAIAccounts = true + break + } + } + if pageHasOpenAIAccounts { + schedulerFilterPool := h.listAccountSchedulerScoreFilterPool(c.Request.Context(), platform, accountType, status, search, groupID, privacyMode) + schedulerScores, schedulerGroupScores = h.buildOpenAIAccountSchedulerScores(c.Request.Context(), accounts, schedulerFilterPool) + } // 始终获取并发数(Redis ZCARD,极低开销) if h.concurrencyService != nil { diff --git a/backend/internal/handler/admin/setting_handler.go b/backend/internal/handler/admin/setting_handler.go index 624fddb1e6..529e46c575 100644 --- a/backend/internal/handler/admin/setting_handler.go +++ b/backend/internal/handler/admin/setting_handler.go @@ -311,6 +311,7 @@ func (h *SettingHandler) GetSettings(c *gin.Context) { PaymentEnabledTypes: paymentCfg.EnabledTypes, PaymentBalanceDisabled: paymentCfg.BalanceDisabled, PaymentBalanceRechargeMultiplier: paymentCfg.BalanceRechargeMultiplier, + PaymentSubscriptionUSDToCNYRate: paymentCfg.SubscriptionUSDToCNYRate, PaymentRechargeFeeRate: paymentCfg.RechargeFeeRate, PaymentLoadBalanceStrat: paymentCfg.LoadBalanceStrategy, PaymentProductNamePrefix: paymentCfg.ProductNamePrefix, @@ -672,6 +673,7 @@ type UpdateSettingsRequest struct { PaymentEnabledTypes []string `json:"payment_enabled_types"` PaymentBalanceDisabled *bool `json:"payment_balance_disabled"` PaymentBalanceRechargeMultiplier *float64 `json:"payment_balance_recharge_multiplier"` + PaymentSubscriptionUSDToCNYRate *float64 `json:"payment_subscription_usd_to_cny_rate"` PaymentRechargeFeeRate *float64 `json:"payment_recharge_fee_rate"` PaymentLoadBalanceStrat *string `json:"payment_load_balance_strategy"` PaymentProductNamePrefix *string `json:"payment_product_name_prefix"` @@ -2015,6 +2017,7 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { EnabledTypes: req.PaymentEnabledTypes, BalanceDisabled: req.PaymentBalanceDisabled, BalanceRechargeMultiplier: req.PaymentBalanceRechargeMultiplier, + SubscriptionUSDToCNYRate: req.PaymentSubscriptionUSDToCNYRate, RechargeFeeRate: req.PaymentRechargeFeeRate, LoadBalanceStrategy: req.PaymentLoadBalanceStrat, ProductNamePrefix: req.PaymentProductNamePrefix, @@ -2258,6 +2261,7 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { PaymentEnabledTypes: updatedPaymentCfg.EnabledTypes, PaymentBalanceDisabled: updatedPaymentCfg.BalanceDisabled, PaymentBalanceRechargeMultiplier: updatedPaymentCfg.BalanceRechargeMultiplier, + PaymentSubscriptionUSDToCNYRate: updatedPaymentCfg.SubscriptionUSDToCNYRate, PaymentRechargeFeeRate: updatedPaymentCfg.RechargeFeeRate, PaymentLoadBalanceStrat: updatedPaymentCfg.LoadBalanceStrategy, PaymentProductNamePrefix: updatedPaymentCfg.ProductNamePrefix, @@ -2316,7 +2320,8 @@ func hasPaymentFields(req UpdateSettingsRequest) bool { req.PaymentMaxAmount != nil || req.PaymentDailyLimit != nil || req.PaymentOrderTimeoutMin != nil || req.PaymentMaxPendingOrders != nil || req.PaymentEnabledTypes != nil || req.PaymentBalanceDisabled != nil || - req.PaymentBalanceRechargeMultiplier != nil || req.PaymentRechargeFeeRate != nil || + req.PaymentBalanceRechargeMultiplier != nil || req.PaymentSubscriptionUSDToCNYRate != nil || + req.PaymentRechargeFeeRate != nil || req.PaymentLoadBalanceStrat != nil || req.PaymentProductNamePrefix != nil || req.PaymentProductNameSuffix != nil || req.PaymentHelpImageURL != nil || req.PaymentHelpText != nil || req.PaymentCancelRateLimitEnabled != nil || diff --git a/backend/internal/handler/dto/api_key_mapper_last_used_test.go b/backend/internal/handler/dto/api_key_mapper_last_used_test.go index 99644ced7f..d63baba91a 100644 --- a/backend/internal/handler/dto/api_key_mapper_last_used_test.go +++ b/backend/internal/handler/dto/api_key_mapper_last_used_test.go @@ -11,18 +11,20 @@ import ( func TestAPIKeyFromService_MapsLastUsedAt(t *testing.T) { lastUsed := time.Now().UTC().Truncate(time.Second) src := &service.APIKey{ - ID: 1, - UserID: 2, - Key: "sk-map-last-used", - Name: "Mapper", - Status: service.StatusActive, - LastUsedAt: &lastUsed, + ID: 1, + UserID: 2, + Key: "sk-map-last-used", + Name: "Mapper", + Status: service.StatusActive, + LastUsedAt: &lastUsed, + CurrentConcurrency: 3, } out := APIKeyFromService(src) require.NotNil(t, out) require.NotNil(t, out.LastUsedAt) require.WithinDuration(t, lastUsed, *out.LastUsedAt, time.Second) + require.Equal(t, 3, out.CurrentConcurrency) } func TestAPIKeyFromService_MapsNilLastUsedAt(t *testing.T) { diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index 7949b278cd..03e4c97309 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -80,31 +80,32 @@ func APIKeyFromService(k *service.APIKey) *APIKey { return nil } out := &APIKey{ - ID: k.ID, - UserID: k.UserID, - Key: k.Key, - Name: k.Name, - GroupID: k.GroupID, - Status: k.Status, - IPWhitelist: k.IPWhitelist, - IPBlacklist: k.IPBlacklist, - LastUsedAt: k.LastUsedAt, - Quota: k.Quota, - QuotaUsed: k.QuotaUsed, - ExpiresAt: k.ExpiresAt, - CreatedAt: k.CreatedAt, - UpdatedAt: k.UpdatedAt, - RateLimit5h: k.RateLimit5h, - RateLimit1d: k.RateLimit1d, - RateLimit7d: k.RateLimit7d, - Usage5h: k.EffectiveUsage5h(), - Usage1d: k.EffectiveUsage1d(), - Usage7d: k.EffectiveUsage7d(), - Window5hStart: k.Window5hStart, - Window1dStart: k.Window1dStart, - Window7dStart: k.Window7dStart, - User: UserFromServiceShallow(k.User), - Group: GroupFromServiceShallow(k.Group), + ID: k.ID, + UserID: k.UserID, + Key: k.Key, + Name: k.Name, + GroupID: k.GroupID, + Status: k.Status, + IPWhitelist: k.IPWhitelist, + IPBlacklist: k.IPBlacklist, + LastUsedAt: k.LastUsedAt, + Quota: k.Quota, + QuotaUsed: k.QuotaUsed, + ExpiresAt: k.ExpiresAt, + CreatedAt: k.CreatedAt, + UpdatedAt: k.UpdatedAt, + CurrentConcurrency: k.CurrentConcurrency, + RateLimit5h: k.RateLimit5h, + RateLimit1d: k.RateLimit1d, + RateLimit7d: k.RateLimit7d, + Usage5h: k.EffectiveUsage5h(), + Usage1d: k.EffectiveUsage1d(), + Usage7d: k.EffectiveUsage7d(), + Window5hStart: k.Window5hStart, + Window1dStart: k.Window1dStart, + Window7dStart: k.Window7dStart, + User: UserFromServiceShallow(k.User), + Group: GroupFromServiceShallow(k.Group), } if k.Window5hStart != nil && !service.IsWindowExpired(k.Window5hStart, service.RateLimitWindow5h) { t := k.Window5hStart.Add(service.RateLimitWindow5h) diff --git a/backend/internal/handler/dto/settings.go b/backend/internal/handler/dto/settings.go index 9a20191a0c..99fba54980 100644 --- a/backend/internal/handler/dto/settings.go +++ b/backend/internal/handler/dto/settings.go @@ -242,6 +242,7 @@ type SystemSettings struct { PaymentEnabledTypes []string `json:"payment_enabled_types"` PaymentBalanceDisabled bool `json:"payment_balance_disabled"` PaymentBalanceRechargeMultiplier float64 `json:"payment_balance_recharge_multiplier"` + PaymentSubscriptionUSDToCNYRate float64 `json:"payment_subscription_usd_to_cny_rate"` PaymentRechargeFeeRate float64 `json:"payment_recharge_fee_rate"` PaymentLoadBalanceStrat string `json:"payment_load_balance_strategy"` PaymentProductNamePrefix string `json:"payment_product_name_prefix"` diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go index 3c705ed4b2..286d2d5459 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -64,6 +64,8 @@ type APIKey struct { ExpiresAt *time.Time `json:"expires_at"` // Expiration time (nil = never expires) CreatedAt time.Time `json:"created_at"` UpdatedAt time.Time `json:"updated_at"` + // CurrentConcurrency is the real-time active request count for this API key. + CurrentConcurrency int `json:"current_concurrency"` // Rate limit fields RateLimit5h float64 `json:"rate_limit_5h"` diff --git a/backend/internal/handler/gateway_helper.go b/backend/internal/handler/gateway_helper.go index b948ac8fc7..48110da93f 100644 --- a/backend/internal/handler/gateway_helper.go +++ b/backend/internal/handler/gateway_helper.go @@ -10,6 +10,7 @@ import ( "sync" "time" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" @@ -211,6 +212,14 @@ func (h *ConcurrencyHelper) TryAcquireUserSlot(ctx context.Context, userID int64 return result.ReleaseFunc, true, nil } +func (h *ConcurrencyHelper) TryAcquireUserSlotForAPIKey(ctx context.Context, userID int64, maxConcurrency int, apiKeyID int64) (func(), bool, error) { + releaseFunc, acquired, err := h.TryAcquireUserSlot(ctx, userID, maxConcurrency) + if err != nil || !acquired { + return releaseFunc, acquired, err + } + return h.withAPIKeySlot(ctx, apiKeyID, releaseFunc), true, nil +} + // TryAcquireAccountSlot 尝试立即获取账号并发槽位。 // 返回值: (releaseFunc, acquired, error) func (h *ConcurrencyHelper) TryAcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int) (func(), bool, error) { @@ -241,7 +250,7 @@ func (h *ConcurrencyHelper) acquireUserSlotWithWaitTimeout(c *gin.Context, userI } if acquired { - return releaseFunc, nil + return h.withAPIKeySlotFromGin(c, releaseFunc), nil } queueLimit := service.CalculateMaxWait(maxConcurrency) - maxConcurrency @@ -258,7 +267,37 @@ func (h *ConcurrencyHelper) acquireUserSlotWithWaitTimeout(c *gin.Context, userI defer h.DecrementWaitCount(ctx, userID) // Need to wait - handle streaming ping if needed - return h.waitForSlotWithPingTimeout(c, "user", userID, maxConcurrency, timeout, isStream, streamStarted, false) + releaseFunc, err = h.waitForSlotWithPingTimeout(c, "user", userID, maxConcurrency, timeout, isStream, streamStarted, false) + if err != nil { + return nil, err + } + return h.withAPIKeySlotFromGin(c, releaseFunc), nil +} + +func (h *ConcurrencyHelper) withAPIKeySlotFromGin(c *gin.Context, releaseFunc func()) func() { + if c == nil { + return releaseFunc + } + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + if !ok || apiKey == nil { + return releaseFunc + } + return h.withAPIKeySlot(c.Request.Context(), apiKey.ID, releaseFunc) +} + +func (h *ConcurrencyHelper) withAPIKeySlot(ctx context.Context, apiKeyID int64, releaseFunc func()) func() { + if h == nil || h.concurrencyService == nil || apiKeyID <= 0 { + return releaseFunc + } + apiKeyReleaseFunc := h.concurrencyService.TrackAPIKeySlot(ctx, apiKeyID) + return func() { + if releaseFunc != nil { + releaseFunc() + } + if apiKeyReleaseFunc != nil { + apiKeyReleaseFunc() + } + } } // AcquireAccountSlotWithWait acquires an account concurrency slot, waiting if necessary. diff --git a/backend/internal/handler/gateway_helper_hotpath_test.go b/backend/internal/handler/gateway_helper_hotpath_test.go index fb17481f1c..5e0697f083 100644 --- a/backend/internal/handler/gateway_helper_hotpath_test.go +++ b/backend/internal/handler/gateway_helper_hotpath_test.go @@ -9,6 +9,7 @@ import ( "testing" "time" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" @@ -29,6 +30,9 @@ type helperConcurrencyCacheStub struct { waitDecrementCalls int waitMaxWait int waitIncrementHook func() + apiKeyTrackCalls int + apiKeyReleaseCalls int + apiKeyTrackIDs []int64 } func (s *helperConcurrencyCacheStub) AcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int, requestID string) (bool, error) { @@ -97,6 +101,29 @@ func (s *helperConcurrencyCacheStub) GetUserConcurrency(ctx context.Context, use return 0, nil } +func (s *helperConcurrencyCacheStub) TrackAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error { + s.mu.Lock() + defer s.mu.Unlock() + s.apiKeyTrackCalls++ + s.apiKeyTrackIDs = append(s.apiKeyTrackIDs, apiKeyID) + return nil +} + +func (s *helperConcurrencyCacheStub) ReleaseAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error { + s.mu.Lock() + defer s.mu.Unlock() + s.apiKeyReleaseCalls++ + return nil +} + +func (s *helperConcurrencyCacheStub) GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) { + out := make(map[int64]int, len(apiKeyIDs)) + for _, apiKeyID := range apiKeyIDs { + out[apiKeyID] = 0 + } + return out, nil +} + func (s *helperConcurrencyCacheStub) IncrementWaitCount(ctx context.Context, userID int64, maxWait int) (bool, error) { s.mu.Lock() s.waitIncrementCalls++ @@ -270,6 +297,48 @@ func TestAcquireUserSlotWithWait_ImmediateAcquireSkipsWaitQueue(t *testing.T) { require.Equal(t, 1, cache.userReleaseCalls) } +func TestAcquireUserSlotWithWait_TracksAPIKeySlot(t *testing.T) { + cache := &helperConcurrencyCacheStub{ + userSeq: []bool{true}, + } + concurrency := service.NewConcurrencyService(cache) + helper := NewConcurrencyHelper(concurrency, SSEPingFormatNone, 5*time.Millisecond) + c, _ := newHelperTestContext(http.MethodPost, "/v1/messages") + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ID: 77}) + streamStarted := false + + release, err := helper.acquireUserSlotWithWaitTimeout(c, 202, 3, time.Second, false, &streamStarted) + require.NoError(t, err) + require.NotNil(t, release) + require.Equal(t, 1, cache.apiKeyTrackCalls) + require.Equal(t, []int64{77}, cache.apiKeyTrackIDs) + + release() + + require.Equal(t, 1, cache.userReleaseCalls) + require.Equal(t, 1, cache.apiKeyReleaseCalls) +} + +func TestTryAcquireUserSlotForAPIKey_TracksAPIKeySlot(t *testing.T) { + cache := &helperConcurrencyCacheStub{ + userSeq: []bool{true}, + } + concurrency := service.NewConcurrencyService(cache) + helper := NewConcurrencyHelper(concurrency, SSEPingFormatNone, 5*time.Millisecond) + + release, acquired, err := helper.TryAcquireUserSlotForAPIKey(context.Background(), 202, 3, 77) + require.NoError(t, err) + require.True(t, acquired) + require.NotNil(t, release) + require.Equal(t, 1, cache.apiKeyTrackCalls) + require.Equal(t, []int64{77}, cache.apiKeyTrackIDs) + + release() + + require.Equal(t, 1, cache.userReleaseCalls) + require.Equal(t, 1, cache.apiKeyReleaseCalls) +} + func TestAcquireUserSlotWithWait_WaitSuccessDecrementsBeforeReturn(t *testing.T) { cache := &helperConcurrencyCacheStub{ userSeq: []bool{false, true}, diff --git a/backend/internal/handler/grok_media.go b/backend/internal/handler/grok_media.go index 8e236ea49f..4fd1411b23 100644 --- a/backend/internal/handler/grok_media.go +++ b/backend/internal/handler/grok_media.go @@ -174,6 +174,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. service.OpenAIUpstreamTransportHTTPSSE, "", false, + false, service.PlatformGrok, ) if err != nil { diff --git a/backend/internal/handler/openai_chat_completions.go b/backend/internal/handler/openai_chat_completions.go index ca43a2cff3..baff1dcbd6 100644 --- a/backend/internal/handler/openai_chat_completions.go +++ b/backend/internal/handler/openai_chat_completions.go @@ -145,6 +145,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, false, + false, requestPlatform, ) if err != nil { diff --git a/backend/internal/handler/openai_embeddings.go b/backend/internal/handler/openai_embeddings.go index a80c7f7d96..8be533c723 100644 --- a/backend/internal/handler/openai_embeddings.go +++ b/backend/internal/handler/openai_embeddings.go @@ -117,6 +117,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) { service.OpenAIUpstreamTransportHTTPSSE, service.OpenAIEndpointCapabilityEmbeddings, false, + false, ) if err != nil { reqLog.Warn("openai_embeddings.account_select_failed", diff --git a/backend/internal/handler/openai_gateway_count_tokens.go b/backend/internal/handler/openai_gateway_count_tokens.go index ec530e8ab6..fc9c4d5df7 100644 --- a/backend/internal/handler/openai_gateway_count_tokens.go +++ b/backend/internal/handler/openai_gateway_count_tokens.go @@ -110,6 +110,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) { service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, false, + false, openAICompatibleRequestPlatform(apiKey), ) service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds()) diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index ccafde7b02..7f097afa4b 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -350,6 +350,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, requireCompact, + false, requestPlatform, ) if err != nil { @@ -783,6 +784,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, false, + false, requestPlatform, ) if err != nil { @@ -1266,8 +1268,8 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, "previous_response_id must be a response.id (resp_*), not a message id") return } - firstMessageToolContext := service.ValidateFunctionCallOutputContextBytes(firstMessage) - previousResponseCanMove := !firstMessageToolContext.HasFunctionCallOutput || firstMessageToolContext.HasToolCallContext + firstMessageToolCoverage := service.AnalyzeToolCallOutputContextCoverageBytes(firstMessage) + previousResponseCanMove := !firstMessageToolCoverage.HasFunctionCallOutput || firstMessageToolCoverage.ContextCoversAllCallIDs reqLog = reqLog.With( zap.Bool("ws_ingress", true), zap.String("model", reqModel), @@ -1320,7 +1322,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { // 必须尽早注册,确保任何 early return 都能释放已获取的并发槽位。 defer releaseTurnSlots() - userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlot(ctx, subject.UserID, subject.Concurrency) + userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlotForAPIKey(ctx, subject.UserID, subject.Concurrency, apiKey.ID) if err != nil { reqLog.Warn("openai.websocket_user_slot_acquire_failed", zap.Error(err)) closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "failed to acquire user concurrency slot") @@ -1335,7 +1337,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { if currentUserRelease != nil { return true } - userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlot(ctx, subject.UserID, subject.Concurrency) + userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlotForAPIKey(ctx, subject.UserID, subject.Concurrency, apiKey.ID) if err != nil { reqLog.Warn("openai.websocket_user_slot_reacquire_failed", zap.Error(err)) closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "failed to acquire user concurrency slot") @@ -1383,13 +1385,8 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { requiredTransport, service.OpenAIEndpointCapabilityChatCompletions, false, + previousResponseCanMove, requestPlatform, - func() string { - if previousResponseCanMove { - return "previous_response_can_move" - } - return "" - }(), ) if err != nil { reqLog.Warn("openai.websocket_account_select_failed", @@ -1492,7 +1489,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { // 防御式清理:避免异常路径下旧槽位覆盖导致泄漏。 releaseTurnSlots() // 非首轮 turn 需要重新抢占并发槽位,避免长连接空闲占槽。 - userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlot(ctx, subject.UserID, subject.Concurrency) + userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlotForAPIKey(ctx, subject.UserID, subject.Concurrency, apiKey.ID) if err != nil { return service.NewOpenAIWSClientCloseError(coderws.StatusInternalError, "failed to acquire user concurrency slot", err) } diff --git a/backend/internal/handler/payment_handler.go b/backend/internal/handler/payment_handler.go index 7cdf73cd3d..a267d73724 100644 --- a/backend/internal/handler/payment_handler.go +++ b/backend/internal/handler/payment_handler.go @@ -150,6 +150,7 @@ func (h *PaymentHandler) GetCheckoutInfo(c *gin.Context) { Plans: planList, BalanceDisabled: cfg.BalanceDisabled, BalanceRechargeMultiplier: cfg.BalanceRechargeMultiplier, + SubscriptionUSDToCNYRate: cfg.SubscriptionUSDToCNYRate, RechargeFeeRate: cfg.RechargeFeeRate, HelpText: cfg.HelpText, HelpImageURL: cfg.HelpImageURL, @@ -165,6 +166,7 @@ type checkoutInfoResponse struct { Plans []checkoutPlan `json:"plans"` BalanceDisabled bool `json:"balance_disabled"` BalanceRechargeMultiplier float64 `json:"balance_recharge_multiplier"` + SubscriptionUSDToCNYRate float64 `json:"subscription_usd_to_cny_rate"` RechargeFeeRate float64 `json:"recharge_fee_rate"` HelpText string `json:"help_text"` HelpImageURL string `json:"help_image_url"` diff --git a/backend/internal/payment/provider/easypay.go b/backend/internal/payment/provider/easypay.go index 32d6b7bebf..f1c17427ad 100644 --- a/backend/internal/payment/provider/easypay.go +++ b/backend/internal/payment/provider/easypay.go @@ -39,6 +39,12 @@ type EasyPay struct { httpClient *http.Client } +type easyPayCustomMethod struct { + Type string `json:"type"` + UpstreamType string `json:"upstreamType"` + DisplayName string `json:"displayName"` +} + // NewEasyPay creates a new EasyPay provider. // config keys: pid, pkey, apiBase, notifyUrl, returnUrl, cid, cidAlipay, cidWxpay func NewEasyPay(instanceID string, config map[string]string) (*EasyPay, error) { @@ -95,7 +101,13 @@ func (e *EasyPay) apiBase() string { func (e *EasyPay) Name() string { return "EasyPay" } func (e *EasyPay) ProviderKey() string { return payment.TypeEasyPay } func (e *EasyPay) SupportedTypes() []payment.PaymentType { - return []payment.PaymentType{payment.TypeAlipay, payment.TypeWxpay} + types := []payment.PaymentType{payment.TypeAlipay, payment.TypeWxpay} + for _, method := range e.customMethods() { + if method.Type != "" { + types = append(types, method.Type) + } + } + return types } func (e *EasyPay) MerchantIdentityMetadata() map[string]string { @@ -124,13 +136,14 @@ func (e *EasyPay) CreatePayment(ctx context.Context, req payment.CreatePaymentRe // TradeNo is empty; it arrives via the notify callback after payment. func (e *EasyPay) createRedirectPayment(req payment.CreatePaymentRequest) (*payment.CreatePaymentResponse, error) { notifyURL, returnURL := e.resolveURLs(req) + paymentType := e.upstreamPaymentType(req.PaymentType) params := map[string]string{ - "pid": e.config["pid"], "type": req.PaymentType, + "pid": e.config["pid"], "type": paymentType, "out_trade_no": req.OrderID, "notify_url": notifyURL, "return_url": returnURL, "name": req.Subject, "money": req.Amount, } - if cid := e.resolveCID(req.PaymentType); cid != "" { + if cid := e.resolveCID(paymentType); cid != "" { params["cid"] = cid } if req.IsMobile { @@ -150,13 +163,14 @@ func (e *EasyPay) createRedirectPayment(req payment.CreatePaymentRequest) (*paym // createAPIPayment calls mapi.php to get payurl/qrcode (existing behavior). func (e *EasyPay) createAPIPayment(ctx context.Context, req payment.CreatePaymentRequest) (*payment.CreatePaymentResponse, error) { notifyURL, returnURL := e.resolveURLs(req) + paymentType := e.upstreamPaymentType(req.PaymentType) params := map[string]string{ - "pid": e.config["pid"], "type": req.PaymentType, + "pid": e.config["pid"], "type": paymentType, "out_trade_no": req.OrderID, "notify_url": notifyURL, "return_url": returnURL, "name": req.Subject, "money": req.Amount, "clientip": req.ClientIP, } - if cid := e.resolveCID(req.PaymentType); cid != "" { + if cid := e.resolveCID(paymentType); cid != "" { params["cid"] = cid } if req.IsMobile { @@ -204,6 +218,41 @@ func (e *EasyPay) resolveURLs(req payment.CreatePaymentRequest) (string, string) return notifyURL, returnURL } +func (e *EasyPay) customMethods() []easyPayCustomMethod { + if e == nil { + return nil + } + raw := strings.TrimSpace(e.config["customMethods"]) + if raw == "" { + return nil + } + var methods []easyPayCustomMethod + if err := json.Unmarshal([]byte(raw), &methods); err != nil { + return nil + } + result := make([]easyPayCustomMethod, 0, len(methods)) + for _, method := range methods { + method.Type = strings.TrimSpace(method.Type) + method.UpstreamType = strings.TrimSpace(method.UpstreamType) + method.DisplayName = strings.TrimSpace(method.DisplayName) + if method.Type == "" || method.UpstreamType == "" { + continue + } + result = append(result, method) + } + return result +} + +func (e *EasyPay) upstreamPaymentType(paymentType string) string { + paymentType = strings.TrimSpace(paymentType) + for _, method := range e.customMethods() { + if paymentType == method.Type { + return method.UpstreamType + } + } + return paymentType +} + func (e *EasyPay) QueryOrder(ctx context.Context, tradeNo string) (*payment.QueryOrderResponse, error) { params := map[string]string{ "act": "order", "pid": e.config["pid"], diff --git a/backend/internal/payment/provider/easypay_refund_test.go b/backend/internal/payment/provider/easypay_refund_test.go index 9e0e4942c2..3b76329870 100644 --- a/backend/internal/payment/provider/easypay_refund_test.go +++ b/backend/internal/payment/provider/easypay_refund_test.go @@ -179,6 +179,102 @@ func TestEasyPayRefundResponseErrors(t *testing.T) { } } +func TestEasyPayCustomMethodsUseConfiguredUpstreamType(t *testing.T) { + t.Parallel() + + provider, err := NewEasyPay("test-instance", map[string]string{ + "pid": "pid-1", + "pkey": "pkey-1", + "apiBase": "https://pay.example.com", + "notifyUrl": "https://example.com/notify", + "returnUrl": "https://example.com/return", + "paymentMode": paymentModePopup, + "customMethods": `[{"type":"ldc","upstreamType":"epay","displayName":"LDC"},{"type":"usdt_trc20","upstreamType":"usdt","displayName":"USDT-TRC20"}]`, + }) + if err != nil { + t.Fatalf("NewEasyPay: %v", err) + } + + resp, err := provider.CreatePayment(context.Background(), payment.CreatePaymentRequest{ + OrderID: "sub2-custom-1", + Amount: "1.00", + PaymentType: "usdt_trc20", + Subject: "Custom EasyPay", + }) + if err != nil { + t.Fatalf("CreatePayment: %v", err) + } + payURL, err := url.Parse(resp.PayURL) + if err != nil { + t.Fatalf("parse pay url: %v", err) + } + if got := payURL.Query().Get("type"); got != "usdt" { + t.Fatalf("pay url type = %q, want usdt (%s)", got, resp.PayURL) + } +} + +func TestEasyPayCustomMethodsResolveCIDFromConfiguredUpstreamType(t *testing.T) { + t.Parallel() + + provider, err := NewEasyPay("test-instance", map[string]string{ + "pid": "pid-1", + "pkey": "pkey-1", + "apiBase": "https://pay.example.com", + "notifyUrl": "https://example.com/notify", + "returnUrl": "https://example.com/return", + "paymentMode": paymentModePopup, + "cidAlipay": "cid-alipay", + "cidWxpay": "cid-wxpay", + "customMethods": `[{"type":"ldc","upstreamType":"alipay","displayName":"LDC"}]`, + }) + if err != nil { + t.Fatalf("NewEasyPay: %v", err) + } + + resp, err := provider.CreatePayment(context.Background(), payment.CreatePaymentRequest{ + OrderID: "sub2-custom-cid", + Amount: "1.00", + PaymentType: "ldc", + Subject: "Custom EasyPay CID", + }) + if err != nil { + t.Fatalf("CreatePayment: %v", err) + } + payURL, err := url.Parse(resp.PayURL) + if err != nil { + t.Fatalf("parse pay url: %v", err) + } + if got := payURL.Query().Get("type"); got != "alipay" { + t.Fatalf("pay url type = %q, want alipay (%s)", got, resp.PayURL) + } + if got := payURL.Query().Get("cid"); got != "cid-alipay" { + t.Fatalf("pay url cid = %q, want cid-alipay (%s)", got, resp.PayURL) + } +} + +func TestEasyPaySupportedTypesIncludeCustomMethods(t *testing.T) { + t.Parallel() + + provider, err := NewEasyPay("test-instance", map[string]string{ + "pid": "pid-1", + "pkey": "pkey-1", + "apiBase": "https://pay.example.com", + "notifyUrl": "https://example.com/notify", + "returnUrl": "https://example.com/return", + "customMethods": `[{"type":"ldc","upstreamType":"epay","displayName":"LDC"},{"type":"usdt_trc20","upstreamType":"usdt","displayName":"USDT-TRC20"}]`, + }) + if err != nil { + t.Fatalf("NewEasyPay: %v", err) + } + + got := strings.Join(provider.SupportedTypes(), ",") + for _, want := range []string{"alipay", "wxpay", "ldc", "usdt_trc20"} { + if !strings.Contains(got, want) { + t.Fatalf("SupportedTypes() = %q, want it to include %q", got, want) + } + } +} + func newTestEasyPay(t *testing.T, apiBase string) *EasyPay { t.Helper() diff --git a/backend/internal/pkg/openai/constants.go b/backend/internal/pkg/openai/constants.go index f658cf0675..c9d391df4e 100644 --- a/backend/internal/pkg/openai/constants.go +++ b/backend/internal/pkg/openai/constants.go @@ -18,6 +18,9 @@ type Model struct { // DefaultModels OpenAI models list var DefaultModels = []Model{ + {ID: "gpt-5.6-sol", Object: "model", Created: 1780876800, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.6 Sol"}, + {ID: "gpt-5.6-terra", Object: "model", Created: 1780876800, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.6 Terra"}, + {ID: "gpt-5.6-luna", Object: "model", Created: 1780876800, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.6 Luna"}, {ID: "gpt-5.5", Object: "model", Created: 1776873600, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.5"}, {ID: "gpt-5.4", Object: "model", Created: 1738368000, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.4"}, {ID: "gpt-5.4-mini", Object: "model", Created: 1738368000, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.4 Mini"}, diff --git a/backend/internal/repository/concurrency_cache.go b/backend/internal/repository/concurrency_cache.go index b530c30b90..f45211f260 100644 --- a/backend/internal/repository/concurrency_cache.go +++ b/backend/internal/repository/concurrency_cache.go @@ -27,6 +27,8 @@ const ( accountSlotKeyPrefix = "concurrency:account:" // 格式: concurrency:user:{userID} userSlotKeyPrefix = "concurrency:user:" + // 格式: concurrency:api_key:{apiKeyID} + apiKeySlotKeyPrefix = "concurrency:api_key:" // 等待队列计数器格式: concurrency:wait:{userID} waitQueueKeyPrefix = "concurrency:wait:" // 账号级等待队列计数器格式: wait:account:{accountID} @@ -99,6 +101,28 @@ var ( return redis.call('ZCARD', key) `) + // trackSlotScript 记录 stats-only 槽位,不做并发上限判断。 + // KEYS[1] = 有序集合键 + // ARGV[1] = TTL(秒) + // ARGV[2] = requestID + trackSlotScript = redis.NewScript(` + -- Redis 3.2-4.x compat: opt into effects replication so redis.call('TIME') + -- replicates correctly. No-op on Redis 5.0+ (effects replication is default). + redis.replicate_commands() + local key = KEYS[1] + local ttl = tonumber(ARGV[1]) + local requestID = ARGV[2] + + local timeResult = redis.call('TIME') + local now = tonumber(timeResult[1]) + local expireBefore = now - ttl + + redis.call('ZREMRANGEBYSCORE', key, '-inf', expireBefore) + redis.call('ZADD', key, now, requestID) + redis.call('EXPIRE', key, ttl) + return 1 + `) + // incrementWaitScript - refreshes TTL on each increment to keep queue depth accurate // KEYS[1] = wait queue key // ARGV[1] = maxWait @@ -254,6 +278,10 @@ func userSlotKey(userID int64) string { return fmt.Sprintf("%s%d", userSlotKeyPrefix, userID) } +func apiKeySlotKey(apiKeyID int64) string { + return fmt.Sprintf("%s%d", apiKeySlotKeyPrefix, apiKeyID) +} + func waitQueueKey(userID int64) string { return fmt.Sprintf("%s%d", waitQueueKeyPrefix, userID) } @@ -353,6 +381,54 @@ func (c *concurrencyCache) GetUserConcurrency(ctx context.Context, userID int64) return result, nil } +func (c *concurrencyCache) TrackAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error { + key := apiKeySlotKey(apiKeyID) + _, err := trackSlotScript.Run(ctx, c.rdb, []string{key}, c.slotTTLSeconds, requestID).Result() + return err +} + +func (c *concurrencyCache) ReleaseAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error { + key := apiKeySlotKey(apiKeyID) + return c.rdb.ZRem(ctx, key, requestID).Err() +} + +func (c *concurrencyCache) GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) { + if len(apiKeyIDs) == 0 { + return map[int64]int{}, nil + } + + now, err := c.rdb.Time(ctx).Result() + if err != nil { + return nil, fmt.Errorf("redis TIME: %w", err) + } + cutoffTime := now.Unix() - int64(c.slotTTLSeconds) + + pipe := c.rdb.Pipeline() + type apiKeyCmd struct { + apiKeyID int64 + zcardCmd *redis.IntCmd + } + cmds := make([]apiKeyCmd, 0, len(apiKeyIDs)) + for _, apiKeyID := range apiKeyIDs { + slotKey := apiKeySlotKeyPrefix + strconv.FormatInt(apiKeyID, 10) + pipe.ZRemRangeByScore(ctx, slotKey, "-inf", strconv.FormatInt(cutoffTime, 10)) + cmds = append(cmds, apiKeyCmd{ + apiKeyID: apiKeyID, + zcardCmd: pipe.ZCard(ctx, slotKey), + }) + } + + if _, err := pipe.Exec(ctx); err != nil && !errors.Is(err, redis.Nil) { + return nil, fmt.Errorf("pipeline exec: %w", err) + } + + result := make(map[int64]int, len(apiKeyIDs)) + for _, cmd := range cmds { + result[cmd.apiKeyID] = int(cmd.zcardCmd.Val()) + } + return result, nil +} + // Wait queue operations func (c *concurrencyCache) IncrementWaitCount(ctx context.Context, userID int64, maxWait int) (bool, error) { @@ -536,7 +612,7 @@ func (c *concurrencyCache) CleanupStaleProcessSlots(ctx context.Context, activeR } // 1. 清理有序集合中非当前进程前缀的成员 - slotPatterns := []string{accountSlotKeyPrefix + "*", userSlotKeyPrefix + "*"} + slotPatterns := []string{accountSlotKeyPrefix + "*", userSlotKeyPrefix + "*", apiKeySlotKeyPrefix + "*"} for _, pattern := range slotPatterns { if err := c.cleanupSlotsByPattern(ctx, pattern, activeRequestPrefix); err != nil { return err diff --git a/backend/internal/repository/concurrency_cache_integration_test.go b/backend/internal/repository/concurrency_cache_integration_test.go index 4aa172b252..8b3e1bc359 100644 --- a/backend/internal/repository/concurrency_cache_integration_test.go +++ b/backend/internal/repository/concurrency_cache_integration_test.go @@ -3,6 +3,7 @@ package repository import ( + "context" "errors" "fmt" "testing" @@ -34,6 +35,18 @@ func (s *ConcurrencyCacheSuite) SetupTest() { s.cache = NewConcurrencyCache(s.rdb, testSlotTTLMinutes, int(testSlotTTL.Seconds())) } +type apiKeyConcurrencyCacheForTest interface { + TrackAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error + ReleaseAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error + GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) +} + +func (s *ConcurrencyCacheSuite) apiKeyConcurrencyCache() apiKeyConcurrencyCacheForTest { + cache, ok := s.cache.(apiKeyConcurrencyCacheForTest) + require.True(s.T(), ok) + return cache +} + func (s *ConcurrencyCacheSuite) TestAccountSlot_AcquireAndRelease() { accountID := int64(10) reqID1, reqID2, reqID3 := "req1", "req2", "req3" @@ -160,6 +173,34 @@ func (s *ConcurrencyCacheSuite) TestUserSlot_TTL() { s.AssertTTLWithin(ttl, 1*time.Second, testSlotTTL) } +func (s *ConcurrencyCacheSuite) TestAPIKeySlot_TrackReleaseAndBatchCount() { + cache := s.apiKeyConcurrencyCache() + apiKeyID := int64(300) + emptyAPIKeyID := int64(301) + slotKey := fmt.Sprintf("%s%d", apiKeySlotKeyPrefix, apiKeyID) + + require.NoError(s.T(), cache.TrackAPIKeySlot(s.ctx, apiKeyID, "req1")) + require.NoError(s.T(), cache.TrackAPIKeySlot(s.ctx, apiKeyID, "req2")) + + counts, err := cache.GetAPIKeyConcurrencyBatch(s.ctx, []int64{apiKeyID, emptyAPIKeyID}) + require.NoError(s.T(), err) + require.Equal(s.T(), map[int64]int{apiKeyID: 2, emptyAPIKeyID: 0}, counts) + + ttl, err := s.rdb.TTL(s.ctx, slotKey).Result() + require.NoError(s.T(), err, "TTL") + s.AssertTTLWithin(ttl, 1*time.Second, testSlotTTL) + + require.NoError(s.T(), cache.ReleaseAPIKeySlot(s.ctx, apiKeyID, "req1")) + counts, err = cache.GetAPIKeyConcurrencyBatch(s.ctx, []int64{apiKeyID}) + require.NoError(s.T(), err) + require.Equal(s.T(), 1, counts[apiKeyID]) + + require.NoError(s.T(), cache.ReleaseAPIKeySlot(s.ctx, apiKeyID, "req2")) + counts, err = cache.GetAPIKeyConcurrencyBatch(s.ctx, []int64{apiKeyID}) + require.NoError(s.T(), err) + require.Equal(s.T(), 0, counts[apiKeyID]) +} + func (s *ConcurrencyCacheSuite) TestWaitQueue_IncrementAndDecrement() { userID := int64(20) waitKey := fmt.Sprintf("%s%d", waitQueueKeyPrefix, userID) @@ -254,8 +295,10 @@ func (s *ConcurrencyCacheSuite) TestAccountWaitQueue_IncrementAndDecrement() { func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots() { accountID := int64(901) userID := int64(902) + apiKeyID := int64(903) accountKey := fmt.Sprintf("%s%d", accountSlotKeyPrefix, accountID) userKey := fmt.Sprintf("%s%d", userSlotKeyPrefix, userID) + apiKeyKey := fmt.Sprintf("%s%d", apiKeySlotKeyPrefix, apiKeyID) userWaitKey := fmt.Sprintf("%s%d", waitQueueKeyPrefix, userID) accountWaitKey := fmt.Sprintf("%s%d", accountWaitKeyPrefix, accountID) @@ -268,6 +311,10 @@ func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots() { redis.Z{Score: float64(now), Member: "oldproc-2"}, redis.Z{Score: float64(now), Member: "keep-2"}, ).Err()) + require.NoError(s.T(), s.rdb.ZAdd(s.ctx, apiKeyKey, + redis.Z{Score: float64(now), Member: "oldproc-3"}, + redis.Z{Score: float64(now), Member: "keep-3"}, + ).Err()) require.NoError(s.T(), s.rdb.Set(s.ctx, userWaitKey, 3, time.Minute).Err()) require.NoError(s.T(), s.rdb.Set(s.ctx, accountWaitKey, 2, time.Minute).Err()) @@ -281,6 +328,10 @@ func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots() { require.NoError(s.T(), err) require.Equal(s.T(), []string{"keep-2"}, userMembers) + apiKeyMembers, err := s.rdb.ZRange(s.ctx, apiKeyKey, 0, -1).Result() + require.NoError(s.T(), err) + require.Equal(s.T(), []string{"keep-3"}, apiKeyMembers) + _, err = s.rdb.Get(s.ctx, userWaitKey).Result() require.True(s.T(), errors.Is(err, redis.Nil)) diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index ee2b320177..278e654834 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -234,6 +234,7 @@ func TestAPIContracts(t *testing.T) { "ip_whitelist": null, "ip_blacklist": null, "last_used_at": null, + "current_concurrency": 0, "quota": 0, "quota_used": 0, "rate_limit_5h": 0, @@ -283,6 +284,7 @@ func TestAPIContracts(t *testing.T) { "ip_whitelist": null, "ip_blacklist": null, "last_used_at": null, + "current_concurrency": 0, "quota": 0, "quota_used": 0, "rate_limit_5h": 0, @@ -903,6 +905,7 @@ func TestAPIContracts(t *testing.T) { "payment_max_pending_orders": 0, "payment_balance_disabled": false, "payment_balance_recharge_multiplier": 0, + "payment_subscription_usd_to_cny_rate": 0, "payment_recharge_fee_rate": 0, "payment_load_balance_strategy": "", "payment_product_name_prefix": "", @@ -1173,6 +1176,7 @@ func TestAPIContracts(t *testing.T) { "payment_enabled_types": null, "payment_balance_disabled": false, "payment_balance_recharge_multiplier": 0, + "payment_subscription_usd_to_cny_rate": 0, "payment_recharge_fee_rate": 0, "payment_load_balance_strategy": "", "payment_product_name_prefix": "", @@ -1740,6 +1744,10 @@ func (s *stubAccountRepo) List(ctx context.Context, params pagination.Pagination return nil, nil, errors.New("not implemented") } +func (s *stubAccountRepo) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]service.Account, error) { + return nil, nil +} + func (s *stubAccountRepo) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]service.Account, *pagination.PaginationResult, error) { return nil, nil, errors.New("not implemented") } diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index 8db5805e9d..ae25bf387d 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -70,6 +70,14 @@ type Account struct { modelMappingCacheRawPtr uintptr modelMappingCacheRawLen int modelMappingCacheRawSig uint64 + + // header_overrides 热路径缓存(非持久化字段,同 model_mapping 缓存先例) + headerOverrideCache map[string]string + headerOverrideCacheReady bool + headerOverrideCacheCredentialsPtr uintptr + headerOverrideCacheRawPtr uintptr + headerOverrideCacheRawLen int + headerOverrideCacheRawSig uint64 } type OpenAIEndpointCapability string diff --git a/backend/internal/service/account_header_override.go b/backend/internal/service/account_header_override.go new file mode 100644 index 0000000000..8882bbef91 --- /dev/null +++ b/backend/internal/service/account_header_override.go @@ -0,0 +1,280 @@ +package service + +import ( + "net/http" + "strings" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + + "golang.org/x/net/http/httpguts" +) + +// 请求头覆写(header override):仅对 Anthropic / OpenAI 平台的 api_key 账号生效。 +// 管理员在账号上配置一组 header name -> value,转发到上游前用配置值覆盖同名请求头 +// (匹配不区分大小写);value 为空的条目视为"未填写",不参与覆盖。 +const ( + credKeyHeaderOverrideEnabled = "header_override_enabled" + credKeyHeaderOverrides = "header_overrides" + + maxHeaderOverrideEntries = 64 + maxHeaderOverrideNameLength = 200 + maxHeaderOverrideValueLength = 8192 +) + +// headerOverrideBlockedNames 禁止覆写的请求头(小写)。 +// - 连接控制/逐跳头:由 HTTP 栈管理,覆写会破坏请求传输; +// - host/content-length:由 Go 的 Request.Host / ContentLength 字段管理,header 覆写不生效或产生冲突; +// - content-type:承载报文框架信息(multipart boundary 为每请求随机值),静态覆写必然与 body 不匹配; +// - authorization/x-api-key/cookie 等:上游认证头由账号凭据统一注入,禁止通过覆写篡改或重新引入; +// - accept-encoding:强制压缩会破坏网关对上游流式响应(SSE/usage)的解析; +// - sec-websocket-*:WebSocket 握手头由拨号器管理(OpenAI WS 模式); +// - session_id/x-claude-code-session-id 等:逐请求会话隔离头,固定值会造成会话串扰。 +var headerOverrideBlockedNames = map[string]struct{}{ + "host": {}, + "content-length": {}, + "content-type": {}, + "transfer-encoding": {}, + "connection": {}, + "keep-alive": {}, + "proxy-authenticate": {}, + "proxy-authorization": {}, + "proxy-connection": {}, + "te": {}, + "trailer": {}, + "upgrade": {}, + "authorization": {}, + "x-api-key": {}, + "x-goog-api-key": {}, + "cookie": {}, + "accept-encoding": {}, + "sec-websocket-key": {}, + "sec-websocket-version": {}, + "sec-websocket-extensions": {}, + "sec-websocket-protocol": {}, + "sec-websocket-accept": {}, + "session_id": {}, + "conversation_id": {}, + "x-codex-turn-state": {}, + "x-codex-turn-metadata": {}, + "chatgpt-account-id": {}, + "x-claude-code-session-id": {}, + "x-client-request-id": {}, +} + +func isHeaderOverrideBlockedName(lowerName string) bool { + _, blocked := headerOverrideBlockedNames[lowerName] + return blocked +} + +// IsHeaderOverrideEligible 报告账号类型是否支持请求头覆写。 +// 目前仅开放 Anthropic / OpenAI 两个平台的 api_key 账号。 +func (a *Account) IsHeaderOverrideEligible() bool { + if a == nil || a.Type != AccountTypeAPIKey { + return false + } + return a.Platform == PlatformAnthropic || a.Platform == PlatformOpenAI +} + +// IsHeaderOverrideEnabled 报告账号是否启用了请求头覆写。 +func (a *Account) IsHeaderOverrideEnabled() bool { + if !a.IsHeaderOverrideEligible() || a.Credentials == nil { + return false + } + enabled, ok := a.Credentials[credKeyHeaderOverrideEnabled].(bool) + return ok && enabled +} + +// GetHeaderOverrides 返回生效的请求头覆写表(key 统一小写)。 +// 未启用、不符合平台/类型条件或配置为空时返回 nil。 +// 空 value 的条目(模板占位)与非法/禁止的 header 名会被跳过。 +// 结果带热路径缓存(同 GetModelMapping 先例):同一 credentials 映射在 +// 一次请求 / 一条 WS 会话内的多次调用只做一次解析与校验。 +func (a *Account) GetHeaderOverrides() map[string]string { + if !a.IsHeaderOverrideEnabled() { + return nil + } + rawMapping, rawIsAnyMap := a.Credentials[credKeyHeaderOverrides].(map[string]any) + if !rawIsAnyMap { + // 非 JSON 反序列化产物(如直接注入的 map[string]string):直接解析,不缓存 + return resolveHeaderOverrides(stringMappingFromRaw(a.Credentials[credKeyHeaderOverrides])) + } + + credentialsPtr := mapPtr(a.Credentials) + rawPtr := mapPtr(rawMapping) + rawLen := len(rawMapping) + rawSig := uint64(0) + rawSigReady := false + + if a.headerOverrideCacheReady && + a.headerOverrideCacheCredentialsPtr == credentialsPtr && + a.headerOverrideCacheRawPtr == rawPtr && + a.headerOverrideCacheRawLen == rawLen { + rawSig = modelMappingSignature(rawMapping) + rawSigReady = true + if a.headerOverrideCacheRawSig == rawSig { + return a.headerOverrideCache + } + } + + overrides := resolveHeaderOverrides(stringMappingFromRaw(rawMapping)) + if !rawSigReady { + rawSig = modelMappingSignature(rawMapping) + } + + a.headerOverrideCache = overrides + a.headerOverrideCacheReady = true + a.headerOverrideCacheCredentialsPtr = credentialsPtr + a.headerOverrideCacheRawPtr = rawPtr + a.headerOverrideCacheRawLen = rawLen + a.headerOverrideCacheRawSig = rawSig + return overrides +} + +// resolveHeaderOverrides 解析并防御性过滤原始覆写表:保存路径已做校验, +// 这里兜底未经 Normalize 落库的数据(含名单扩充前保存的旧配置),非法条目直接跳过。 +func resolveHeaderOverrides(raw map[string]string) map[string]string { + if len(raw) == 0 { + return nil + } + result := make(map[string]string, len(raw)) + for name, value := range raw { + lowerName, value, err := normalizeHeaderOverrideEntry(name, value) + if err != nil || lowerName == "" || value == "" { + continue + } + result[lowerName] = value + } + if len(result) == 0 { + return nil + } + return result +} + +// HeaderOverrideValue 返回指定 header(小写名)的生效覆写值。 +// 供转发链路在 header 写入前感知覆写结果(如 anthropic-beta 需要参与 body 净化)。 +func (a *Account) HeaderOverrideValue(lowerName string) (string, bool) { + value, ok := a.GetHeaderOverrides()[lowerName] + return value, ok +} + +// ApplyHeaderOverrides 将账号配置的请求头覆写应用到出站请求头。 +// 对每个覆写条目:先删除所有大小写变体(转发链路会以 wire casing 直接写入 map, +// 可能存在非 canonical key),再按已知 wire casing 写入,避免产生重复头。 +// 账号未启用或不符合条件时为 no-op,可安全地在 OAuth/api_key 共用的构建器中调用。 +func (a *Account) ApplyHeaderOverrides(h http.Header) { + if h == nil { + return + } + overrides := a.GetHeaderOverrides() + if len(overrides) == 0 { + return + } + // 覆写名两两不同(大小写不敏感)且各自只操作同名键,应用顺序不影响结果。 + // 全量 EqualFold 扫描兜底删除任意 casing 的既有键:透传链路可能保留客户端 + // 原始 casing,非 canonical/wire casing 的键 deleteHeaderAllForms 覆盖不到。 + for name, value := range overrides { + for existing := range h { + if strings.EqualFold(existing, name) { + delete(h, existing) + } + } + h[resolveWireCasing(name)] = []string{value} + } +} + +// NormalizeHeaderOverrideCredentials 校验并原地规范化 credentials 中的请求头覆写字段。 +// 供账号创建/更新/批量更新的保存路径调用;credentials 未携带相关字段时为 no-op。 +// 规范化内容:header 名转小写并去除首尾空白,value 去除首尾空白,丢弃名和值均为空的条目。 +func NormalizeHeaderOverrideCredentials(credentials map[string]any) error { + if credentials == nil { + return nil + } + if raw, ok := credentials[credKeyHeaderOverrideEnabled]; ok && raw != nil { + if _, isBool := raw.(bool); !isBool { + return infraerrors.New(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header_override_enabled must be a boolean") + } + } + raw, ok := credentials[credKeyHeaderOverrides] + if !ok || raw == nil { + return nil + } + + var entries map[string]any + switch m := raw.(type) { + case map[string]any: + entries = m + case map[string]string: + entries = make(map[string]any, len(m)) + for k, v := range m { + entries[k] = v + } + default: + return infraerrors.New(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header_overrides must be an object of header name to string value") + } + + if len(entries) > maxHeaderOverrideEntries { + return infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header_overrides supports at most %d entries", maxHeaderOverrideEntries) + } + + normalized := make(map[string]any, len(entries)) + for name, rawValue := range entries { + value, isString := rawValue.(string) + if !isString { + return infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header %q value must be a string", name) + } + lowerName, value, err := normalizeHeaderOverrideEntry(name, value) + if err != nil { + return err + } + if lowerName == "" { + continue // 丢弃完全为空的占位行 + } + if _, dup := normalized[lowerName]; dup { + return infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "duplicate header name %q (matching is case-insensitive)", lowerName) + } + normalized[lowerName] = value + } + credentials[credKeyHeaderOverrides] = normalized + return nil +} + +// normalizeHeaderOverrideEntry 校验并规范化单个覆写条目,保存路径(Normalize,err → 400) +// 与应用路径(resolveHeaderOverrides,err → 跳过)共用同一套规则,避免两处校验漂移。 +// 名和值均为空表示空占位行,返回 ("", "", nil);空 value 的具名条目合法(模板占位)。 +func normalizeHeaderOverrideEntry(name, value string) (string, string, error) { + lowerName := strings.ToLower(strings.TrimSpace(name)) + value = strings.TrimSpace(value) + if lowerName == "" { + if value == "" { + return "", "", nil + } + return "", "", infraerrors.New(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header name must not be empty") + } + if len(lowerName) > maxHeaderOverrideNameLength { + return "", "", infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header name %q exceeds %d characters", lowerName, maxHeaderOverrideNameLength) + } + if !httpguts.ValidHeaderFieldName(lowerName) { + return "", "", infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "invalid header name %q", lowerName) + } + if isHeaderOverrideBlockedName(lowerName) { + return "", "", infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header %q is not allowed to be overridden", lowerName) + } + if len(value) > maxHeaderOverrideValueLength { + return "", "", infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header %q value exceeds %d characters", lowerName, maxHeaderOverrideValueLength) + } + if !httpguts.ValidHeaderFieldValue(value) { + return "", "", infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE", + "header %q has an invalid value", lowerName) + } + return lowerName, value, nil +} diff --git a/backend/internal/service/account_header_override_test.go b/backend/internal/service/account_header_override_test.go new file mode 100644 index 0000000000..c89b5e0587 --- /dev/null +++ b/backend/internal/service/account_header_override_test.go @@ -0,0 +1,339 @@ +//go:build unit + +package service + +import ( + "net/http" + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func headerOverrideTestAccount(platform, accountType string, credentials map[string]any) *Account { + return &Account{ + Platform: platform, + Type: accountType, + Credentials: credentials, + } +} + +func TestIsHeaderOverrideEligible(t *testing.T) { + tests := []struct { + name string + platform string + accType string + want bool + }{ + {"anthropic apikey", PlatformAnthropic, AccountTypeAPIKey, true}, + {"openai apikey", PlatformOpenAI, AccountTypeAPIKey, true}, + {"anthropic oauth", PlatformAnthropic, AccountTypeOAuth, false}, + {"openai oauth", PlatformOpenAI, AccountTypeOAuth, false}, + {"gemini apikey", PlatformGemini, AccountTypeAPIKey, false}, + {"grok apikey", PlatformGrok, AccountTypeAPIKey, false}, + {"antigravity apikey", PlatformAntigravity, AccountTypeAPIKey, false}, + {"anthropic bedrock", PlatformAnthropic, AccountTypeBedrock, false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + acc := headerOverrideTestAccount(tt.platform, tt.accType, nil) + require.Equal(t, tt.want, acc.IsHeaderOverrideEligible()) + }) + } + + var nilAccount *Account + require.False(t, nilAccount.IsHeaderOverrideEligible()) + require.False(t, nilAccount.IsHeaderOverrideEnabled()) + require.Nil(t, nilAccount.GetHeaderOverrides()) +} + +func TestIsHeaderOverrideEnabled(t *testing.T) { + acc := headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: true, + }) + require.True(t, acc.IsHeaderOverrideEnabled()) + + // 未配置 / 非 bool / false 均视为未启用 + require.False(t, headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, nil).IsHeaderOverrideEnabled()) + require.False(t, headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: "true", + }).IsHeaderOverrideEnabled()) + require.False(t, headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: false, + }).IsHeaderOverrideEnabled()) + + // 不符合平台/类型条件时即使配置了 true 也不启用 + require.False(t, headerOverrideTestAccount(PlatformAnthropic, AccountTypeOAuth, map[string]any{ + credKeyHeaderOverrideEnabled: true, + }).IsHeaderOverrideEnabled()) + require.False(t, headerOverrideTestAccount(PlatformGemini, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: true, + }).IsHeaderOverrideEnabled()) +} + +func TestGetHeaderOverrides(t *testing.T) { + acc := headerOverrideTestAccount(PlatformOpenAI, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{ + "User-Agent": "my-agent/1.0", // 大写 key 归一化为小写 + " X-App ": "cli", // 名称去空白 + "x-empty": "", // 空 value(模板占位)跳过 + "authorization": "Bearer leaked", // 禁止覆写的头跳过 + "bad name": "value", // 非法 header 名跳过 + "x-padded": " padded ", // value 去空白 + }, + }) + overrides := acc.GetHeaderOverrides() + require.Equal(t, map[string]string{ + "user-agent": "my-agent/1.0", + "x-app": "cli", + "x-padded": "padded", + }, overrides) + + // 未启用时返回 nil + disabled := headerOverrideTestAccount(PlatformOpenAI, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrides: map[string]any{"user-agent": "x"}, + }) + require.Nil(t, disabled.GetHeaderOverrides()) + + // 启用但全部为空 value 时返回 nil + empty := headerOverrideTestAccount(PlatformOpenAI, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{"user-agent": ""}, + }) + require.Nil(t, empty.GetHeaderOverrides()) + + // 未经 Normalize 落库的超长数据 / WebSocket 握手头在应用时被防御性跳过 + oversizedValue := strings.Repeat("a", maxHeaderOverrideValueLength+1) + defensive := headerOverrideTestAccount(PlatformOpenAI, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{ + "x-big": oversizedValue, + "sec-websocket-key": "forged", + "content-type": "application/json", // 名单扩充前落库的数据也要被拦截 + "x-claude-code-session-id": "pinned-session", + "x-ok": "ok", + }, + }) + require.Equal(t, map[string]string{"x-ok": "ok"}, defensive.GetHeaderOverrides()) +} + +func TestApplyHeaderOverrides(t *testing.T) { + acc := headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{ + "user-agent": "override-agent/2.0", + "anthropic-beta": "custom-beta-1", + "x-custom": "custom-value", + }, + }) + + h := http.Header{} + // 模拟转发链路:canonical key 与 wire casing 原样 key 混合存在 + h.Set("User-Agent", "claude-cli/2.1.161 (external, cli)") + h["anthropic-beta"] = []string{"claude-code-20250219,oauth-2025-04-20"} // 非 canonical 原样 key + h.Set("Content-Type", "application/json") + + acc.ApplyHeaderOverrides(h) + + // user-agent 覆盖且只有一个值(已知头恢复 wire casing) + require.Equal(t, []string{"override-agent/2.0"}, h["User-Agent"]) + // anthropic-beta:非 canonical 旧值被清除,写入 wire casing(小写) + require.Equal(t, []string{"custom-beta-1"}, h["anthropic-beta"]) + require.Empty(t, h["Anthropic-Beta"]) + // 新增头(未知头以小写原样键写入,与转发链路 wire casing 约定一致) + require.Equal(t, []string{"custom-value"}, h["x-custom"]) + require.Equal(t, "custom-value", getHeaderRaw(h, "x-custom")) + // 未覆写的头不受影响 + require.Equal(t, "application/json", h.Get("Content-Type")) + + // 覆盖后不存在任何大小写重复 + count := 0 + for k := range h { + if k == "anthropic-beta" || k == "Anthropic-Beta" { + count++ + } + } + require.Equal(t, 1, count) +} + +func TestApplyHeaderOverridesNoOpPaths(t *testing.T) { + baseline := func() http.Header { + h := http.Header{} + h.Set("User-Agent", "orig") + return h + } + + // OAuth 账号:即使配置了覆写也不生效 + oauth := headerOverrideTestAccount(PlatformAnthropic, AccountTypeOAuth, map[string]any{ + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{"user-agent": "hacked"}, + }) + h := baseline() + oauth.ApplyHeaderOverrides(h) + require.Equal(t, "orig", h.Get("User-Agent")) + + // 未启用开关 + off := headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrides: map[string]any{"user-agent": "hacked"}, + }) + h = baseline() + off.ApplyHeaderOverrides(h) + require.Equal(t, "orig", h.Get("User-Agent")) + + // 禁止覆写的头(authorization / x-api-key / host 等)不会被应用 + blocked := headerOverrideTestAccount(PlatformOpenAI, AccountTypeAPIKey, map[string]any{ + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{ + "Authorization": "Bearer evil", + "X-Api-Key": "evil", + "Host": "evil.example.com", + "Content-Length": "0", + }, + }) + h = http.Header{} + h.Set("Authorization", "Bearer real-key") + blocked.ApplyHeaderOverrides(h) + require.Equal(t, "Bearer real-key", h.Get("Authorization")) + require.Empty(t, h.Get("X-Api-Key")) + require.Empty(t, h.Get("Host")) + + // nil header 不 panic + blocked.ApplyHeaderOverrides(nil) +} + +func TestNormalizeHeaderOverrideCredentials(t *testing.T) { + t.Run("nil credentials no-op", func(t *testing.T) { + require.NoError(t, NormalizeHeaderOverrideCredentials(nil)) + }) + + t.Run("missing keys no-op", func(t *testing.T) { + creds := map[string]any{"api_key": "sk-xxx"} + require.NoError(t, NormalizeHeaderOverrideCredentials(creds)) + _, exists := creds[credKeyHeaderOverrides] + require.False(t, exists) + }) + + t.Run("normalizes names and values", func(t *testing.T) { + creds := map[string]any{ + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{ + " User-Agent ": " my-agent ", + "X-App": "", + "": "", // 完全空行被丢弃 + }, + } + require.NoError(t, NormalizeHeaderOverrideCredentials(creds)) + require.Equal(t, map[string]any{ + "user-agent": "my-agent", + "x-app": "", + }, creds[credKeyHeaderOverrides]) + }) + + t.Run("accepts map[string]string input", func(t *testing.T) { + creds := map[string]any{ + credKeyHeaderOverrides: map[string]string{"X-App": "cli"}, + } + require.NoError(t, NormalizeHeaderOverrideCredentials(creds)) + require.Equal(t, map[string]any{"x-app": "cli"}, creds[credKeyHeaderOverrides]) + }) + + t.Run("rejects non-bool enabled", func(t *testing.T) { + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrideEnabled: "yes", + }) + require.Error(t, err) + }) + + t.Run("rejects non-object overrides", func(t *testing.T) { + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: []any{"user-agent"}, + }) + require.Error(t, err) + }) + + t.Run("rejects non-string value", func(t *testing.T) { + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: map[string]any{"x-app": 123}, + }) + require.Error(t, err) + }) + + t.Run("rejects invalid header name", func(t *testing.T) { + for _, name := range []string{"bad name", "bad:name", "bad\nname", "值"} { + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: map[string]any{name: "v"}, + }) + require.Error(t, err, "name %q should be rejected", name) + } + }) + + t.Run("rejects empty name with value", func(t *testing.T) { + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: map[string]any{" ": "v"}, + }) + require.Error(t, err) + }) + + t.Run("rejects blocked headers", func(t *testing.T) { + for _, name := range []string{ + "Authorization", "x-api-key", "Host", "content-length", "Transfer-Encoding", + "connection", "accept-encoding", "Sec-WebSocket-Key", "session_id", + "conversation_id", "x-codex-turn-state", "chatgpt-account-id", + "Content-Type", "Cookie", "x-goog-api-key", + "X-Claude-Code-Session-Id", "x-client-request-id", + } { + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: map[string]any{name: "v"}, + }) + require.Error(t, err, "blocked header %q should be rejected", name) + } + }) + + t.Run("allows tab inside value", func(t *testing.T) { + creds := map[string]any{ + credKeyHeaderOverrides: map[string]any{"x-app": "a\tb"}, + } + require.NoError(t, NormalizeHeaderOverrideCredentials(creds)) + require.Equal(t, map[string]any{"x-app": "a\tb"}, creds[credKeyHeaderOverrides]) + }) + + t.Run("rejects invalid value", func(t *testing.T) { + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: map[string]any{"x-app": "bad\nvalue"}, + }) + require.Error(t, err) + }) + + t.Run("rejects duplicate names case-insensitively", func(t *testing.T) { + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: map[string]any{ + "User-Agent": "a", + "user-agent": "b", + }, + }) + require.Error(t, err) + }) + + t.Run("rejects too many entries", func(t *testing.T) { + entries := make(map[string]any, maxHeaderOverrideEntries+1) + for i := 0; i <= maxHeaderOverrideEntries; i++ { + entries["x-h-"+string(rune('a'+i%26))+string(rune('a'+(i/26)%26))+string(rune('a'+(i/676)%26))] = "v" + } + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: entries, + }) + require.Error(t, err) + }) + + t.Run("rejects oversized value", func(t *testing.T) { + big := make([]byte, maxHeaderOverrideValueLength+1) + for i := range big { + big[i] = 'a' + } + err := NormalizeHeaderOverrideCredentials(map[string]any{ + credKeyHeaderOverrides: map[string]any{"x-app": string(big)}, + }) + require.Error(t, err) + }) +} diff --git a/backend/internal/service/account_service.go b/backend/internal/service/account_service.go index dcba614c2c..5956684f98 100644 --- a/backend/internal/service/account_service.go +++ b/backend/internal/service/account_service.go @@ -39,6 +39,9 @@ type AccountRepository interface { List(ctx context.Context, params pagination.PaginationParams) ([]Account, *pagination.PaginationResult, error) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) + // ListAllWithFilters 返回符合过滤条件的全部账号(不分页),用于账号列表页 + // 计算 OpenAI 调度分数的过滤范围池。 + ListAllWithFilters(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) ListByGroup(ctx context.Context, groupID int64) ([]Account, error) ListActive(ctx context.Context) ([]Account, error) ListOAuthRefreshCandidates(ctx context.Context) ([]Account, error) diff --git a/backend/internal/service/account_service_delete_test.go b/backend/internal/service/account_service_delete_test.go index a304356c09..ee6163239e 100644 --- a/backend/internal/service/account_service_delete_test.go +++ b/backend/internal/service/account_service_delete_test.go @@ -79,6 +79,10 @@ func (s *accountRepoStub) List(ctx context.Context, params pagination.Pagination panic("unexpected List call") } +func (s *accountRepoStub) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]Account, error) { + return nil, nil +} + func (s *accountRepoStub) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) { panic("unexpected ListWithFilters call") } diff --git a/backend/internal/service/account_test_service.go b/backend/internal/service/account_test_service.go index 80a862b971..b598ac8ced 100644 --- a/backend/internal/service/account_test_service.go +++ b/backend/internal/service/account_test_service.go @@ -295,6 +295,9 @@ func (s *AccountTestService) testClaudeAccountConnection(c *gin.Context, account setAnthropicAPIKeyAuthHeader(req.Header, account, authToken) } + // 账号级请求头覆写:测试请求与真实转发保持一致的最终头 + account.ApplyHeaderOverrides(req.Header) + // Get proxy URL proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { @@ -603,6 +606,9 @@ func (s *AccountTestService) testOpenAIAccountConnection(c *gin.Context, account setOpenAIChatGPTAccountHeaders(req.Header, credentialAccount) } + // 账号级请求头覆写:测试请求与真实转发保持一致的最终头 + credentialAccount.ApplyHeaderOverrides(req.Header) + // Get proxy URL proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { @@ -756,6 +762,9 @@ func (s *AccountTestService) testOpenAIChatCompletionsConnection( req.Header.Set("Accept", "text/event-stream") req.Header.Set("Authorization", "Bearer "+authToken) + // 账号级请求头覆写:测试请求与真实转发保持一致的最终头 + account.ApplyHeaderOverrides(req.Header) + proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { proxyURL = account.Proxy.URL() @@ -848,6 +857,9 @@ func (s *AccountTestService) testOpenAICompactConnection(c *gin.Context, account setOpenAIChatGPTAccountHeaders(req.Header, account) } + // 账号级请求头覆写:测试请求与真实转发保持一致的最终头 + account.ApplyHeaderOverrides(req.Header) + proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { proxyURL = account.Proxy.URL() @@ -1599,6 +1611,9 @@ func (s *AccountTestService) testOpenAIImageAPIKey(c *gin.Context, ctx context.C req.Header.Set("Content-Type", "application/json") req.Header.Set("Authorization", "Bearer "+authToken) + // 账号级请求头覆写:测试请求与真实转发保持一致的最终头 + account.ApplyHeaderOverrides(req.Header) + proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { proxyURL = account.Proxy.URL() diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index acdfe08202..435164154f 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -78,6 +78,12 @@ type AdminService interface { // Account management ListAccounts(ctx context.Context, page, pageSize int, platform, accountType, status, search string, groupID int64, privacyMode string, sortBy, sortOrder string) ([]Account, int64, error) + // ListAccountsForSchedulerScoreFilter 返回符合过滤条件的全部账号(不分页), + // 作为账号列表页计算 OpenAI 调度分数的过滤范围池。 + ListAccountsForSchedulerScoreFilter(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) + // ListOpenAISchedulableAccountsForSchedulerScore 返回指定分组(nil 为未分组)内 + // 可调度的 OpenAI 账号,用于按组计算调度分数。 + ListOpenAISchedulableAccountsForSchedulerScore(ctx context.Context, groupID *int64) ([]Account, error) GetAccount(ctx context.Context, id int64) (*Account, error) GetAccountsByIDs(ctx context.Context, ids []int64) ([]*Account, error) CreateAccount(ctx context.Context, input *CreateAccountInput) (*Account, error) @@ -2664,13 +2670,7 @@ func (s *adminServiceImpl) ListAccountsForSchedulerScoreFilter(ctx context.Conte if s == nil || s.accountRepo == nil { return nil, nil } - lister, ok := s.accountRepo.(interface { - ListAllWithFilters(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) - }) - if !ok { - return nil, nil - } - return lister.ListAllWithFilters(ctx, platform, accountType, status, search, groupID, privacyMode) + return s.accountRepo.ListAllWithFilters(ctx, platform, accountType, status, search, groupID, privacyMode) } func (s *adminServiceImpl) ListOpenAISchedulableAccountsForSchedulerScore(ctx context.Context, groupID *int64) ([]Account, error) { @@ -2733,6 +2733,11 @@ func (s *adminServiceImpl) CreateAccount(ctx context.Context, input *CreateAccou } } + // 校验并规范化请求头覆写配置(header 名小写化、格式检查) + if err := NormalizeHeaderOverrideCredentials(input.Credentials); err != nil { + return nil, err + } + account := &Account{ Name: input.Name, Notes: normalizeAccountNotes(input.Notes), @@ -2863,6 +2868,10 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U // 敏感子键采用"incoming 没提供就保留"的合并语义:前端响应已脱敏, // 全对象 PUT 编辑时不会再带回 token,避免覆盖时清空已有凭证。 account.Credentials = MergePreservingSensitiveCreds(account.Credentials, input.Credentials) + // 校验并规范化请求头覆写配置(header 名小写化、格式检查) + if err := NormalizeHeaderOverrideCredentials(account.Credentials); err != nil { + return nil, err + } } // Extra 使用 map:需要区分“未提供(nil)”与“显式清空({})”。 // 关闭配额限制时前端会删除 quota_* 键并提交 extra:{},此时也必须落库。 @@ -3081,6 +3090,11 @@ func (s *adminServiceImpl) BulkUpdateAccounts(ctx context.Context, input *BulkUp } } + // 校验并规范化请求头覆写配置(批量路径为 JSONB 顶层 key 合并,直接校验增量即可) + if err := NormalizeHeaderOverrideCredentials(input.Credentials); err != nil { + return nil, err + } + // Prepare bulk updates for columns and JSONB fields. repoUpdates := AccountBulkUpdate{ Credentials: input.Credentials, diff --git a/backend/internal/service/admin_service_bulk_update_test.go b/backend/internal/service/admin_service_bulk_update_test.go index df415295b1..2f44b1741d 100644 --- a/backend/internal/service/admin_service_bulk_update_test.go +++ b/backend/internal/service/admin_service_bulk_update_test.go @@ -88,6 +88,10 @@ func (s *accountRepoStubForBulkUpdate) ListByGroup(_ context.Context, groupID in return nil, nil } +func (s *accountRepoStubForBulkUpdate) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]Account, error) { + return nil, nil +} + func (s *accountRepoStubForBulkUpdate) ListWithFilters(_ context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) { s.listCalled = true s.lastListParams = params diff --git a/backend/internal/service/admin_service_search_test.go b/backend/internal/service/admin_service_search_test.go index 595e99e344..76acd1b5e2 100644 --- a/backend/internal/service/admin_service_search_test.go +++ b/backend/internal/service/admin_service_search_test.go @@ -25,6 +25,10 @@ type accountRepoStubForAdminList struct { listWithFiltersErr error } +func (s *accountRepoStubForAdminList) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]Account, error) { + return nil, nil +} + func (s *accountRepoStubForAdminList) ListWithFilters(_ context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) { s.listWithFiltersCalls++ s.listWithFiltersParams = params diff --git a/backend/internal/service/api_key.go b/backend/internal/service/api_key.go index ec20b0a9bf..dfc3ec1c5a 100644 --- a/backend/internal/service/api_key.go +++ b/backend/internal/service/api_key.go @@ -44,6 +44,7 @@ type APIKey struct { UpdatedAt time.Time User *User Group *Group + CurrentConcurrency int // Quota fields Quota float64 // Quota limit in USD (0 = unlimited) diff --git a/backend/internal/service/api_key_service.go b/backend/internal/service/api_key_service.go index de9b908dc7..8903be65ee 100644 --- a/backend/internal/service/api_key_service.go +++ b/backend/internal/service/api_key_service.go @@ -203,6 +203,7 @@ type APIKeyService struct { userGroupRateRepo UserGroupRateRepository cache APIKeyCache rateLimitCacheInvalid RateLimitCacheInvalidator // optional: invalidate Redis rate limit cache + concurrencyService *ConcurrencyService cfg *config.Config authCacheL1 *ristretto.Cache authCfg apiKeyAuthCacheConfig @@ -240,6 +241,10 @@ func (s *APIKeyService) SetRateLimitCacheInvalidator(inv RateLimitCacheInvalidat s.rateLimitCacheInvalid = inv } +func (s *APIKeyService) SetConcurrencyService(concurrencyService *ConcurrencyService) { + s.concurrencyService = concurrencyService +} + func (s *APIKeyService) compileAPIKeyIPRules(apiKey *APIKey) { if apiKey == nil { return @@ -436,9 +441,40 @@ func (s *APIKeyService) List(ctx context.Context, userID int64, params paginatio if err != nil { return nil, nil, fmt.Errorf("list api keys: %w", err) } + s.fillCurrentConcurrency(ctx, keys) return keys, pagination, nil } +func (s *APIKeyService) fillCurrentConcurrency(ctx context.Context, keys []APIKey) { + if s == nil || s.concurrencyService == nil || len(keys) == 0 { + return + } + ids := make([]int64, 0, len(keys)) + for i := range keys { + if keys[i].ID > 0 { + ids = append(ids, keys[i].ID) + } + } + counts, err := s.concurrencyService.GetAPIKeyConcurrencyBatch(ctx, ids) + if err != nil { + return + } + for i := range keys { + keys[i].CurrentConcurrency = counts[keys[i].ID] + } +} + +func (s *APIKeyService) currentConcurrencyForAPIKey(ctx context.Context, apiKeyID int64) int { + if s == nil || s.concurrencyService == nil || apiKeyID <= 0 { + return 0 + } + counts, err := s.concurrencyService.GetAPIKeyConcurrencyBatch(ctx, []int64{apiKeyID}) + if err != nil { + return 0 + } + return counts[apiKeyID] +} + func (s *APIKeyService) VerifyOwnership(ctx context.Context, userID int64, apiKeyIDs []int64) ([]int64, error) { if len(apiKeyIDs) == 0 { return []int64{}, nil @@ -458,6 +494,9 @@ func (s *APIKeyService) GetByID(ctx context.Context, id int64) (*APIKey, error) return nil, fmt.Errorf("get api key: %w", err) } s.compileAPIKeyIPRules(apiKey) + if apiKey != nil { + apiKey.CurrentConcurrency = s.currentConcurrencyForAPIKey(ctx, apiKey.ID) + } return apiKey, nil } diff --git a/backend/internal/service/api_key_service_delete_test.go b/backend/internal/service/api_key_service_delete_test.go index 8664c03bd7..25ad1edb15 100644 --- a/backend/internal/service/api_key_service_delete_test.go +++ b/backend/internal/service/api_key_service_delete_test.go @@ -300,6 +300,40 @@ func TestApiKeyService_Delete_NotFound(t *testing.T) { require.Empty(t, cache.deleteAuthKeys) } +func TestAPIKeyService_List_FillsCurrentConcurrency(t *testing.T) { + repo := &apiKeyRepoStub{ + allowListByUserID: true, + listByUserIDKeys: []APIKey{ + {ID: 10, UserID: 7, Key: "sk-10", Name: "key-10"}, + {ID: 11, UserID: 7, Key: "sk-11", Name: "key-11"}, + }, + } + concurrency := NewConcurrencyService(&stubConcurrencyCacheForTest{ + apiKeyConcurrency: map[int64]int{10: 2, 11: 0}, + }) + svc := &APIKeyService{apiKeyRepo: repo, concurrencyService: concurrency} + + keys, _, err := svc.List(context.Background(), 7, pagination.PaginationParams{Page: 1, PageSize: 20}, APIKeyListFilters{}) + require.NoError(t, err) + require.Len(t, keys, 2) + require.Equal(t, 2, keys[0].CurrentConcurrency) + require.Equal(t, 0, keys[1].CurrentConcurrency) +} + +func TestAPIKeyService_GetByID_FillsCurrentConcurrency(t *testing.T) { + repo := &apiKeyRepoStub{ + apiKey: &APIKey{ID: 10, UserID: 7, Key: "sk-10", Name: "key-10"}, + } + concurrency := NewConcurrencyService(&stubConcurrencyCacheForTest{ + apiKeyConcurrency: map[int64]int{10: 4}, + }) + svc := &APIKeyService{apiKeyRepo: repo, concurrencyService: concurrency} + + key, err := svc.GetByID(context.Background(), 10) + require.NoError(t, err) + require.Equal(t, 4, key.CurrentConcurrency) +} + // TestApiKeyService_Delete_DeleteFails 测试删除操作失败时的错误处理。 // 预期行为: // - GetKeyAndOwnerID 返回正确的所有者 ID diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index a781936598..dc54a1b1f3 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -280,6 +280,11 @@ func (s *BillingService) initFallbackPricing() { s.fallbackPrices["gpt-5.5"] = s.fallbackPrices["gpt-5.4"] s.fallbackPrices["gpt-5.5-pro"] = s.fallbackPrices["gpt-5.4"] + // GPT-5.6(sol / terra / luna)暂无独立定价,回退到 GPT-5.4。 + s.fallbackPrices["gpt-5.6-sol"] = s.fallbackPrices["gpt-5.4"] + s.fallbackPrices["gpt-5.6-terra"] = s.fallbackPrices["gpt-5.4"] + s.fallbackPrices["gpt-5.6-luna"] = s.fallbackPrices["gpt-5.4"] + s.fallbackPrices["gpt-5.4-mini"] = &ModelPricing{ InputPricePerToken: 7.5e-7, OutputPricePerToken: 4.5e-6, @@ -667,6 +672,12 @@ func (s *BillingService) getFallbackPricing(model string) *ModelPricing { // OpenAI(GPT-5 / Codex 族):仅匹配已知型号,避免未知 OpenAI 型号误计价。 if normalized := normalizeKnownOpenAICodexModel(modelLower); normalized != "" { switch normalized { + case "gpt-5.6-sol": + return s.fallbackPrices["gpt-5.6-sol"] + case "gpt-5.6-terra": + return s.fallbackPrices["gpt-5.6-terra"] + case "gpt-5.6-luna": + return s.fallbackPrices["gpt-5.6-luna"] case "gpt-5.5-pro": return s.fallbackPrices["gpt-5.5-pro"] case "gpt-5.5": @@ -1060,7 +1071,8 @@ func isOpenAIGPT54Model(model string) bool { // normalizeCodexModel 的默认兜底把非 OpenAI 模型(claude-*、gemini-*、gpt-4o) // 误识别为 gpt-5.4。 normalized := normalizeKnownOpenAICodexModel(model) - return normalized == "gpt-5.4" || normalized == "gpt-5.5" || normalized == "gpt-5.5-pro" + return normalized == "gpt-5.4" || normalized == "gpt-5.5" || normalized == "gpt-5.5-pro" || + normalized == "gpt-5.6-sol" || normalized == "gpt-5.6-terra" || normalized == "gpt-5.6-luna" } // CalculateCostWithConfig 使用配置中的默认倍率计算费用 diff --git a/backend/internal/service/concurrency_service.go b/backend/internal/service/concurrency_service.go index 50e0077b94..c4ec6c5f98 100644 --- a/backend/internal/service/concurrency_service.go +++ b/backend/internal/service/concurrency_service.go @@ -53,6 +53,12 @@ type ConcurrencyCache interface { CleanupStaleProcessSlots(ctx context.Context, activeRequestPrefix string) error } +type APIKeyConcurrencyCache interface { + TrackAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error + ReleaseAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error + GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) +} + var ( requestIDPrefix = initRequestIDPrefix() requestIDCounter atomic.Uint64 @@ -90,6 +96,8 @@ const ( defaultAccountLoadBatchCacheTTL = 200 * time.Millisecond accountLoadBatchFetchTimeout = 3 * time.Second maxAccountLoadBatchCacheEntries = 256 + apiKeyConcurrencyFetchTimeout = 3 * time.Second + apiKeySlotTrackTimeout = 2 * time.Second ) // ConcurrencyService 管理账号和用户的并发限制。 @@ -238,6 +246,77 @@ func (s *ConcurrencyService) AcquireUserSlot(ctx context.Context, userID int64, }, nil } +// TrackAPIKeySlot records one active request slot for an API key without +// applying key-level concurrency limits. It is fail-open: Redis errors are +// logged and return a no-op release function. +func (s *ConcurrencyService) TrackAPIKeySlot(ctx context.Context, apiKeyID int64) func() { + if s == nil || s.cache == nil || apiKeyID <= 0 { + return func() {} + } + cache, ok := s.cache.(APIKeyConcurrencyCache) + if !ok { + return func() {} + } + + requestID := generateRequestID() + baseCtx := context.Background() + if ctx != nil { + baseCtx = context.WithoutCancel(ctx) + } + trackCtx, cancel := context.WithTimeout(baseCtx, apiKeySlotTrackTimeout) + err := cache.TrackAPIKeySlot(trackCtx, apiKeyID, requestID) + cancel() + if err != nil { + logger.LegacyPrintf("service.concurrency", "Warning: failed to track api key slot for %d (req=%s): %v", apiKeyID, requestID, err) + return func() {} + } + + return func() { + bgCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + if err := cache.ReleaseAPIKeySlot(bgCtx, apiKeyID, requestID); err != nil { + logger.LegacyPrintf("service.concurrency", "Warning: failed to release api key slot for %d (req=%s): %v", apiKeyID, requestID, err) + } + } +} + +// GetAPIKeyConcurrencyBatch gets real-time active request counts for API keys. +// Stats are best-effort: missing Redis support or Redis errors return zeroes. +func (s *ConcurrencyService) GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) { + result := zeroAPIKeyConcurrencyMap(apiKeyIDs) + if len(apiKeyIDs) == 0 { + return result, nil + } + if s == nil || s.cache == nil { + return result, nil + } + cache, ok := s.cache.(APIKeyConcurrencyCache) + if !ok { + return result, nil + } + + redisCtx, cancel := context.WithTimeout(context.Background(), apiKeyConcurrencyFetchTimeout) + defer cancel() + + counts, err := cache.GetAPIKeyConcurrencyBatch(redisCtx, apiKeyIDs) + if err != nil { + logger.LegacyPrintf("service.concurrency", "Warning: get api key concurrency batch failed: %v", err) + return result, nil + } + for _, apiKeyID := range apiKeyIDs { + result[apiKeyID] = counts[apiKeyID] + } + return result, nil +} + +func zeroAPIKeyConcurrencyMap(apiKeyIDs []int64) map[int64]int { + result := make(map[int64]int, len(apiKeyIDs)) + for _, apiKeyID := range apiKeyIDs { + result[apiKeyID] = 0 + } + return result +} + // ============================================ // Wait Queue Count Methods // ============================================ diff --git a/backend/internal/service/concurrency_service_test.go b/backend/internal/service/concurrency_service_test.go index bacad0245e..3f358bbe6a 100644 --- a/backend/internal/service/concurrency_service_test.go +++ b/backend/internal/service/concurrency_service_test.go @@ -16,25 +16,33 @@ import ( // stubConcurrencyCacheForTest 用于并发服务单元测试的缓存桩 type stubConcurrencyCacheForTest struct { - acquireResult bool - acquireErr error - releaseErr error - concurrency int - concurrencyErr error - waitAllowed bool - waitErr error - waitCount int - waitCountErr error - loadBatch map[int64]*AccountLoadInfo - loadBatchErr error - usersLoadBatch map[int64]*UserLoadInfo - usersLoadErr error - cleanupErr error + acquireResult bool + acquireErr error + releaseErr error + concurrency int + concurrencyErr error + waitAllowed bool + waitErr error + waitCount int + waitCountErr error + loadBatch map[int64]*AccountLoadInfo + loadBatchErr error + usersLoadBatch map[int64]*UserLoadInfo + usersLoadErr error + cleanupErr error + apiKeyTrackErr error + apiKeyReleaseErr error + apiKeyConcurrency map[int64]int + apiKeyConcurrencyErr error // 记录调用 - releasedAccountIDs []int64 - releasedRequestIDs []string - loadBatchCalls atomic.Int64 + releasedAccountIDs []int64 + releasedRequestIDs []string + loadBatchCalls atomic.Int64 + trackedAPIKeyIDs []int64 + trackedAPIKeyRequestIDs []string + releasedAPIKeyIDs []int64 + releasedAPIKeyRequestIDs []string } var _ ConcurrencyCache = (*stubConcurrencyCacheForTest)(nil) @@ -78,6 +86,26 @@ func (c *stubConcurrencyCacheForTest) ReleaseUserSlot(_ context.Context, _ int64 func (c *stubConcurrencyCacheForTest) GetUserConcurrency(_ context.Context, _ int64) (int, error) { return c.concurrency, c.concurrencyErr } +func (c *stubConcurrencyCacheForTest) TrackAPIKeySlot(_ context.Context, apiKeyID int64, requestID string) error { + c.trackedAPIKeyIDs = append(c.trackedAPIKeyIDs, apiKeyID) + c.trackedAPIKeyRequestIDs = append(c.trackedAPIKeyRequestIDs, requestID) + return c.apiKeyTrackErr +} +func (c *stubConcurrencyCacheForTest) ReleaseAPIKeySlot(_ context.Context, apiKeyID int64, requestID string) error { + c.releasedAPIKeyIDs = append(c.releasedAPIKeyIDs, apiKeyID) + c.releasedAPIKeyRequestIDs = append(c.releasedAPIKeyRequestIDs, requestID) + return c.apiKeyReleaseErr +} +func (c *stubConcurrencyCacheForTest) GetAPIKeyConcurrencyBatch(_ context.Context, apiKeyIDs []int64) (map[int64]int, error) { + if c.apiKeyConcurrencyErr != nil { + return nil, c.apiKeyConcurrencyErr + } + result := make(map[int64]int, len(apiKeyIDs)) + for _, apiKeyID := range apiKeyIDs { + result[apiKeyID] = c.apiKeyConcurrency[apiKeyID] + } + return result, nil +} func (c *stubConcurrencyCacheForTest) IncrementWaitCount(_ context.Context, _ int64, _ int) (bool, error) { return c.waitAllowed, c.waitErr } @@ -201,6 +229,62 @@ func TestAcquireUserSlot_UnlimitedConcurrency(t *testing.T) { require.True(t, result.Acquired) } +func TestTrackAPIKeySlot_ReleaseDecrements(t *testing.T) { + cache := &stubConcurrencyCacheForTest{} + svc := NewConcurrencyService(cache) + + release := svc.TrackAPIKeySlot(context.Background(), 88) + require.NotNil(t, release) + require.Equal(t, []int64{88}, cache.trackedAPIKeyIDs) + require.Len(t, cache.trackedAPIKeyRequestIDs, 1) + require.NotEmpty(t, cache.trackedAPIKeyRequestIDs[0]) + + release() + + require.Equal(t, []int64{88}, cache.releasedAPIKeyIDs) + require.Equal(t, cache.trackedAPIKeyRequestIDs, cache.releasedAPIKeyRequestIDs) +} + +func TestTrackAPIKeySlot_FailOpen(t *testing.T) { + cache := &stubConcurrencyCacheForTest{apiKeyTrackErr: errors.New("redis down")} + svc := NewConcurrencyService(cache) + + release := svc.TrackAPIKeySlot(context.Background(), 88) + require.NotNil(t, release) + require.Equal(t, []int64{88}, cache.trackedAPIKeyIDs) + + require.NotPanics(t, release) + require.Empty(t, cache.releasedAPIKeyIDs) +} + +func TestGetAPIKeyConcurrencyBatch_Fallbacks(t *testing.T) { + t.Run("nil cache returns zeroes", func(t *testing.T) { + svc := &ConcurrencyService{cache: nil} + + counts, err := svc.GetAPIKeyConcurrencyBatch(context.Background(), []int64{1, 2}) + require.NoError(t, err) + require.Equal(t, map[int64]int{1: 0, 2: 0}, counts) + }) + + t.Run("redis error returns zeroes", func(t *testing.T) { + cache := &stubConcurrencyCacheForTest{apiKeyConcurrencyErr: errors.New("redis down")} + svc := NewConcurrencyService(cache) + + counts, err := svc.GetAPIKeyConcurrencyBatch(context.Background(), []int64{1, 2}) + require.NoError(t, err) + require.Equal(t, map[int64]int{1: 0, 2: 0}, counts) + }) + + t.Run("success returns counts", func(t *testing.T) { + cache := &stubConcurrencyCacheForTest{apiKeyConcurrency: map[int64]int{1: 3, 2: 0}} + svc := NewConcurrencyService(cache) + + counts, err := svc.GetAPIKeyConcurrencyBatch(context.Background(), []int64{1, 2}) + require.NoError(t, err) + require.Equal(t, map[int64]int{1: 3, 2: 0}, counts) + }) +} + func TestGenerateRequestID_UsesStablePrefixAndMonotonicCounter(t *testing.T) { id1 := generateRequestID() id2 := generateRequestID() diff --git a/backend/internal/service/gateway_multiplatform_test.go b/backend/internal/service/gateway_multiplatform_test.go index f843ba3e45..35a60c8124 100644 --- a/backend/internal/service/gateway_multiplatform_test.go +++ b/backend/internal/service/gateway_multiplatform_test.go @@ -95,6 +95,9 @@ func (m *mockAccountRepoForPlatform) List(ctx context.Context, params pagination func (m *mockAccountRepoForPlatform) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) { return nil, nil, nil } +func (m *mockAccountRepoForPlatform) ListAllWithFilters(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) { + return nil, nil +} func (m *mockAccountRepoForPlatform) ListByGroup(ctx context.Context, groupID int64) ([]Account, error) { return nil, nil } diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 54035345d9..dcaf3a645c 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -5920,6 +5920,10 @@ func (s *GatewayService) buildUpstreamRequestAnthropicAPIKeyPassthrough( if c != nil && c.Request != nil { clientBeta = getHeaderRaw(c.Request.Header, "anthropic-beta") } + // 账号覆写了 anthropic-beta 时,覆写值即最终上游值:净化以覆写值为准 + if beta, ok := account.HeaderOverrideValue("anthropic-beta"); ok { + clientBeta = beta + } if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta); changed { body = sanitized } @@ -5956,6 +5960,9 @@ func (s *GatewayService) buildUpstreamRequestAnthropicAPIKeyPassthrough( setHeaderRaw(req.Header, "anthropic-version", "2023-06-01") } + // 账号级请求头覆写(最终生效,覆盖上面所有来源的同名头) + account.ApplyHeaderOverrides(req.Header) + return req, body, nil } @@ -6886,6 +6893,12 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex tokenType, mimicClaudeCode, modelID, clientHeaders, body, effectiveDropSet, ) + // 账号覆写了 anthropic-beta 时,覆写值即最终上游值(由下方 ApplyHeaderOverrides 写入): + // body 能力净化必须以覆写值为准,否则 header/body 不对称会被上游 400。 + if beta, ok := account.HeaderOverrideValue("anthropic-beta"); ok { + finalBetaHeader, finalBetaShouldSet = beta, true + } + // 能力维度 body sanitize:与最终 anthropic-beta header 对称 if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, finalBetaHeader); changed { body = sanitized @@ -6959,6 +6972,10 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex } } + // 账号级请求头覆写(仅 anthropic/openai api_key 账号启用时生效;OAuth 路径 no-op)。 + // 放在所有 header 逻辑之后,确保配置值对同名头拥有最终决定权。 + account.ApplyHeaderOverrides(req.Header) + // === DEBUG: 打印上游转发请求(headers + body 摘要),与 CLIENT_ORIGINAL 对比 === s.debugLogGatewaySnapshot("UPSTREAM_FORWARD", req.Header, body, map[string]string{ "url": req.URL.String(), @@ -10410,6 +10427,10 @@ func (s *GatewayService) buildCountTokensRequestAnthropicAPIKeyPassthrough( if c != nil && c.Request != nil { clientBeta = getHeaderRaw(c.Request.Header, "anthropic-beta") } + // 账号覆写了 anthropic-beta 时,覆写值即最终上游值:净化以覆写值为准 + if beta, ok := account.HeaderOverrideValue("anthropic-beta"); ok { + clientBeta = beta + } if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta); changed { body = sanitized } @@ -10445,6 +10466,9 @@ func (s *GatewayService) buildCountTokensRequestAnthropicAPIKeyPassthrough( req.Header.Set("anthropic-version", "2023-06-01") } + // 账号级请求头覆写(最终生效,覆盖上面所有来源的同名头) + account.ApplyHeaderOverrides(req.Header) + return req, nil } @@ -10512,6 +10536,11 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con tokenType, mimicClaudeCode, modelID, clientHeaders, body, ctEffectiveDropSet, ) + // 账号覆写了 anthropic-beta 时,覆写值即最终上游值:净化以覆写值为准 + if beta, ok := account.HeaderOverrideValue("anthropic-beta"); ok { + finalBetaHeader, finalBetaShouldSet = beta, true + } + // 能力维度 body sanitize:与最终 anthropic-beta header 对称 if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, finalBetaHeader); changed { body = sanitized @@ -10578,6 +10607,9 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con } } + // 账号级请求头覆写(仅 anthropic/openai api_key 账号启用时生效;OAuth 路径 no-op) + account.ApplyHeaderOverrides(req.Header) + if c != nil && tokenType == "oauth" { c.Set(claudeMimicDebugInfoKey, buildClaudeMimicDebugLine(req, body, account, tokenType, mimicClaudeCode)) } diff --git a/backend/internal/service/gemini_multiplatform_test.go b/backend/internal/service/gemini_multiplatform_test.go index c021e88edf..7d5ed0ec9e 100644 --- a/backend/internal/service/gemini_multiplatform_test.go +++ b/backend/internal/service/gemini_multiplatform_test.go @@ -82,6 +82,9 @@ func (m *mockAccountRepoForGemini) List(ctx context.Context, params pagination.P func (m *mockAccountRepoForGemini) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) { return nil, nil, nil } +func (m *mockAccountRepoForGemini) ListAllWithFilters(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) { + return nil, nil +} func (m *mockAccountRepoForGemini) ListByGroup(ctx context.Context, groupID int64) ([]Account, error) { return nil, nil } diff --git a/backend/internal/service/openai_account_scheduler.go b/backend/internal/service/openai_account_scheduler.go index dd65163abc..a0a4fbff3e 100644 --- a/backend/internal/service/openai_account_scheduler.go +++ b/backend/internal/service/openai_account_scheduler.go @@ -1037,6 +1037,14 @@ func (s *defaultOpenAIAccountScheduler) tryFallbackToWeightedSticky( if account == nil || !s.isAccountRequestCompatible(ctx, account, req) || !s.isAccountTransportCompatible(account, req.RequiredTransport) { continue } + // 粘性绑定只证明绑定时账号在分组内;账号被移出分组后绑定仍会在 TTL 内存活, + // 必须与 selectBySessionHash 一样重验分组归属,否则会把分组流量泄漏到组外账号。 + if !openAIStickyAccountMatchesGroup(account, req.GroupID) { + if accountID == req.StickyAccountID && strings.TrimSpace(req.SessionHash) != "" { + _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, req.SessionHash) + } + continue + } if req.RequireCompact && openAICompactSupportTier(account) == 0 { continue } @@ -1145,13 +1153,29 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance( } if len(regularAccounts) > 0 { regularAttempt := s.trySelectByLoadBalancePool(ctx, req, regularAccounts, loadMap) - if regularAttempt.err != nil { + if regularAttempt.err != nil && !regularAttempt.noCompactCandidates { return nil, regularAttempt.candidateCount, regularAttempt.topK, regularAttempt.loadSkew, regularAttempt.err } if regularAttempt.result != nil { return regularAttempt.result, regularAttempt.candidateCount, regularAttempt.topK, regularAttempt.loadSkew, nil } - return s.finishLoadBalanceSelectionFallback(ctx, req, regularAttempt) + var result *AccountSelectionResult + candidateCount, topK, loadSkew := regularAttempt.candidateCount, regularAttempt.topK, regularAttempt.loadSkew + fallbackErr := regularAttempt.err + if regularAttempt.err == nil { + result, candidateCount, topK, loadSkew, fallbackErr = s.finishLoadBalanceSelectionFallback(ctx, req, regularAttempt) + if fallbackErr == nil && result != nil { + return result, candidateCount, topK, loadSkew, nil + } + } + // 常规池既无法获取也无法排队(含仅剩不支持 compact 的候选)时, + // 回退到订阅池的等待计划:busy-but-waitable 的订阅账号不应因常规池存在 + // 而被丢弃,否则开启订阅优先反而让本可排队成功的请求硬失败。 + subResult, subCandidateCount, subTopK, subLoadSkew, subErr := s.finishLoadBalanceSelectionFallback(ctx, req, attempt) + if subErr == nil && subResult != nil { + return subResult, subCandidateCount, subTopK, subLoadSkew, nil + } + return result, candidateCount, topK, loadSkew, fallbackErr } return s.finishLoadBalanceSelectionFallback(ctx, req, attempt) } @@ -1464,15 +1488,20 @@ func (s *OpenAIGatewayService) openAIAdvancedSchedulerRuntimeSettings(ctx contex lbTopKOverride = parsePositiveIntOverride(values[SettingKeyOpenAIAdvancedSchedulerLBTopK]) weightOverrides = parseOpenAIAdvancedSchedulerWeightOverrides(values) } else { - if value, err := repo.GetValue(dbCtx, openAIAdvancedSchedulerSettingKey); err == nil { - enabled = strings.EqualFold(strings.TrimSpace(value), "true") - } - if value, err := repo.GetValue(dbCtx, SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled); err == nil { - stickyWeightedEnabled = strings.EqualFold(strings.TrimSpace(value), "true") - } - if value, err := repo.GetValue(dbCtx, SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled); err == nil { - subscriptionPriorityEnabled = strings.EqualFold(strings.TrimSpace(value), "true") + // 批量读取失败时逐键降级,覆盖全部键(含 TopK/权重),避免只加载布尔开关 + // 而静默丢弃管理员配置的覆盖值;降级状态会被缓存一个 TTL,必须留痕。 + slog.Warn("openai_advanced_scheduler_settings_batch_load_failed", "error", err) + fallbackValues := make(map[string]string) + for _, key := range openAIAdvancedSchedulerRuntimeSettingKeys() { + if value, valueErr := repo.GetValue(dbCtx, key); valueErr == nil { + fallbackValues[key] = value + } } + enabled = strings.EqualFold(strings.TrimSpace(fallbackValues[openAIAdvancedSchedulerSettingKey]), "true") + stickyWeightedEnabled = strings.EqualFold(strings.TrimSpace(fallbackValues[SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled]), "true") + subscriptionPriorityEnabled = strings.EqualFold(strings.TrimSpace(fallbackValues[SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled]), "true") + lbTopKOverride = parsePositiveIntOverride(fallbackValues[SettingKeyOpenAIAdvancedSchedulerLBTopK]) + weightOverrides = parseOpenAIAdvancedSchedulerWeightOverrides(fallbackValues) } } @@ -1618,6 +1647,9 @@ func (s *OpenAIGatewayService) SelectAccountWithScheduler( return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, "", "", requireCompact, PlatformOpenAI, false) } +// SelectAccountWithSchedulerForCapability 按能力要求调度账号。 +// previousResponseCanMove 表示首包 input 可自行重建工具续链,previous_response_id 允许跨账号迁移 +// (粘性加权模式下改为加权偏好而非硬粘连)。 func (s *OpenAIGatewayService) SelectAccountWithSchedulerForCapability( ctx context.Context, groupID *int64, @@ -1628,16 +1660,13 @@ func (s *OpenAIGatewayService) SelectAccountWithSchedulerForCapability( requiredTransport OpenAIUpstreamTransport, requiredCapability OpenAIEndpointCapability, requireCompact bool, + previousResponseCanMove bool, platformOverride ...string, ) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) { platform := PlatformOpenAI - previousResponseCanMove := false if len(platformOverride) > 0 { platform = platformOverride[0] } - if len(platformOverride) > 1 { - previousResponseCanMove = strings.EqualFold(platformOverride[1], "previous_response_can_move") - } return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, requiredCapability, "", requireCompact, platform, previousResponseCanMove) } @@ -1853,6 +1882,11 @@ func (s *OpenAIGatewayService) openAIWSLBTopK() int { func (s *OpenAIGatewayService) openAIWSLBTopKForRequest(ctx context.Context) int { base := s.openAIWSLBTopK() settings := s.openAIAdvancedSchedulerRuntimeSettings(ctx) + // DB 覆盖值与 stickyWeighted/subscriptionPriority 一样受总开关门控: + // 关闭高级调度器后所有调用方(含管理页分数快照)都应回到配置/默认行为。 + if !settings.enabled { + return base + } if settings.lbTopKOverride > 0 { return settings.lbTopKOverride } @@ -1920,6 +1954,10 @@ func (s *OpenAIGatewayService) openAIWSSchedulerWeights() GatewayOpenAIWSSchedul func (s *OpenAIGatewayService) openAIWSSchedulerWeightsForRequest(ctx context.Context) GatewayOpenAIWSSchedulerScoreWeightsView { weights := s.openAIWSSchedulerWeights() settings := s.openAIAdvancedSchedulerRuntimeSettings(ctx) + // 同 openAIWSLBTopKForRequest:总开关关闭时不应用 DB 覆盖值。 + if !settings.enabled { + return weights + } return applyOpenAIAdvancedSchedulerWeightOverrides(weights, settings.weightOverrides) } diff --git a/backend/internal/service/openai_account_scheduler_test.go b/backend/internal/service/openai_account_scheduler_test.go index a61a923053..2ff7c25e5d 100644 --- a/backend/internal/service/openai_account_scheduler_test.go +++ b/backend/internal/service/openai_account_scheduler_test.go @@ -529,6 +529,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_Embeddi OpenAIUpstreamTransportHTTPSSE, OpenAIEndpointCapabilityEmbeddings, false, + false, ) require.NoError(t, err) require.NotNil(t, selection) @@ -572,6 +573,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_AllowsG OpenAIUpstreamTransportAny, OpenAIEndpointCapabilityChatCompletions, false, + false, PlatformGrok, ) require.NoError(t, err) @@ -774,6 +776,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_StickyWeightedPreviousR OpenAIUpstreamTransportAny, OpenAIEndpointCapabilityChatCompletions, false, + false, PlatformOpenAI, ) require.NoError(t, err) @@ -796,8 +799,8 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_StickyWeightedPreviousR OpenAIUpstreamTransportAny, OpenAIEndpointCapabilityChatCompletions, false, + true, PlatformOpenAI, - "previous_response_can_move", ) require.NoError(t, err) require.NotNil(t, selection) @@ -938,6 +941,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_Enabled_EmbeddingsSkips OpenAIUpstreamTransportHTTPSSE, OpenAIEndpointCapabilityEmbeddings, false, + false, ) require.NoError(t, err) require.NotNil(t, selection) @@ -1011,6 +1015,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_Enabled_EmbeddingsSkips OpenAIUpstreamTransportHTTPSSE, OpenAIEndpointCapabilityEmbeddings, false, + false, ) require.NoError(t, err) require.NotNil(t, selection) @@ -2935,3 +2940,141 @@ func TestDefaultOpenAIAccountScheduler_IsAccountTransportCompatible_Branches(t * func int64PtrForTest(v int64) *int64 { return &v } + +func TestOpenAIGatewayService_SelectAccountWithScheduler_StickyWeightedFallbackSkipsOutOfGroupStickyAccount(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + + ctx := context.Background() + groupID := int64(101081) + otherGroupID := int64(101082) + accounts := []Account{ + { + ID: 38001, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 10, + GroupIDs: []int64{groupID}, + }, + { + // 会话粘连绑定指向的账号已被移出请求分组(绑定 TTL 内账号改组的场景)。 + ID: 38002, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + GroupIDs: []int64{otherGroupID}, + }, + } + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.LBTopK = 2 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 0.7 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate = 0.8 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT = 0.5 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.SessionSticky = 3 + cache := &schedulerTestGatewayCache{sessionBindings: map[string]int64{ + "openai:session_weighted_out_of_group": 38002, + }} + concurrencyCache := schedulerTestConcurrencyCache{ + acquireResults: map[int64]bool{38001: false, 38002: true}, + } + svc := &OpenAIGatewayService{ + accountRepo: schedulerGroupAwareOpenAIAccountRepo{schedulerTestOpenAIAccountRepo{accounts: accounts}}, + cache: cache, + cfg: cfg, + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "true"), + concurrencyService: NewConcurrencyService(concurrencyCache), + } + + selection, decision, err := svc.SelectAccountWithScheduler( + ctx, + &groupID, + "", + "session_weighted_out_of_group", + "gpt-5.1", + nil, + OpenAIUpstreamTransportAny, + false, + ) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + // 组内唯一候选 38001 满并发:必须返回其等待计划,绝不能把请求泄漏到组外的粘连账号 38002。 + require.Equal(t, int64(38001), selection.Account.ID) + require.False(t, selection.Acquired) + require.NotNil(t, selection.WaitPlan) + require.Equal(t, int64(38001), selection.WaitPlan.AccountID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) + // 失效的粘连绑定应被清理,避免后续请求反复走同一条泄漏路径。 + require.Positive(t, cache.deletedSessions["openai:session_weighted_out_of_group"]) +} + +func TestOpenAIGatewayService_SelectAccountWithScheduler_SubscriptionPriorityWaitsOnBusySubscriptionWhenRegularUnusable(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + + ctx := context.Background() + groupID := int64(101091) + accounts := []Account{ + { + // 订阅账号:支持 compact,但并发已满(busy-but-waitable)。 + ID: 38011, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + GroupIDs: []int64{groupID}, + Credentials: map[string]any{"plan_type": "team"}, + Extra: map[string]any{"openai_compact_supported": true}, + }, + { + // 常规账号:明确不支持 compact,无法服务本次请求。 + ID: 38012, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 9, + GroupIDs: []int64{groupID}, + Extra: map[string]any{"openai_compact_supported": false}, + }, + } + concurrencyCache := schedulerTestConcurrencyCache{ + acquireResults: map[int64]bool{38011: false, 38012: true}, + } + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts}, + cache: &schedulerTestGatewayCache{}, + cfg: newSchedulerTestSubscriptionPriorityConfig(), + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "", "true"), + concurrencyService: NewConcurrencyService(concurrencyCache), + } + + selection, decision, err := svc.SelectAccountWithScheduler( + ctx, + &groupID, + "", + "session_subscription_wait", + "gpt-5.1", + nil, + OpenAIUpstreamTransportAny, + true, + ) + // 常规池无可用候选时,忙碌的订阅账号应产生等待计划,而不是直接返回 no available accounts。 + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(38011), selection.Account.ID) + require.False(t, selection.Acquired) + require.NotNil(t, selection.WaitPlan) + require.Equal(t, int64(38011), selection.WaitPlan.AccountID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) +} diff --git a/backend/internal/service/openai_apikey_responses_probe.go b/backend/internal/service/openai_apikey_responses_probe.go index 64f963ab9b..10cf050029 100644 --- a/backend/internal/service/openai_apikey_responses_probe.go +++ b/backend/internal/service/openai_apikey_responses_probe.go @@ -149,6 +149,9 @@ func (s *AccountTestService) ProbeOpenAIAPIKeyResponsesSupport(ctx context.Conte req.Header.Set("Authorization", "Bearer "+apiKey) req.Header.Set("Accept", "application/json") + // 账号级请求头覆写:能力探测与真实转发保持一致的最终头 + account.ApplyHeaderOverrides(req.Header) + proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { proxyURL = account.Proxy.URL() diff --git a/backend/internal/service/openai_codex_transform.go b/backend/internal/service/openai_codex_transform.go index c33ac97f0f..0666293deb 100644 --- a/backend/internal/service/openai_codex_transform.go +++ b/backend/internal/service/openai_codex_transform.go @@ -9,6 +9,9 @@ import ( ) var codexModelMap = map[string]string{ + "gpt-5.6-sol": "gpt-5.6-sol", + "gpt-5.6-terra": "gpt-5.6-terra", + "gpt-5.6-luna": "gpt-5.6-luna", "gpt-5.5": "gpt-5.5", "gpt-5.5-pro": "gpt-5.5-pro", "codex-auto-review": "codex-auto-review", @@ -54,6 +57,9 @@ var codexVersionModelPrefixes = []struct { prefix string target string }{ + {prefix: "gpt-5.6-sol", target: "gpt-5.6-sol"}, + {prefix: "gpt-5.6-terra", target: "gpt-5.6-terra"}, + {prefix: "gpt-5.6-luna", target: "gpt-5.6-luna"}, {prefix: "gpt-5.3-codex-spark", target: "gpt-5.3-codex-spark"}, {prefix: "gpt-5.3-codex", target: "gpt-5.3-codex"}, {prefix: "gpt-5.4-mini", target: "gpt-5.4-mini"}, diff --git a/backend/internal/service/openai_embeddings.go b/backend/internal/service/openai_embeddings.go index 0fb3fff1f7..fb2dc5ccbb 100644 --- a/backend/internal/service/openai_embeddings.go +++ b/backend/internal/service/openai_embeddings.go @@ -82,6 +82,9 @@ func (s *OpenAIGatewayService) ForwardEmbeddings( upstreamReq.Header.Set("user-agent", customUA) } + // 账号级请求头覆写(仅 openai api_key 账号启用时生效) + account.ApplyHeaderOverrides(upstreamReq.Header) + proxyURL := "" if account.Proxy != nil { proxyURL = account.Proxy.URL() diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go index 6bcb6718b7..348213a992 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw.go @@ -166,6 +166,9 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( upstreamReq.Header.Set("user-agent", "sub2api-grok/1.0") } + // 账号级请求头覆写(仅 openai api_key 账号启用时生效) + account.ApplyHeaderOverrides(upstreamReq.Header) + // 6. Send request proxyURL := "" if account.Proxy != nil { diff --git a/backend/internal/service/openai_gateway_count_tokens.go b/backend/internal/service/openai_gateway_count_tokens.go index 4a01b143e9..7518a6073a 100644 --- a/backend/internal/service/openai_gateway_count_tokens.go +++ b/backend/internal/service/openai_gateway_count_tokens.go @@ -231,6 +231,9 @@ func (s *OpenAIGatewayService) buildInputTokensUpstreamRequest( } } + // 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op) + account.ApplyHeaderOverrides(req.Header) + return req, nil } diff --git a/backend/internal/service/openai_gateway_responses_chat_fallback.go b/backend/internal/service/openai_gateway_responses_chat_fallback.go index d33df4c19d..c499bec778 100644 --- a/backend/internal/service/openai_gateway_responses_chat_fallback.go +++ b/backend/internal/service/openai_gateway_responses_chat_fallback.go @@ -138,6 +138,9 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions( upstreamReq.Header.Set("user-agent", customUA) } + // 账号级请求头覆写(仅 openai api_key 账号启用时生效) + account.ApplyHeaderOverrides(upstreamReq.Header) + proxyURL := "" if account.Proxy != nil { proxyURL = account.Proxy.URL() diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index dd3d48aed0..f13c44f3a6 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -3783,6 +3783,9 @@ func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough( req.Header.Set("content-type", "application/json") } + // 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op) + account.ApplyHeaderOverrides(req.Header) + return req, nil } @@ -4568,6 +4571,9 @@ func (s *OpenAIGatewayService) buildUpstreamRequest(ctx context.Context, c *gin. req.Header.Set("content-type", "application/json") } + // 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op) + account.ApplyHeaderOverrides(req.Header) + return req, nil } diff --git a/backend/internal/service/openai_images.go b/backend/internal/service/openai_images.go index 7081653d80..09472fbaf1 100644 --- a/backend/internal/service/openai_images.go +++ b/backend/internal/service/openai_images.go @@ -760,6 +760,8 @@ func (s *OpenAIGatewayService) buildOpenAIImagesRequest( if strings.TrimSpace(contentType) != "" { req.Header.Set("Content-Type", contentType) } + // 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op) + account.ApplyHeaderOverrides(req.Header) return req, nil } diff --git a/backend/internal/service/openai_model_alias.go b/backend/internal/service/openai_model_alias.go index ac2a8cf942..4e3d3b2d9a 100644 --- a/backend/internal/service/openai_model_alias.go +++ b/backend/internal/service/openai_model_alias.go @@ -65,6 +65,12 @@ func normalizeKnownOpenAICodexModel(model string) string { } switch { + case strings.Contains(normalized, "gpt-5.6-sol"): + return "gpt-5.6-sol" + case strings.Contains(normalized, "gpt-5.6-terra"): + return "gpt-5.6-terra" + case strings.Contains(normalized, "gpt-5.6-luna"): + return "gpt-5.6-luna" case strings.Contains(normalized, "gpt-5.5-pro"): return "gpt-5.5-pro" case strings.Contains(normalized, "gpt-5.5"): diff --git a/backend/internal/service/openai_tool_continuation.go b/backend/internal/service/openai_tool_continuation.go index 6515c0c4e5..a507213701 100644 --- a/backend/internal/service/openai_tool_continuation.go +++ b/backend/internal/service/openai_tool_continuation.go @@ -215,6 +215,84 @@ func ValidateFunctionCallOutputContextBytes(body []byte) FunctionCallOutputValid return result } +// ToolCallOutputContextCoverage 描述 input 中工具输出与可重建上下文的覆盖关系, +// 用于判断剥离 previous_response_id 后上游能否仅凭 input 重建工具续链。 +type ToolCallOutputContextCoverage struct { + HasFunctionCallOutput bool + // ContextCoversAllCallIDs 表示每个工具输出的 call_id 都能在 input 内找到 + // 同 call_id 的工具调用上下文项或同 id 的 item_reference,且不存在缺失 call_id 的输出。 + // 任一输出无法由 input 自身重建时为 false,此时剥离 previous_response_id 会导致 + // 上游以 "No tool call found for function call output" 拒绝请求。 + ContextCoversAllCallIDs bool +} + +// AnalyzeToolCallOutputContextCoverageBytes 全量扫描 input,按 call_id 精确匹配工具输出 +// 与可重建上下文。不能复用 ValidateFunctionCallOutputContextBytes 的 HasToolCallContext: +// 该标志只代表"存在某一个上下文项",部分覆盖的续链仍会被上游拒绝。 +func AnalyzeToolCallOutputContextCoverageBytes(body []byte) ToolCallOutputContextCoverage { + coverage := ToolCallOutputContextCoverage{} + if len(body) == 0 { + return coverage + } + input := parseRawJSONView(body).Get("input") + if !input.IsArray() { + return coverage + } + + missingCallID := false + var outputCallIDs map[string]struct{} + var contextIDs map[string]struct{} + input.ForEach(func(_, item gjson.Result) bool { + if !item.IsObject() { + return true + } + itemType := item.Get("type").String() + switch { + case isCodexToolCallOutputItemType(itemType): + coverage.HasFunctionCallOutput = true + callID := strings.TrimSpace(item.Get("call_id").String()) + if callID == "" { + missingCallID = true + return true + } + if outputCallIDs == nil { + outputCallIDs = make(map[string]struct{}) + } + outputCallIDs[callID] = struct{}{} + case isCodexToolCallContextItemType(itemType): + callID := strings.TrimSpace(item.Get("call_id").String()) + if callID == "" { + return true + } + if contextIDs == nil { + contextIDs = make(map[string]struct{}) + } + contextIDs[callID] = struct{}{} + case itemType == "item_reference": + idValue := strings.TrimSpace(item.Get("id").String()) + if idValue == "" { + return true + } + if contextIDs == nil { + contextIDs = make(map[string]struct{}) + } + contextIDs[idValue] = struct{}{} + } + return true + }) + + if !coverage.HasFunctionCallOutput || missingCallID { + return coverage + } + for callID := range outputCallIDs { + if _, ok := contextIDs[callID]; !ok { + return coverage + } + } + coverage.ContextCoversAllCallIDs = true + return coverage +} + // ValidateFunctionCallOutputContext 为 handler 提供低开销校验结果: // 1) 无工具输出直接返回 // 2) 若已存在工具调用上下文则提前返回 diff --git a/backend/internal/service/openai_tool_continuation_test.go b/backend/internal/service/openai_tool_continuation_test.go index 4610652b6c..569d89eff0 100644 --- a/backend/internal/service/openai_tool_continuation_test.go +++ b/backend/internal/service/openai_tool_continuation_test.go @@ -184,3 +184,109 @@ func TestValidateFunctionCallOutputContextBytesMatchesMapValidation(t *testing.T }) } } + +func TestAnalyzeToolCallOutputContextCoverageBytes(t *testing.T) { + cases := []struct { + name string + body map[string]any + hasOutput bool + coversAllIDs bool + }{ + { + name: "no_input", + body: map[string]any{"model": "gpt-5.1"}, + hasOutput: false, + coversAllIDs: false, + }, + { + name: "no_tool_output", + body: map[string]any{"input": []any{ + map[string]any{"type": "message", "content": "hi"}, + }}, + hasOutput: false, + coversAllIDs: false, + }, + { + name: "all_outputs_covered_by_context", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call", "call_id": "call_a"}, + map[string]any{"type": "function_call_output", "call_id": "call_a"}, + }}, + hasOutput: true, + coversAllIDs: true, + }, + { + name: "all_outputs_covered_by_item_reference", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call_output", "call_id": "call_a"}, + map[string]any{"type": "item_reference", "id": "call_a"}, + }}, + hasOutput: true, + coversAllIDs: true, + }, + { + // 关键回归用例:input 内存在某一个上下文项,但另一个输出的 call_id + // 只能由上游会话链(previous_response_id)解析——不可剥离。 + name: "partial_coverage_not_movable", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call", "call_id": "call_a"}, + map[string]any{"type": "function_call_output", "call_id": "call_a"}, + map[string]any{"type": "function_call_output", "call_id": "call_b"}, + }}, + hasOutput: true, + coversAllIDs: false, + }, + { + name: "unrelated_context_does_not_cover", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call", "call_id": "call_x"}, + map[string]any{"type": "function_call_output", "call_id": "call_b"}, + }}, + hasOutput: true, + coversAllIDs: false, + }, + { + name: "output_missing_call_id_not_movable", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call", "call_id": "call_a"}, + map[string]any{"type": "function_call_output"}, + map[string]any{"type": "function_call_output", "call_id": "call_a"}, + }}, + hasOutput: true, + coversAllIDs: false, + }, + { + name: "mixed_context_and_reference_cover_all", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call", "call_id": "call_a"}, + map[string]any{"type": "function_call_output", "call_id": "call_a"}, + map[string]any{"type": "function_call_output", "call_id": "call_b"}, + map[string]any{"type": "item_reference", "id": "call_b"}, + }}, + hasOutput: true, + coversAllIDs: true, + }, + { + name: "all_codex_output_types_covered", + body: map[string]any{"input": []any{ + map[string]any{"type": "tool_search_output", "call_id": "call_s"}, + map[string]any{"type": "tool_search_call", "call_id": "call_s"}, + map[string]any{"type": "mcp_tool_call_output", "call_id": "call_m"}, + map[string]any{"type": "mcp_tool_call", "call_id": "call_m"}, + }}, + hasOutput: true, + coversAllIDs: true, + }, + } + + for _, tt := range cases { + t.Run(tt.name, func(t *testing.T) { + bodyBytes, err := json.Marshal(tt.body) + require.NoError(t, err) + + coverage := AnalyzeToolCallOutputContextCoverageBytes(bodyBytes) + require.Equal(t, tt.hasOutput, coverage.HasFunctionCallOutput, "HasFunctionCallOutput") + require.Equal(t, tt.coversAllIDs, coverage.ContextCoversAllCallIDs, "ContextCoversAllCallIDs") + }) + } +} diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go index fcc4b98064..bbca9776ab 100644 --- a/backend/internal/service/openai_ws_forwarder.go +++ b/backend/internal/service/openai_ws_forwarder.go @@ -1183,6 +1183,10 @@ func (s *OpenAIGatewayService) buildOpenAIWSHeaders( headers.Set("user-agent", codexCLIUserAgent) } + // 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op)。 + // 覆盖所有 WS 模式(ctx_pool/dedicated/passthrough)的握手头。 + account.ApplyHeaderOverrides(headers) + return headers, sessionResolution, nil } diff --git a/backend/internal/service/payment_amounts.go b/backend/internal/service/payment_amounts.go index a7f620d33e..2fd00c5957 100644 --- a/backend/internal/service/payment_amounts.go +++ b/backend/internal/service/payment_amounts.go @@ -16,6 +16,15 @@ func normalizeBalanceRechargeMultiplier(multiplier float64) float64 { return multiplier } +// normalizeSubscriptionUSDToCNYRate 将非法值归一为 0(换算关闭)。 +// 与余额倍率不同,0 是合法状态:表示订阅保持 price 直付的存量行为。 +func normalizeSubscriptionUSDToCNYRate(rate float64) float64 { + if math.IsNaN(rate) || math.IsInf(rate, 0) || rate < 0 { + return 0 + } + return rate +} + func calculateCreditedBalance(paymentAmount, multiplier float64) float64 { return decimal.NewFromFloat(paymentAmount). Mul(decimal.NewFromFloat(normalizeBalanceRechargeMultiplier(multiplier))). diff --git a/backend/internal/service/payment_config_limits.go b/backend/internal/service/payment_config_limits.go index 45b24bfce7..202eea9f26 100644 --- a/backend/internal/service/payment_config_limits.go +++ b/backend/internal/service/payment_config_limits.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "fmt" + "strings" dbent "github.com/Wei-Shaw/sub2api/ent" "github.com/Wei-Shaw/sub2api/ent/paymentproviderinstance" @@ -31,6 +32,7 @@ func (s *PaymentConfigService) GetAvailableMethodLimits(ctx context.Context) (*M continue } ml := pcAggregateMethodLimits(pt, insts) + ml.DisplayName = s.pcAggregateMethodDisplayName(pt, insts) ml.Currency = currency resp.Methods[ml.PaymentType] = ml } @@ -93,6 +95,7 @@ func (s *PaymentConfigService) GetMethodLimits(ctx context.Context, types []stri continue } ml := pcAggregateMethodLimits(pt, matching) + ml.DisplayName = s.pcAggregateMethodDisplayName(pt, matching) ml.Currency = currency result = append(result, ml) } @@ -163,6 +166,53 @@ func (s *PaymentConfigService) pcInstancePaymentCurrency(inst *dbent.PaymentProv return paymentProviderConfigCurrency(inst.ProviderKey, cfg) } +type easyPayCustomMethodDisplayConfig struct { + Type string `json:"type"` + DisplayName string `json:"displayName"` +} + +func (s *PaymentConfigService) pcAggregateMethodDisplayName(pt string, instances []*dbent.PaymentProviderInstance) string { + pt = strings.TrimSpace(pt) + if pt == "" { + return "" + } + for _, inst := range instances { + displayName := s.pcInstanceEasyPayCustomMethodDisplayName(inst, pt) + if displayName != "" { + return displayName + } + } + return "" +} + +func (s *PaymentConfigService) pcInstanceEasyPayCustomMethodDisplayName(inst *dbent.PaymentProviderInstance, pt string) string { + if inst == nil || inst.ProviderKey != payment.TypeEasyPay { + return "" + } + cfg := map[string]string{} + if s != nil { + decrypted, err := s.decryptConfig(inst.Config) + if err == nil && decrypted != nil { + cfg = decrypted + } + } + raw := strings.TrimSpace(cfg["customMethods"]) + if raw == "" { + return "" + } + + var methods []easyPayCustomMethodDisplayConfig + if err := json.Unmarshal([]byte(raw), &methods); err != nil { + return "" + } + for _, method := range methods { + if strings.TrimSpace(method.Type) == pt { + return strings.TrimSpace(method.DisplayName) + } + } + return "" +} + // pcGroupByPaymentType groups instances by user-facing payment type. // For Stripe providers, ALL sub-types (card, link, alipay, wxpay) map to "stripe" // because the user sees a single "Stripe" button, not individual sub-methods. diff --git a/backend/internal/service/payment_config_limits_test.go b/backend/internal/service/payment_config_limits_test.go index c0aa2b27a5..a70bc90a29 100644 --- a/backend/internal/service/payment_config_limits_test.go +++ b/backend/internal/service/payment_config_limits_test.go @@ -255,6 +255,28 @@ func TestGetAvailableMethodLimitsOmitsMixedCurrencyMethod(t *testing.T) { require.Equal(t, "PAYMENT_METHOD_CURRENCY_CONFLICT", appErr.Reason) } +func TestGetAvailableMethodLimitsIncludesEasyPayCustomMethodDisplayName(t *testing.T) { + ctx := context.Background() + client := newPaymentConfigServiceTestClient(t) + + _, err := client.PaymentProviderInstance.Create(). + SetProviderKey(payment.TypeEasyPay). + SetName("EasyPay Custom"). + SetConfig(`{"customMethods":"[{\"type\":\"ldc\",\"upstreamType\":\"ldc\",\"displayName\":\"LDC Pay\"}]"}`). + SetSupportedTypes("alipay,wxpay,ldc"). + SetEnabled(true). + Save(ctx) + require.NoError(t, err) + + svc := &PaymentConfigService{entClient: client} + resp, err := svc.GetAvailableMethodLimits(ctx) + require.NoError(t, err) + + limits, ok := resp.Methods["ldc"] + require.True(t, ok, "expected custom EasyPay method limits to be visible") + require.Equal(t, "LDC Pay", limits.DisplayName) +} + func TestPcComputeGlobalRange(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/payment_config_providers.go b/backend/internal/service/payment_config_providers.go index 7e92558568..d1bf2de7aa 100644 --- a/backend/internal/service/payment_config_providers.go +++ b/backend/internal/service/payment_config_providers.go @@ -5,6 +5,7 @@ import ( "encoding/json" "fmt" "log/slog" + "regexp" "strconv" "strings" @@ -185,6 +186,11 @@ func (s *PaymentConfigService) CreateProviderInstance(ctx context.Context, req C if err := validateProviderRequest(req.ProviderKey, req.Name, typesStr); err != nil { return nil, err } + if req.ProviderKey == payment.TypeEasyPay { + if err := validateEasyPayCustomMethods(req.Config, typesStr); err != nil { + return nil, err + } + } if err := s.validateVisibleMethodEnablementConflicts(ctx, 0, req.ProviderKey, typesStr, req.Enabled); err != nil { return nil, err } @@ -217,6 +223,67 @@ func validateProviderRequest(providerKey, name, supportedTypes string) error { return nil } +var easyPayCustomMethodCodePattern = regexp.MustCompile(`^[a-z0-9_-]+$`) + +type easyPayCustomMethodConfig struct { + Type string `json:"type"` + UpstreamType string `json:"upstreamType"` + DisplayName string `json:"displayName"` +} + +func validateEasyPayCustomMethods(config map[string]string, supportedTypes string) error { + if config == nil { + config = map[string]string{} + } + raw := strings.TrimSpace(config["customMethods"]) + methods := make([]easyPayCustomMethodConfig, 0) + if raw != "" { + if err := json.Unmarshal([]byte(raw), &methods); err != nil { + return infraerrors.BadRequest("VALIDATION_ERROR", "customMethods must be a JSON array") + } + } + + customTypes := make(map[string]struct{}, len(methods)) + for _, method := range methods { + method.Type = strings.TrimSpace(method.Type) + method.UpstreamType = strings.TrimSpace(method.UpstreamType) + if method.Type == "" || method.UpstreamType == "" { + return infraerrors.BadRequest("VALIDATION_ERROR", "customMethods upstreamType is required") + } + if !easyPayCustomMethodCodePattern.MatchString(method.Type) { + return infraerrors.BadRequest("VALIDATION_ERROR", "customMethods type may only contain lowercase letters, digits, underscores, and hyphens") + } + if !easyPayCustomMethodCodePattern.MatchString(method.UpstreamType) { + return infraerrors.BadRequest("VALIDATION_ERROR", "customMethods upstreamType may only contain lowercase letters, digits, underscores, and hyphens") + } + if easyPayCustomMethodTypeConflictsWithBuiltin(method.Type) { + return infraerrors.BadRequest("VALIDATION_ERROR", "customMethods type cannot start with alipay or wxpay") + } + if _, exists := customTypes[method.Type]; exists { + return infraerrors.BadRequest("VALIDATION_ERROR", "duplicate customMethods type") + } + customTypes[method.Type] = struct{}{} + } + + for _, supportedType := range splitTypes(supportedTypes) { + supportedType = strings.TrimSpace(supportedType) + if supportedType == "" || supportedType == payment.TypeAlipay || supportedType == payment.TypeWxpay { + continue + } + if !easyPayCustomMethodCodePattern.MatchString(supportedType) { + return infraerrors.BadRequest("VALIDATION_ERROR", fmt.Sprintf("supported EasyPay custom type %s may only contain lowercase letters, digits, underscores, and hyphens", supportedType)) + } + if _, exists := customTypes[supportedType]; !exists { + return infraerrors.BadRequest("VALIDATION_ERROR", fmt.Sprintf("supported EasyPay custom type %s has no customMethods mapping", supportedType)) + } + } + return nil +} + +func easyPayCustomMethodTypeConflictsWithBuiltin(methodType string) bool { + return strings.HasPrefix(methodType, payment.TypeAlipay) || strings.HasPrefix(methodType, payment.TypeWxpay) +} + // UpdateProviderInstance updates a provider instance by ID (patch semantics). // NOTE: This function exceeds 30 lines due to per-field nil-check patch update // boilerplate and pending-order safety checks. @@ -279,6 +346,18 @@ func (s *PaymentConfigService) UpdateProviderInstance(ctx context.Context, id in WithMetadata(map[string]string{"count": strconv.Itoa(count)}) } } + configToValidate := mergedConfig + if configToValidate == nil { + configToValidate, err = s.decryptConfig(current.Config) + if err != nil { + return nil, fmt.Errorf("decrypt existing config: %w", err) + } + } + if current.ProviderKey == payment.TypeEasyPay { + if err := validateEasyPayCustomMethods(configToValidate, nextSupportedTypes); err != nil { + return nil, err + } + } // Validate merged config when the instance will end up enabled. // This surfaces provider-level errors (e.g. wxpay missing certSerial) at save time, // so admins see them in the dialog instead of only when an order is created. @@ -287,13 +366,6 @@ func (s *PaymentConfigService) UpdateProviderInstance(ctx context.Context, id in finalEnabled = *req.Enabled } if finalEnabled { - configToValidate := mergedConfig - if configToValidate == nil { - configToValidate, err = s.decryptConfig(current.Config) - if err != nil { - return nil, fmt.Errorf("decrypt existing config: %w", err) - } - } if err := s.validateProviderConfig(current.ProviderKey, configToValidate); err != nil { return nil, err } diff --git a/backend/internal/service/payment_config_providers_test.go b/backend/internal/service/payment_config_providers_test.go index 43708de73d..74fd2a3467 100644 --- a/backend/internal/service/payment_config_providers_test.go +++ b/backend/internal/service/payment_config_providers_test.go @@ -114,6 +114,92 @@ func TestValidateProviderRequest(t *testing.T) { } } +func TestValidateEasyPayCustomMethods(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + config map[string]string + supportedTypes string + wantErr string + }{ + { + name: "valid custom methods", + config: map[string]string{"customMethods": `[{"type":"ldc","upstreamType":"epay","displayName":"LDC"}]`}, + supportedTypes: "alipay,wxpay,ldc", + }, + { + name: "malformed custom methods json", + config: map[string]string{"customMethods": `not-json`}, + supportedTypes: "alipay,wxpay,ldc", + wantErr: "customMethods must be a JSON array", + }, + { + name: "missing upstream type", + config: map[string]string{"customMethods": `[{"type":"ldc","displayName":"LDC"}]`}, + supportedTypes: "alipay,wxpay,ldc", + wantErr: "customMethods upstreamType is required", + }, + { + name: "duplicate custom type", + config: map[string]string{"customMethods": `[{"type":"ldc","upstreamType":"epay"},{"type":"ldc","upstreamType":"epay2"}]`}, + supportedTypes: "alipay,wxpay,ldc", + wantErr: "duplicate customMethods type", + }, + { + name: "custom type must already be lowercase", + config: map[string]string{"customMethods": `[{"type":"LDC","upstreamType":"epay"}]`}, + supportedTypes: "alipay,wxpay,ldc", + wantErr: "customMethods type may only contain lowercase letters", + }, + { + name: "upstream type must already be lowercase", + config: map[string]string{"customMethods": `[{"type":"ldc","upstreamType":"ALIPAY"}]`}, + supportedTypes: "alipay,wxpay,ldc", + wantErr: "customMethods upstreamType may only contain lowercase letters", + }, + { + name: "custom type uses alipay prefix", + config: map[string]string{"customMethods": `[{"type":"alipay_hk","upstreamType":"hkpay"}]`}, + supportedTypes: "alipay,wxpay,alipay_hk", + wantErr: "customMethods type cannot start with alipay or wxpay", + }, + { + name: "custom type uses wxpay prefix", + config: map[string]string{"customMethods": `[{"type":"wxpay_usdt","upstreamType":"usdt"}]`}, + supportedTypes: "alipay,wxpay,wxpay_usdt", + wantErr: "customMethods type cannot start with alipay or wxpay", + }, + { + name: "supported custom type missing mapping", + config: map[string]string{"customMethods": `[{"type":"ldc","upstreamType":"epay"}]`}, + supportedTypes: "alipay,wxpay,ldc,usdt_trc20", + wantErr: "supported EasyPay custom type usdt_trc20 has no customMethods mapping", + }, + { + name: "supported custom type must already be lowercase", + config: map[string]string{"customMethods": `[{"type":"ldc","upstreamType":"epay"}]`}, + supportedTypes: "alipay,wxpay,LDC", + wantErr: "supported EasyPay custom type LDC may only contain lowercase letters", + }, + } + + for _, tc := range tests { + tc := tc + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + err := validateEasyPayCustomMethods(tc.config, tc.supportedTypes) + if tc.wantErr == "" { + require.NoError(t, err) + return + } + require.Error(t, err) + require.Contains(t, err.Error(), tc.wantErr) + }) + } +} + func TestIsSensitiveProviderConfigField(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/payment_config_service.go b/backend/internal/service/payment_config_service.go index 022b1b0156..52a7ddc67e 100644 --- a/backend/internal/service/payment_config_service.go +++ b/backend/internal/service/payment_config_service.go @@ -24,17 +24,20 @@ const ( SettingLoadBalanceStrategy = "LOAD_BALANCE_STRATEGY" SettingBalancePayDisabled = "BALANCE_PAYMENT_DISABLED" SettingBalanceRechargeMult = "BALANCE_RECHARGE_MULTIPLIER" - SettingRechargeFeeRate = "RECHARGE_FEE_RATE" - SettingProductNamePrefix = "PRODUCT_NAME_PREFIX" - SettingProductNameSuffix = "PRODUCT_NAME_SUFFIX" - SettingHelpImageURL = "PAYMENT_HELP_IMAGE_URL" - SettingHelpText = "PAYMENT_HELP_TEXT" - SettingCancelRateLimitOn = "CANCEL_RATE_LIMIT_ENABLED" - SettingCancelRateLimitMax = "CANCEL_RATE_LIMIT_MAX" - SettingCancelWindowSize = "CANCEL_RATE_LIMIT_WINDOW" - SettingCancelWindowUnit = "CANCEL_RATE_LIMIT_UNIT" - SettingCancelWindowMode = "CANCEL_RATE_LIMIT_WINDOW_MODE" - SettingAlipayForceQRCode = "ALIPAY_FORCE_QRCODE" + // SettingSubscriptionUSDToCNYRate 是订阅 CNY 换算汇率(1 USD = X CNY)。 + // 0/未配置 = 关闭换算(订阅按 price 数值直付),显式配置后 CNY 通道订阅按 price × rate 收款。 + SettingSubscriptionUSDToCNYRate = "SUBSCRIPTION_USD_TO_CNY_RATE" + SettingRechargeFeeRate = "RECHARGE_FEE_RATE" + SettingProductNamePrefix = "PRODUCT_NAME_PREFIX" + SettingProductNameSuffix = "PRODUCT_NAME_SUFFIX" + SettingHelpImageURL = "PAYMENT_HELP_IMAGE_URL" + SettingHelpText = "PAYMENT_HELP_TEXT" + SettingCancelRateLimitOn = "CANCEL_RATE_LIMIT_ENABLED" + SettingCancelRateLimitMax = "CANCEL_RATE_LIMIT_MAX" + SettingCancelWindowSize = "CANCEL_RATE_LIMIT_WINDOW" + SettingCancelWindowUnit = "CANCEL_RATE_LIMIT_UNIT" + SettingCancelWindowMode = "CANCEL_RATE_LIMIT_WINDOW_MODE" + SettingAlipayForceQRCode = "ALIPAY_FORCE_QRCODE" ) // Default values for payment configuration settings. @@ -54,13 +57,15 @@ type PaymentConfig struct { EnabledTypes []string `json:"enabled_payment_types"` BalanceDisabled bool `json:"balance_disabled"` BalanceRechargeMultiplier float64 `json:"balance_recharge_multiplier"` - RechargeFeeRate float64 `json:"recharge_fee_rate"` - LoadBalanceStrategy string `json:"load_balance_strategy"` - ProductNamePrefix string `json:"product_name_prefix"` - ProductNameSuffix string `json:"product_name_suffix"` - HelpImageURL string `json:"help_image_url"` - HelpText string `json:"help_text"` - StripePublishableKey string `json:"stripe_publishable_key,omitempty"` + // SubscriptionUSDToCNYRate 为 0 时订阅换算关闭(兼容存量行为)。 + SubscriptionUSDToCNYRate float64 `json:"subscription_usd_to_cny_rate"` + RechargeFeeRate float64 `json:"recharge_fee_rate"` + LoadBalanceStrategy string `json:"load_balance_strategy"` + ProductNamePrefix string `json:"product_name_prefix"` + ProductNameSuffix string `json:"product_name_suffix"` + HelpImageURL string `json:"help_image_url"` + HelpText string `json:"help_text"` + StripePublishableKey string `json:"stripe_publishable_key,omitempty"` // Cancel rate limit settings CancelRateLimitEnabled bool `json:"cancel_rate_limit_enabled"` @@ -84,6 +89,7 @@ type UpdatePaymentConfigRequest struct { EnabledTypes []string `json:"enabled_payment_types"` BalanceDisabled *bool `json:"balance_disabled"` BalanceRechargeMultiplier *float64 `json:"balance_recharge_multiplier"` + SubscriptionUSDToCNYRate *float64 `json:"subscription_usd_to_cny_rate"` RechargeFeeRate *float64 `json:"recharge_fee_rate"` LoadBalanceStrategy *string `json:"load_balance_strategy"` ProductNamePrefix *string `json:"product_name_prefix"` @@ -110,6 +116,7 @@ type UpdatePaymentConfigRequest struct { // MethodLimits holds per-payment-type limits. type MethodLimits struct { PaymentType string `json:"payment_type"` + DisplayName string `json:"display_name,omitempty"` Currency string `json:"currency"` FeeRate float64 `json:"fee_rate"` DailyLimit float64 `json:"daily_limit"` @@ -204,7 +211,7 @@ func (s *PaymentConfigService) GetPaymentConfig(ctx context.Context) (*PaymentCo keys := []string{ SettingPaymentEnabled, SettingMinRechargeAmount, SettingMaxRechargeAmount, SettingDailyRechargeLimit, SettingOrderTimeoutMinutes, SettingMaxPendingOrders, - SettingEnabledPaymentTypes, SettingBalancePayDisabled, SettingBalanceRechargeMult, SettingRechargeFeeRate, SettingLoadBalanceStrategy, + SettingEnabledPaymentTypes, SettingBalancePayDisabled, SettingBalanceRechargeMult, SettingSubscriptionUSDToCNYRate, SettingRechargeFeeRate, SettingLoadBalanceStrategy, SettingProductNamePrefix, SettingProductNameSuffix, SettingHelpImageURL, SettingHelpText, SettingCancelRateLimitOn, SettingCancelRateLimitMax, @@ -233,6 +240,7 @@ func (s *PaymentConfigService) parsePaymentConfig(vals map[string]string) *Payme MaxPendingOrders: pcParseInt(vals[SettingMaxPendingOrders], defaultMaxPendingOrders), BalanceDisabled: vals[SettingBalancePayDisabled] == "true", BalanceRechargeMultiplier: normalizeBalanceRechargeMultiplier(pcParseFloat(vals[SettingBalanceRechargeMult], defaultBalanceRechargeMultiplier)), + SubscriptionUSDToCNYRate: normalizeSubscriptionUSDToCNYRate(pcParseFloat(vals[SettingSubscriptionUSDToCNYRate], 0)), RechargeFeeRate: pcParseFloat(vals[SettingRechargeFeeRate], 0), LoadBalanceStrategy: vals[SettingLoadBalanceStrategy], ProductNamePrefix: vals[SettingProductNamePrefix], @@ -294,6 +302,12 @@ func (s *PaymentConfigService) UpdatePaymentConfig(ctx context.Context, req Upda return infraerrors.BadRequest("INVALID_BALANCE_RECHARGE_MULTIPLIER", "balance recharge multiplier must be greater than 0") } } + if req.SubscriptionUSDToCNYRate != nil { + v := *req.SubscriptionUSDToCNYRate + if math.IsNaN(v) || math.IsInf(v, 0) || v < 0 { + return infraerrors.BadRequest("INVALID_SUBSCRIPTION_USD_TO_CNY_RATE", "subscription USD to CNY rate must be 0 (disabled) or a positive number") + } + } if req.RechargeFeeRate != nil { v := *req.RechargeFeeRate if math.IsNaN(v) || math.IsInf(v, 0) || v < 0 || v > 100 { @@ -313,6 +327,7 @@ func (s *PaymentConfigService) UpdatePaymentConfig(ctx context.Context, req Upda SettingMaxPendingOrders: formatPositiveInt(req.MaxPendingOrders), SettingBalancePayDisabled: formatBoolOrEmpty(req.BalanceDisabled), SettingBalanceRechargeMult: formatPositiveFloat(req.BalanceRechargeMultiplier), + SettingSubscriptionUSDToCNYRate: formatPositiveFloatExact(req.SubscriptionUSDToCNYRate), SettingRechargeFeeRate: formatNonNegativeFloat(req.RechargeFeeRate), SettingLoadBalanceStrategy: derefStr(req.LoadBalanceStrategy), SettingProductNamePrefix: derefStr(req.ProductNamePrefix), @@ -352,6 +367,14 @@ func formatPositiveFloat(v *float64) string { return strconv.FormatFloat(*v, 'f', 2, 64) } +// formatPositiveFloatExact 保留完整精度,用于汇率等对小数位敏感的配置。 +func formatPositiveFloatExact(v *float64) string { + if v == nil || *v <= 0 { + return "" // empty → parsePaymentConfig 视为未配置(换算关闭) + } + return strconv.FormatFloat(*v, 'f', -1, 64) +} + func formatNonNegativeFloat(v *float64) string { if v == nil || *v < 0 { return "" diff --git a/backend/internal/service/payment_config_service_test.go b/backend/internal/service/payment_config_service_test.go index f04f4697b1..bfc69d1705 100644 --- a/backend/internal/service/payment_config_service_test.go +++ b/backend/internal/service/payment_config_service_test.go @@ -187,6 +187,23 @@ func TestParsePaymentConfig(t *testing.T) { } }) + t.Run("custom enabled types are preserved", func(t *testing.T) { + t.Parallel() + vals := map[string]string{ + SettingEnabledPaymentTypes: "alipay,ldc,usdt_trc20", + } + cfg := svc.parsePaymentConfig(vals) + want := []string{"alipay", "ldc", "usdt_trc20"} + if len(cfg.EnabledTypes) != len(want) { + t.Fatalf("EnabledTypes len = %d, want %d (%v)", len(cfg.EnabledTypes), len(want), cfg.EnabledTypes) + } + for i := range want { + if cfg.EnabledTypes[i] != want[i] { + t.Fatalf("EnabledTypes[%d] = %q, want %q (full=%v)", i, cfg.EnabledTypes[i], want[i], cfg.EnabledTypes) + } + } + }) + t.Run("empty enabled types string", func(t *testing.T) { t.Parallel() vals := map[string]string{ diff --git a/backend/internal/service/payment_fulfillment_test.go b/backend/internal/service/payment_fulfillment_test.go index b46d6a1fc8..a8c78d713c 100644 --- a/backend/internal/service/payment_fulfillment_test.go +++ b/backend/internal/service/payment_fulfillment_test.go @@ -602,8 +602,8 @@ func TestExecuteSubscriptionFulfillmentAppliesAffiliateRebate(t *testing.T) { SetUserID(user.ID). SetUserEmail(user.Email). SetUserName(user.Username). - SetAmount(120). - SetPayAmount(120). + SetAmount(9.99). + SetPayAmount(71.36). SetFeeRate(0). SetRechargeCode("PAY-SUB-AFFILIATE"). SetOutTradeNo("sub2_subscription_affiliate"). @@ -636,7 +636,7 @@ func TestExecuteSubscriptionFulfillmentAppliesAffiliateRebate(t *testing.T) { } settingSvc := NewSettingService(&paymentFulfillmentSettingRepoStub{values: map[string]string{ SettingKeyAffiliateEnabled: "true", - SettingKeyAffiliateRebateRate: "20", + SettingKeyAffiliateRebateRate: "15", SettingKeyAffiliateRebateFreezeHours: "0", }}, nil) subRepo := newSubscriptionUserSubRepoStub() @@ -659,7 +659,7 @@ func TestExecuteSubscriptionFulfillmentAppliesAffiliateRebate(t *testing.T) { require.Len(t, affiliateRepo.accrueCalls, 1) require.Equal(t, inviterID, affiliateRepo.accrueCalls[0].inviterID) require.Equal(t, user.ID, affiliateRepo.accrueCalls[0].inviteeUserID) - require.Equal(t, 24.0, affiliateRepo.accrueCalls[0].amount) + require.InDelta(t, 1.4985, affiliateRepo.accrueCalls[0].amount, 0.00000001) require.NotNil(t, affiliateRepo.accrueCalls[0].sourceOrderID) require.Equal(t, order.ID, *affiliateRepo.accrueCalls[0].sourceOrderID) require.Equal(t, 1, subRepo.createCalls) @@ -668,8 +668,8 @@ func TestExecuteSubscriptionFulfillmentAppliesAffiliateRebate(t *testing.T) { Where(paymentauditlog.OrderIDEQ(strconv.FormatInt(order.ID, 10)), paymentauditlog.ActionEQ("AFFILIATE_REBATE_APPLIED")). Only(ctx) require.NoError(t, err) - require.Contains(t, applied.Detail, `"baseAmount":120`) - require.Contains(t, applied.Detail, `"rebateAmount":24`) + require.Contains(t, applied.Detail, `"baseAmount":9.99`) + require.Contains(t, applied.Detail, `"rebateAmount":1.4985`) } func TestExecuteSubscriptionFulfillmentDoesNotDuplicateWorkAfterLegacySuccessAudit(t *testing.T) { diff --git a/backend/internal/service/payment_order.go b/backend/internal/service/payment_order.go index 154159b932..04feb8002a 100644 --- a/backend/internal/service/payment_order.go +++ b/backend/internal/service/payment_order.go @@ -16,6 +16,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/payment" "github.com/Wei-Shaw/sub2api/internal/payment/provider" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/shopspring/decimal" ) // --- Order Creation --- @@ -67,8 +68,7 @@ func (s *PaymentService) CreateOrder(ctx context.Context, req CreateOrderRequest return nil, err } } - // 订阅套餐 price 是直付价,余额充值倍率只影响余额充值到账,不参与订阅 pay_amount 计算。 - payAmountStr, payAmount, err := calculateCreateOrderPayAmount(limitAmount, feeRate, methodCurrency) + payAmountStr, payAmount, err := calculateCreateOrderPayAmountForOrderType(limitAmount, feeRate, methodCurrency, req.OrderType, cfg.SubscriptionUSDToCNYRate) if err != nil { return nil, err } @@ -84,7 +84,7 @@ func (s *PaymentService) CreateOrder(ctx context.Context, req CreateOrderRequest selectedCurrency = paymentProviderConfigCurrency(sel.ProviderKey, sel.Config) } if selectedCurrency != methodCurrency { - payAmountStr, payAmount, err = calculateCreateOrderPayAmount(limitAmount, feeRate, selectedCurrency) + payAmountStr, payAmount, err = calculateCreateOrderPayAmountForOrderType(limitAmount, feeRate, selectedCurrency, req.OrderType, cfg.SubscriptionUSDToCNYRate) if err != nil { return nil, err } @@ -630,6 +630,28 @@ func calculateCreateOrderPayAmount(limitAmount, feeRate float64, currency string return payAmountStr, payAmount, nil } +func calculateCreateOrderPayAmountForOrderType(limitAmount, feeRate float64, currency, orderType string, usdToCnyRate float64) (string, float64, error) { + paymentAmount := limitAmount + if orderType == payment.OrderTypeSubscription { + paymentAmount = calculateSubscriptionGatewayBaseAmount(limitAmount, usdToCnyRate, currency) + } + return calculateCreateOrderPayAmount(paymentAmount, feeRate, currency) +} + +// calculateSubscriptionGatewayBaseAmount 计算订阅订单的网关扣款基数。 +// 换算是显式 opt-in:仅当管理员配置了订阅汇率(rate > 0,1 USD = rate CNY) +// 且网关币种为 CNY 时,按 price × rate 换算;未配置时保持 price 直付的存量行为。 +func calculateSubscriptionGatewayBaseAmount(amount, usdToCnyRate float64, currency string) float64 { + rate := normalizeSubscriptionUSDToCNYRate(usdToCnyRate) + if rate <= 0 || currency != payment.DefaultPaymentCurrency { + return amount + } + return decimal.NewFromFloat(amount). + Mul(decimal.NewFromFloat(rate)). + Round(int32(payment.CurrencyMaxFractionDigits(currency))). + InexactFloat64() +} + func validateCreateOrderAmountCurrency(amount float64, currency string) error { amountStr := strconv.FormatFloat(amount, 'f', -1, 64) if _, err := payment.AmountToMinorUnit(amountStr, currency); err != nil { diff --git a/backend/internal/service/payment_order_result_test.go b/backend/internal/service/payment_order_result_test.go index 14192bd6cd..ac439ee6f2 100644 --- a/backend/internal/service/payment_order_result_test.go +++ b/backend/internal/service/payment_order_result_test.go @@ -161,27 +161,66 @@ func TestCalculateCreateOrderPayAmountUsesCurrencyPrecision(t *testing.T) { } } -func TestCalculateCreateOrderPayAmountForSubscriptionKeepsDirectPrice(t *testing.T) { +func TestCalculateCreateOrderPayAmountForSubscriptionConvertsCNYPriceWhenRateConfigured(t *testing.T) { t.Parallel() - amountStr, amount, err := calculateCreateOrderPayAmount(5, 0, "CNY") + amountStr, amount, err := calculateCreateOrderPayAmountForOrderType(9.99, 0, "CNY", payment.OrderTypeSubscription, 7.15) if err != nil { t.Fatalf("unexpected error: %v", err) } - if amountStr != "5.00" || amount != 5 { - t.Fatalf("subscription CNY pay amount = (%q, %v), want (5.00, 5)", amountStr, amount) + if amountStr != "71.43" || amount != 71.43 { + t.Fatalf("subscription CNY pay amount = (%q, %v), want (71.43, 71.43)", amountStr, amount) } } -func TestCalculateCreateOrderPayAmountForSubscriptionAppliesFeeToDirectPrice(t *testing.T) { +func TestCalculateCreateOrderPayAmountForSubscriptionAppliesFeeAfterCNYConversion(t *testing.T) { t.Parallel() - amountStr, amount, err := calculateCreateOrderPayAmount(5, 2.5, "CNY") + amountStr, amount, err := calculateCreateOrderPayAmountForOrderType(9.99, 2.5, "CNY", payment.OrderTypeSubscription, 7.15) if err != nil { t.Fatalf("unexpected error: %v", err) } - if amountStr != "5.13" || amount != 5.13 { - t.Fatalf("subscription CNY pay amount with fee = (%q, %v), want (5.13, 5.13)", amountStr, amount) + if amountStr != "73.22" || amount != 73.22 { + t.Fatalf("subscription CNY pay amount with fee = (%q, %v), want (73.22, 73.22)", amountStr, amount) + } +} + +func TestCalculateCreateOrderPayAmountForSubscriptionKeepsNonCNYPrice(t *testing.T) { + t.Parallel() + + amountStr, amount, err := calculateCreateOrderPayAmountForOrderType(9.99, 0, "USD", payment.OrderTypeSubscription, 7.15) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if amountStr != "9.99" || amount != 9.99 { + t.Fatalf("subscription USD pay amount = (%q, %v), want (9.99, 9.99)", amountStr, amount) + } +} + +// 换算是 opt-in:未配置汇率(rate=0)时,CNY 订阅保持 price 直付的存量行为。 +// 该测试锁住存量部署升级后行为不变的兼容承诺。 +func TestCalculateCreateOrderPayAmountForSubscriptionKeepsDirectPriceWhenRateDisabled(t *testing.T) { + t.Parallel() + + amountStr, amount, err := calculateCreateOrderPayAmountForOrderType(9.99, 0, "CNY", payment.OrderTypeSubscription, 0) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if amountStr != "9.99" || amount != 9.99 { + t.Fatalf("subscription CNY pay amount without rate = (%q, %v), want (9.99, 9.99)", amountStr, amount) + } +} + +// 汇率只作用于订阅订单,余额充值订单不受影响。 +func TestCalculateCreateOrderPayAmountForBalanceIgnoresSubscriptionRate(t *testing.T) { + t.Parallel() + + amountStr, amount, err := calculateCreateOrderPayAmountForOrderType(50, 0, "CNY", payment.OrderTypeBalance, 7.15) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if amountStr != "50.00" || amount != 50 { + t.Fatalf("balance CNY pay amount = (%q, %v), want (50.00, 50)", amountStr, amount) } } diff --git a/backend/internal/service/payment_resume_service_test.go b/backend/internal/service/payment_resume_service_test.go index 7e0adc2de8..17b637fa23 100644 --- a/backend/internal/service/payment_resume_service_test.go +++ b/backend/internal/service/payment_resume_service_test.go @@ -26,9 +26,10 @@ func TestNormalizeVisibleMethods(t *testing.T) { " wxpay_direct ", "wxpay", "stripe", + "ldc", }) - want := []string{"alipay", "wxpay", "stripe"} + want := []string{"alipay", "wxpay", "stripe", "ldc"} if len(got) != len(want) { t.Fatalf("NormalizeVisibleMethods len = %d, want %d (%v)", len(got), len(want), got) } @@ -39,6 +40,21 @@ func TestNormalizeVisibleMethods(t *testing.T) { } } +func TestEnabledVisibleMethodsForEasyPayIncludesCustomSupportedTypes(t *testing.T) { + t.Parallel() + + got := enabledVisibleMethodsForProvider(payment.TypeEasyPay, "alipay,ldc,usdt_trc20") + want := []string{"alipay", "ldc", "usdt_trc20"} + if len(got) != len(want) { + t.Fatalf("enabledVisibleMethodsForProvider len = %d, want %d (%v)", len(got), len(want), got) + } + for i := range want { + if got[i] != want[i] { + t.Fatalf("enabledVisibleMethodsForProvider[%d] = %q, want %q (full=%v)", i, got[i], want[i], got) + } + } +} + func TestNormalizePaymentSource(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/payment_visible_method_instances.go b/backend/internal/service/payment_visible_method_instances.go index 899bd7a020..97b3b1ef66 100644 --- a/backend/internal/service/payment_visible_method_instances.go +++ b/backend/internal/service/payment_visible_method_instances.go @@ -16,8 +16,7 @@ func enabledVisibleMethodsForProvider(providerKey, supportedTypes string) []stri methodSet := make(map[string]struct{}, 2) addMethod := func(method string) { method = NormalizeVisibleMethod(method) - switch method { - case payment.TypeAlipay, payment.TypeWxpay: + if method != "" { methodSet[method] = struct{}{} } } @@ -55,6 +54,14 @@ func enabledVisibleMethodsForProvider(providerKey, supportedTypes string) []stri for _, method := range []string{payment.TypeAlipay, payment.TypeWxpay} { if _, ok := methodSet[method]; ok { methods = append(methods, method) + delete(methodSet, method) + } + } + for _, supportedType := range splitTypes(supportedTypes) { + method := NormalizeVisibleMethod(supportedType) + if _, ok := methodSet[method]; ok { + methods = append(methods, method) + delete(methodSet, method) } } return methods @@ -215,7 +222,7 @@ func (s *PaymentConfigService) resolveEnabledVisibleMethodInstance( } method = NormalizeVisibleMethod(method) - if method != payment.TypeAlipay && method != payment.TypeWxpay { + if method == "" { return nil, nil } diff --git a/backend/internal/service/pricing_service.go b/backend/internal/service/pricing_service.go index cc62248d08..f3fb36d73a 100644 --- a/backend/internal/service/pricing_service.go +++ b/backend/internal/service/pricing_service.go @@ -831,6 +831,13 @@ func (s *PricingService) matchOpenAIModel(model string) *LiteLLMModelPricing { } } + // GPT-5.6(sol / terra / luna)回退到 GPT-5.4 定价 + if strings.HasPrefix(model, "gpt-5.6") { + logger.With(zap.String("component", "service.pricing")). + Info(fmt.Sprintf("[Pricing] OpenAI fallback matched %s -> %s", model, "gpt-5.4(static)")) + return openAIGPT54FallbackPricing + } + // GPT-5.5 回退到 GPT-5.4 定价 if strings.HasPrefix(model, "gpt-5.5") { logger.With(zap.String("component", "service.pricing")). diff --git a/backend/internal/service/ratelimit_session_window_test.go b/backend/internal/service/ratelimit_session_window_test.go index cb19227e54..279a31ccdc 100644 --- a/backend/internal/service/ratelimit_session_window_test.go +++ b/backend/internal/service/ratelimit_session_window_test.go @@ -87,6 +87,9 @@ func (m *sessionWindowMockRepo) List(context.Context, pagination.PaginationParam func (m *sessionWindowMockRepo) ListWithFilters(context.Context, pagination.PaginationParams, string, string, string, string, int64, string) ([]Account, *pagination.PaginationResult, error) { panic("unexpected") } +func (m *sessionWindowMockRepo) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]Account, error) { + panic("unexpected") +} func (m *sessionWindowMockRepo) ListByGroup(context.Context, int64) ([]Account, error) { panic("unexpected") } diff --git a/backend/internal/service/setting_service.go b/backend/internal/service/setting_service.go index 1024243bea..3aeb611418 100644 --- a/backend/internal/service/setting_service.go +++ b/backend/internal/service/setting_service.go @@ -1921,7 +1921,7 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting if err != nil { return nil, err } - if err := normalizeOpenAIAdvancedSchedulerOverrides(settings); err != nil { + if err := s.normalizeOpenAIAdvancedSchedulerOverrides(settings); err != nil { return nil, err } settings.PaymentVisibleMethodAlipaySource = alipaySource @@ -3940,7 +3940,7 @@ func formatOpenAIAdvancedSchedulerFloat(value float64) string { return strconv.FormatFloat(value, 'f', -1, 64) } -func normalizeOpenAIAdvancedSchedulerOverrides(settings *SystemSettings) error { +func (s *SettingService) normalizeOpenAIAdvancedSchedulerOverrides(settings *SystemSettings) error { lbTopK, err := normalizeOptionalPositiveIntString(settings.OpenAIAdvancedSchedulerLBTopK) if err != nil { return infraerrors.BadRequest("INVALID_OPENAI_ADVANCED_SCHEDULER_LB_TOP_K", "openai advanced scheduler TopK must be a positive integer or empty") @@ -3965,9 +3965,35 @@ func normalizeOpenAIAdvancedSchedulerOverrides(settings *SystemSettings) error { } *target = normalized } + + // 与 config.Validate 的 "scheduler_score_weights must not all be zero" 保持一致: + // 覆盖值(空则回退到生效的配置值)叠加后的基础权重和不允许为 0, + // 否则调度会静默退化为 TopK 内均匀随机。 + effective := s.openAIAdvancedSchedulerEffectiveWeights() + baseSum := resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightPriority, effective.Priority) + + resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightLoad, effective.Load) + + resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightQueue, effective.Queue) + + resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightErrorRate, effective.ErrorRate) + + resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightTTFT, effective.TTFT) + + resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightQuotaHeadroom, effective.QuotaHeadroom) + if baseSum <= 0 { + return infraerrors.BadRequest("INVALID_OPENAI_ADVANCED_SCHEDULER_WEIGHT", "openai advanced scheduler base weights must not all be zero") + } return nil } +// resolveOpenAIAdvancedSchedulerWeight 返回覆盖值(已归一化的非空字符串),空则回退默认值。 +func resolveOpenAIAdvancedSchedulerWeight(normalized string, fallback float64) float64 { + if normalized == "" { + return fallback + } + value, err := strconv.ParseFloat(normalized, 64) + if err != nil { + return fallback + } + return value +} + func normalizeOptionalPositiveIntString(raw string) (string, error) { raw = strings.TrimSpace(raw) if raw == "" { diff --git a/backend/internal/service/upstream_models.go b/backend/internal/service/upstream_models.go index ee3e6bfc04..e9fa7de451 100644 --- a/backend/internal/service/upstream_models.go +++ b/backend/internal/service/upstream_models.go @@ -208,6 +208,8 @@ func (s *AccountTestService) buildAnthropicUpstreamModelsRequest(ctx context.Con } else { setAnthropicAPIKeyAuthHeader(req.Header, account, apiKeyAuthToken) } + // 账号级请求头覆写:模型列表探测与真实转发保持一致的最终头 + account.ApplyHeaderOverrides(req.Header) return req, nil } @@ -277,6 +279,8 @@ func (s *AccountTestService) buildOpenAIUpstreamModelsRequest(ctx context.Contex } req.Header.Set("Accept", "application/json") req.Header.Set("Authorization", "Bearer "+apiKey) + // 账号级请求头覆写:模型列表探测与真实转发保持一致的最终头 + account.ApplyHeaderOverrides(req.Header) return req, nil } diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index 6ecc40aafd..7258ff05a3 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -547,9 +547,11 @@ func ProvideAPIKeyService( cache APIKeyCache, cfg *config.Config, billingCacheService *BillingCacheService, + concurrencyService *ConcurrencyService, ) *APIKeyService { svc := NewAPIKeyService(apiKeyRepo, userRepo, groupRepo, userSubRepo, userGroupRateRepo, cache, cfg) svc.SetRateLimitCacheInvalidator(billingCacheService) + svc.SetConcurrencyService(concurrencyService) return svc } diff --git a/backend/resources/model-pricing/model_prices_and_context_window.json b/backend/resources/model-pricing/model_prices_and_context_window.json index f35a91220e..b9d18f8bf5 100644 --- a/backend/resources/model-pricing/model_prices_and_context_window.json +++ b/backend/resources/model-pricing/model_prices_and_context_window.json @@ -4960,6 +4960,150 @@ "supports_web_search": true, "supports_xhigh_reasoning_effort": true }, + "gpt-5.6-sol": { + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_272k_tokens": 1e-06, + "cache_read_input_token_cost_flex": 2.5e-07, + "cache_read_input_token_cost_priority": 1e-06, + "input_cost_per_token": 5e-06, + "input_cost_per_token_above_272k_tokens": 1e-05, + "input_cost_per_token_batches": 2.5e-06, + "input_cost_per_token_flex": 2.5e-06, + "input_cost_per_token_priority": 1e-05, + "litellm_provider": "openai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 3e-05, + "output_cost_per_token_above_272k_tokens": 4.5e-05, + "output_cost_per_token_batches": 1.5e-05, + "output_cost_per_token_flex": 1.5e-05, + "output_cost_per_token_priority": 6e-05, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "gpt-5.6-terra": { + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_272k_tokens": 1e-06, + "cache_read_input_token_cost_flex": 2.5e-07, + "cache_read_input_token_cost_priority": 1e-06, + "input_cost_per_token": 5e-06, + "input_cost_per_token_above_272k_tokens": 1e-05, + "input_cost_per_token_batches": 2.5e-06, + "input_cost_per_token_flex": 2.5e-06, + "input_cost_per_token_priority": 1e-05, + "litellm_provider": "openai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 3e-05, + "output_cost_per_token_above_272k_tokens": 4.5e-05, + "output_cost_per_token_batches": 1.5e-05, + "output_cost_per_token_flex": 1.5e-05, + "output_cost_per_token_priority": 6e-05, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "gpt-5.6-luna": { + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_272k_tokens": 1e-06, + "cache_read_input_token_cost_flex": 2.5e-07, + "cache_read_input_token_cost_priority": 1e-06, + "input_cost_per_token": 5e-06, + "input_cost_per_token_above_272k_tokens": 1e-05, + "input_cost_per_token_batches": 2.5e-06, + "input_cost_per_token_flex": 2.5e-06, + "input_cost_per_token_priority": 1e-05, + "litellm_provider": "openai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 3e-05, + "output_cost_per_token_above_272k_tokens": 4.5e-05, + "output_cost_per_token_batches": 1.5e-05, + "output_cost_per_token_flex": 1.5e-05, + "output_cost_per_token_priority": 6e-05, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, "gpt-5.5": { "cache_read_input_token_cost": 5e-07, "cache_read_input_token_cost_above_272k_tokens": 1e-06, diff --git a/frontend/src/api/admin/payment.ts b/frontend/src/api/admin/payment.ts index 49efcc355d..9bab627218 100644 --- a/frontend/src/api/admin/payment.ts +++ b/frontend/src/api/admin/payment.ts @@ -24,6 +24,7 @@ export interface AdminPaymentConfig { enabled_payment_types: string[] balance_disabled: boolean balance_recharge_multiplier: number + subscription_usd_to_cny_rate: number load_balance_strategy: string product_name_prefix: string product_name_suffix: string @@ -42,6 +43,7 @@ export interface UpdatePaymentConfigRequest { enabled_payment_types?: string[] balance_disabled?: boolean balance_recharge_multiplier?: number + subscription_usd_to_cny_rate?: number load_balance_strategy?: string product_name_prefix?: string product_name_suffix?: string diff --git a/frontend/src/api/admin/settings.ts b/frontend/src/api/admin/settings.ts index 457775870d..f5da990930 100644 --- a/frontend/src/api/admin/settings.ts +++ b/frontend/src/api/admin/settings.ts @@ -589,6 +589,7 @@ export interface SystemSettings { payment_enabled_types: string[]; payment_balance_disabled: boolean; payment_balance_recharge_multiplier: number; + payment_subscription_usd_to_cny_rate: number; payment_recharge_fee_rate: number; payment_load_balance_strategy: string; payment_product_name_prefix: string; @@ -860,6 +861,7 @@ export interface UpdateSettingsRequest { payment_enabled_types?: string[]; payment_balance_disabled?: boolean; payment_balance_recharge_multiplier?: number; + payment_subscription_usd_to_cny_rate?: number; payment_recharge_fee_rate?: number; payment_load_balance_strategy?: string; payment_product_name_prefix?: string; diff --git a/frontend/src/assets/icons/payment.svg b/frontend/src/assets/icons/payment.svg new file mode 100644 index 0000000000..c78bea4cf7 --- /dev/null +++ b/frontend/src/assets/icons/payment.svg @@ -0,0 +1,9 @@ + diff --git a/frontend/src/components/account/BulkEditAccountModal.vue b/frontend/src/components/account/BulkEditAccountModal.vue index d24488b47a..91656b196d 100644 --- a/frontend/src/components/account/BulkEditAccountModal.vue +++ b/frontend/src/components/account/BulkEditAccountModal.vue @@ -486,6 +486,129 @@ + +
+ {{ t('admin.accounts.headerOverride.hint') }} +
+
+
+ {{ t('admin.accounts.headerOverride.bulkReplaceHint') }} +
+ ++ {{ t('admin.accounts.headerOverride.emptyValueHint') }} +
++ {{ t('admin.accounts.headerOverride.bulkDisableHint') }} +
++ {{ t('admin.accounts.headerOverride.hint') }} +
+
+
+ {{ t('admin.accounts.headerOverride.emptyValueHint') }} +
++ {{ t('admin.accounts.headerOverride.hint') }} +
+
+
+ {{ t('admin.accounts.headerOverride.emptyValueHint') }} +
++ {{ t('admin.settings.payment.easypayCustomMethodsHint') }} +
++ {{ + t("admin.settings.payment.subscriptionUsdToCnyRateHint") + }} +
+