diff --git a/backend/cmd/server/VERSION b/backend/cmd/server/VERSION index 70a88088c2..64405cb6fb 100644 --- a/backend/cmd/server/VERSION +++ b/backend/cmd/server/VERSION @@ -1 +1 @@ -0.1.105.14 +0.1.105.15 diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index e8cb928cda..2b93e2a84c 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -178,10 +178,9 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { digestSessionStore := service.NewDigestSessionStore() channelService := service.NewChannelService(channelRepository, apiKeyAuthCacheInvalidator) modelPricingResolver := service.NewModelPricingResolver(channelService, billingService) - _ = modelPricingResolver // Phase 4: 已注册,后续 Gateway 迁移时使用 - gatewayService := service.NewGatewayService(accountRepository, groupRepository, usageLogRepository, usageBillingRepository, userRepository, userSubscriptionRepository, userGroupRateRepository, gatewayCache, configConfig, schedulerSnapshotService, concurrencyService, billingService, rateLimitService, billingCacheService, identityService, httpUpstream, deferredService, claudeTokenProvider, sessionLimitCache, rpmCache, digestSessionStore, settingService, tlsFingerprintProfileService, channelService) + gatewayService := service.NewGatewayService(accountRepository, groupRepository, usageLogRepository, usageBillingRepository, userRepository, userSubscriptionRepository, userGroupRateRepository, gatewayCache, configConfig, schedulerSnapshotService, concurrencyService, billingService, rateLimitService, billingCacheService, identityService, httpUpstream, deferredService, claudeTokenProvider, sessionLimitCache, rpmCache, digestSessionStore, settingService, tlsFingerprintProfileService, channelService, modelPricingResolver) openAITokenProvider := service.ProvideOpenAITokenProvider(accountRepository, geminiTokenCache, openAIOAuthService, oauthRefreshAPI) - openAIGatewayService := service.NewOpenAIGatewayService(accountRepository, usageLogRepository, usageBillingRepository, userRepository, userSubscriptionRepository, userGroupRateRepository, gatewayCache, configConfig, schedulerSnapshotService, concurrencyService, billingService, rateLimitService, billingCacheService, httpUpstream, deferredService, openAITokenProvider) + openAIGatewayService := service.NewOpenAIGatewayService(accountRepository, usageLogRepository, usageBillingRepository, userRepository, userSubscriptionRepository, userGroupRateRepository, gatewayCache, configConfig, schedulerSnapshotService, concurrencyService, billingService, rateLimitService, billingCacheService, httpUpstream, deferredService, openAITokenProvider, modelPricingResolver) geminiMessagesCompatService := service.NewGeminiMessagesCompatService(accountRepository, groupRepository, gatewayCache, schedulerSnapshotService, geminiTokenProvider, rateLimitService, httpUpstream, antigravityGatewayService, configConfig) opsSystemLogSink := service.ProvideOpsSystemLogSink(opsRepository) opsService := service.NewOpsService(opsRepository, settingRepository, configConfig, accountRepository, userRepository, concurrencyService, gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, opsSystemLogSink) diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index 2a91601ecb..0618905300 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -164,14 +164,10 @@ func (h *GatewayHandler) Messages(c *gin.Context) { channelMapping = h.gatewayService.ResolveChannelMapping(c.Request.Context(), *apiKey.GroupID, reqModel) } - // 渠道模型限制检查 + // 渠道模型限制检查:使用原始请求模型名,因为定价列表中注册的是用户请求的模型名 if apiKey.GroupID != nil { - checkModel := reqModel - if channelMapping.Mapped { - checkModel = channelMapping.MappedModel - } - if h.gatewayService.IsModelRestricted(c.Request.Context(), *apiKey.GroupID, checkModel) { - h.errorResponse(c, http.StatusForbidden, "invalid_request_error", "Model not available in current channel: "+reqModel) + if h.gatewayService.IsModelRestricted(c.Request.Context(), *apiKey.GroupID, reqModel) { + h.errorResponse(c, http.StatusServiceUnavailable, "api_error", "No available accounts") return } } diff --git a/backend/internal/handler/gateway_handler_warmup_intercept_unit_test.go b/backend/internal/handler/gateway_handler_warmup_intercept_unit_test.go index 7dc062df0a..4caef9551b 100644 --- a/backend/internal/handler/gateway_handler_warmup_intercept_unit_test.go +++ b/backend/internal/handler/gateway_handler_warmup_intercept_unit_test.go @@ -162,6 +162,7 @@ func newTestGatewayHandler(t *testing.T, group *service.Group, accounts []*servi nil, // settingService nil, // tlsFPProfileService nil, // channelService + nil, // resolver ) // RunModeSimple:跳过计费检查,避免引入 repo/cache 依赖。 diff --git a/backend/internal/handler/sora_client_handler_test.go b/backend/internal/handler/sora_client_handler_test.go index 5bacfc06d2..5705578660 100644 --- a/backend/internal/handler/sora_client_handler_test.go +++ b/backend/internal/handler/sora_client_handler_test.go @@ -124,9 +124,6 @@ func (r *stubSoraGenRepo) CountByUserAndStatus(_ context.Context, _ int64, _ []s } return r.countValue, nil } -func (r *stubSoraGenRepo) CountByStorageType(_ context.Context, _ string, _ []string) (int64, error) { - return 0, nil -} // ==================== 辅助函数 ==================== @@ -1661,7 +1658,7 @@ func TestStoreMediaWithDegradation_S3SuccessSingleURL(t *testing.T) { defer fakeS3.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{objectStorage: s3Storage} + h := &SoraClientHandler{s3Storage: s3Storage} storedURL, storedURLs, storageType, s3Keys, fileSize := h.storeMediaWithDegradation( context.Background(), 1, "video", sourceServer.URL+"/v.mp4", nil, @@ -1683,7 +1680,7 @@ func TestStoreMediaWithDegradation_S3SuccessMultiURL(t *testing.T) { defer fakeS3.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{objectStorage: s3Storage} + h := &SoraClientHandler{s3Storage: s3Storage} urls := []string{sourceServer.URL + "/a.mp4", sourceServer.URL + "/b.mp4"} storedURL, storedURLs, storageType, s3Keys, fileSize := h.storeMediaWithDegradation( @@ -1708,7 +1705,7 @@ func TestStoreMediaWithDegradation_S3DownloadFails(t *testing.T) { defer badSource.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{objectStorage: s3Storage} + h := &SoraClientHandler{s3Storage: s3Storage} _, _, storageType, _, _ := h.storeMediaWithDegradation( context.Background(), 1, "video", badSource.URL+"/missing.mp4", nil, @@ -1723,7 +1720,7 @@ func TestStoreMediaWithDegradation_S3FailsSingleURL(t *testing.T) { defer fakeS3.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{objectStorage: s3Storage} + h := &SoraClientHandler{s3Storage: s3Storage} _, _, storageType, s3Keys, _ := h.storeMediaWithDegradation( context.Background(), 1, "video", sourceServer.URL+"/v.mp4", nil, @@ -1740,7 +1737,7 @@ func TestStoreMediaWithDegradation_S3PartialFailureCleanup(t *testing.T) { defer fakeS3.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{objectStorage: s3Storage} + h := &SoraClientHandler{s3Storage: s3Storage} urls := []string{sourceServer.URL + "/a.mp4", sourceServer.URL + "/b.mp4"} _, _, storageType, s3Keys, _ := h.storeMediaWithDegradation( @@ -1824,8 +1821,8 @@ func TestStoreMediaWithDegradation_S3FailsFallbackToLocal(t *testing.T) { } mediaStorage := service.NewSoraMediaStorage(cfg) h := &SoraClientHandler{ - objectStorage: s3Storage, - mediaStorage: mediaStorage, + s3Storage: s3Storage, + mediaStorage: mediaStorage, } _, _, storageType, _, _ := h.storeMediaWithDegradation( @@ -1851,14 +1848,14 @@ func TestSaveToStorage_S3EnabledButUploadFails(t *testing.T) { } s3Storage := newS3StorageForHandler(fakeS3.URL) genService := service.NewSoraGenerationService(repo, nil, nil) - h := &SoraClientHandler{genService: genService, objectStorage: s3Storage} + h := &SoraClientHandler{genService: genService, s3Storage: s3Storage} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} h.SaveToStorage(c) require.Equal(t, http.StatusInternalServerError, rec.Code) resp := parseResponse(t, rec) - require.Contains(t, resp["message"], "上传到存储失败") + require.Contains(t, resp["message"], "S3") } func TestSaveToStorage_UpstreamURLExpired(t *testing.T) { @@ -1877,7 +1874,7 @@ func TestSaveToStorage_UpstreamURLExpired(t *testing.T) { } s3Storage := newS3StorageForHandler(fakeS3.URL) genService := service.NewSoraGenerationService(repo, nil, nil) - h := &SoraClientHandler{genService: genService, objectStorage: s3Storage} + h := &SoraClientHandler{genService: genService, s3Storage: s3Storage} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -1901,7 +1898,7 @@ func TestSaveToStorage_S3EnabledUploadSuccess(t *testing.T) { } s3Storage := newS3StorageForHandler(fakeS3.URL) genService := service.NewSoraGenerationService(repo, nil, nil) - h := &SoraClientHandler{genService: genService, objectStorage: s3Storage} + h := &SoraClientHandler{genService: genService, s3Storage: s3Storage} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -1909,7 +1906,7 @@ func TestSaveToStorage_S3EnabledUploadSuccess(t *testing.T) { require.Equal(t, http.StatusOK, rec.Code) resp := parseResponse(t, rec) data := resp["data"].(map[string]any) - require.Contains(t, data["message"], "已保存到云存储") + require.Contains(t, data["message"], "S3") require.NotEmpty(t, data["object_key"]) // 验证记录已更新为 S3 存储 require.Equal(t, service.SoraStorageTypeS3, repo.gens[1].StorageType) @@ -1933,7 +1930,7 @@ func TestSaveToStorage_S3EnabledUploadSuccess_MultiMediaURLs(t *testing.T) { } s3Storage := newS3StorageForHandler(fakeS3.URL) genService := service.NewSoraGenerationService(repo, nil, nil) - h := &SoraClientHandler{genService: genService, objectStorage: s3Storage} + h := &SoraClientHandler{genService: genService, s3Storage: s3Storage} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -1969,7 +1966,7 @@ func TestSaveToStorage_S3EnabledUploadSuccessWithQuota(t *testing.T) { SoraStorageUsedBytes: 0, } quotaService := service.NewSoraQuotaService(userRepo, nil, nil) - h := &SoraClientHandler{genService: genService, objectStorage: s3Storage, quotaService: quotaService} + h := &SoraClientHandler{genService: genService, s3Storage: s3Storage, quotaService: quotaService} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -1995,7 +1992,7 @@ func TestSaveToStorage_S3UploadSuccessMarkCompletedFails(t *testing.T) { repo.updateErr = fmt.Errorf("db error") s3Storage := newS3StorageForHandler(fakeS3.URL) genService := service.NewSoraGenerationService(repo, nil, nil) - h := &SoraClientHandler{genService: genService, objectStorage: s3Storage} + h := &SoraClientHandler{genService: genService, s3Storage: s3Storage} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -2011,7 +2008,7 @@ func TestGetStorageStatus_S3EnabledNotHealthy(t *testing.T) { defer fakeS3.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{objectStorage: s3Storage} + h := &SoraClientHandler{s3Storage: s3Storage} c, rec := makeGinContext("GET", "/api/v1/sora/storage-status", "", 0) h.GetStorageStatus(c) @@ -2027,7 +2024,7 @@ func TestGetStorageStatus_S3EnabledHealthy(t *testing.T) { defer fakeS3.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{objectStorage: s3Storage} + h := &SoraClientHandler{s3Storage: s3Storage} c, rec := makeGinContext("GET", "/api/v1/sora/storage-status", "", 0) h.GetStorageStatus(c) @@ -2227,7 +2224,8 @@ func (s *stubSoraClientForHandler) GetVideoTask(_ context.Context, _ *service.Ac func newMinimalGatewayService(accountRepo service.AccountRepository) *service.GatewayService { return service.NewGatewayService( accountRepo, nil, nil, nil, nil, nil, nil, nil, nil, - nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, + nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, + nil, nil, ) } @@ -2467,7 +2465,7 @@ func TestProcessGeneration_FullSuccessWithS3(t *testing.T) { genService: genService, gatewayService: gatewayService, soraGatewayService: soraGatewayService, - objectStorage: s3Storage, + s3Storage: s3Storage, quotaService: quotaService, } @@ -2517,7 +2515,7 @@ func TestProcessGeneration_MarkCompletedFails(t *testing.T) { // ==================== cleanupStoredMedia 直接测试 ==================== func TestCleanupStoredMedia_S3Path(t *testing.T) { - // S3 清理路径:objectStorage 为 nil 时不 panic + // S3 清理路径:s3Storage 为 nil 时不 panic h := &SoraClientHandler{} // 不应 panic h.cleanupStoredMedia(context.Background(), service.SoraStorageTypeS3, []string{"key1"}, nil) @@ -2975,7 +2973,7 @@ func TestSaveToStorage_QuotaExceeded(t *testing.T) { SoraStorageUsedBytes: 10, } quotaService := service.NewSoraQuotaService(userRepo, nil, nil) - h := &SoraClientHandler{genService: genService, objectStorage: s3Storage, quotaService: quotaService} + h := &SoraClientHandler{genService: genService, s3Storage: s3Storage, quotaService: quotaService} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -3003,7 +3001,7 @@ func TestSaveToStorage_QuotaNonQuotaError(t *testing.T) { // 用户不存在 → GetByID 失败 → AddUsage 返回普通 error userRepo := newStubUserRepoForHandler() quotaService := service.NewSoraQuotaService(userRepo, nil, nil) - h := &SoraClientHandler{genService: genService, objectStorage: s3Storage, quotaService: quotaService} + h := &SoraClientHandler{genService: genService, s3Storage: s3Storage, quotaService: quotaService} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -3026,7 +3024,7 @@ func TestSaveToStorage_EmptyMediaURLs(t *testing.T) { } s3Storage := newS3StorageForHandler(fakeS3.URL) genService := service.NewSoraGenerationService(repo, nil, nil) - h := &SoraClientHandler{genService: genService, objectStorage: s3Storage} + h := &SoraClientHandler{genService: genService, s3Storage: s3Storage} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -3053,7 +3051,7 @@ func TestSaveToStorage_MultiURL_SecondUploadFails(t *testing.T) { } s3Storage := newS3StorageForHandler(fakeS3.URL) genService := service.NewSoraGenerationService(repo, nil, nil) - h := &SoraClientHandler{genService: genService, objectStorage: s3Storage} + h := &SoraClientHandler{genService: genService, s3Storage: s3Storage} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -3086,7 +3084,7 @@ func TestSaveToStorage_MarkCompletedFailsWithQuotaRollback(t *testing.T) { SoraStorageUsedBytes: 0, } quotaService := service.NewSoraQuotaService(userRepo, nil, nil) - h := &SoraClientHandler{genService: genService, objectStorage: s3Storage, quotaService: quotaService} + h := &SoraClientHandler{genService: genService, s3Storage: s3Storage, quotaService: quotaService} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -3100,7 +3098,7 @@ func TestCleanupStoredMedia_WithS3Storage_ActualDelete(t *testing.T) { fakeS3 := newFakeS3Server("ok") defer fakeS3.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{objectStorage: s3Storage} + h := &SoraClientHandler{s3Storage: s3Storage} h.cleanupStoredMedia(context.Background(), service.SoraStorageTypeS3, []string{"key1", "key2"}, nil) } @@ -3109,7 +3107,7 @@ func TestCleanupStoredMedia_S3DeleteFails_LogOnly(t *testing.T) { fakeS3 := newFakeS3Server("fail") defer fakeS3.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{objectStorage: s3Storage} + h := &SoraClientHandler{s3Storage: s3Storage} h.cleanupStoredMedia(context.Background(), service.SoraStorageTypeS3, []string{"key1"}, nil) } diff --git a/backend/internal/handler/sora_gateway_handler_test.go b/backend/internal/handler/sora_gateway_handler_test.go index 18e6e92971..e053b668d3 100644 --- a/backend/internal/handler/sora_gateway_handler_test.go +++ b/backend/internal/handler/sora_gateway_handler_test.go @@ -466,6 +466,7 @@ func TestSoraGatewayHandler_ChatCompletions(t *testing.T) { nil, // settingService nil, // tlsFPProfileService nil, // channelService + nil, // resolver ) soraClient := &stubSoraClient{imageURLs: []string{"https://example.com/a.png"}} diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index 7deb1cf991..d256102cc1 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -104,6 +104,7 @@ type CostBreakdown struct { CacheReadCost float64 TotalCost float64 ActualCost float64 // 应用倍率后的实际费用 + BillingMode string // 计费模式("token"/"per_request"/"image"),由 CalculateCostUnified 填充 } // BillingService 计费服务 @@ -439,12 +440,21 @@ func (s *BillingService) CalculateCostUnified(input CostInput) (*CostBreakdown, input.RateMultiplier = 1.0 } + var breakdown *CostBreakdown + var err error switch resolved.Mode { case BillingModePerRequest, BillingModeImage: - return s.calculatePerRequestCost(resolved, input) + breakdown, err = s.calculatePerRequestCost(resolved, input) default: // BillingModeToken - return s.calculateTokenCost(resolved, input) + breakdown, err = s.calculateTokenCost(resolved, input) } + if err == nil && breakdown != nil { + breakdown.BillingMode = string(resolved.Mode) + if breakdown.BillingMode == "" { + breakdown.BillingMode = string(BillingModeToken) + } + } + return breakdown, err } // calculateTokenCost 按 token 区间计费 diff --git a/backend/internal/service/gateway_record_usage_test.go b/backend/internal/service/gateway_record_usage_test.go index 5df0b58c57..97703a9d5e 100644 --- a/backend/internal/service/gateway_record_usage_test.go +++ b/backend/internal/service/gateway_record_usage_test.go @@ -42,6 +42,7 @@ func newGatewayRecordUsageServiceForTest(usageRepo UsageLogRepository, userRepo nil, nil, nil, + nil, ) } diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 5cc943e793..c9ce1f8b9b 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -631,6 +631,7 @@ type GatewayService struct { debugModelRouting atomic.Bool debugClaudeMimic atomic.Bool channelService *ChannelService + resolver *ModelPricingResolver debugGatewayBodyFile atomic.Pointer[os.File] // non-nil when SUB2API_DEBUG_GATEWAY_BODY is set tlsFPProfileService *TLSFingerprintProfileService } @@ -661,6 +662,7 @@ func NewGatewayService( settingService *SettingService, tlsFPProfileService *TLSFingerprintProfileService, channelService *ChannelService, + resolver *ModelPricingResolver, ) *GatewayService { userGroupRateTTL := resolveUserGroupRateCacheTTL(cfg) modelsListTTL := resolveModelsListCacheTTL(cfg) @@ -694,6 +696,7 @@ func NewGatewayService( responseHeaderFilter: compileResponseHeaderFilter(cfg), tlsFPProfileService: tlsFPProfileService, channelService: channelService, + resolver: resolver, } svc.userGroupRateResolver = newUserGroupRateResolver( userGroupRateRepo, @@ -7959,13 +7962,21 @@ func (s *GatewayService) RecordUsage(ctx context.Context, input *RecordUsageInpu CacheCreation1hTokens: result.Usage.CacheCreation1hTokens, } var err error - // 渠道定价覆盖 - var chPricing *ChannelModelPricing - if s.channelService != nil && apiKey.Group != nil { - chPricing = s.channelService.GetChannelModelPricing(ctx, apiKey.Group.ID, billingModel) - } - if chPricing != nil { - cost, err = s.billingService.CalculateCostWithChannel(billingModel, tokens, multiplier, chPricing) + if s.resolver != nil && apiKey.Group != nil { + var groupID *int64 + if apiKey.Group != nil { + gid := apiKey.Group.ID + groupID = &gid + } + cost, err = s.billingService.CalculateCostUnified(CostInput{ + Ctx: ctx, + Model: billingModel, + GroupID: groupID, + Tokens: tokens, + RequestCount: 1, + RateMultiplier: multiplier, + Resolver: s.resolver, + }) } else { cost, err = s.billingService.CalculateCost(billingModel, tokens, multiplier) } @@ -8037,6 +8048,9 @@ func (s *GatewayService) RecordUsage(ctx context.Context, input *RecordUsageInpu if result.ImageCount > 0 { billingMode := "image" usageLog.BillingMode = &billingMode + } else if cost != nil && cost.BillingMode != "" { + billingMode := cost.BillingMode + usageLog.BillingMode = &billingMode } else { billingMode := "token" usageLog.BillingMode = &billingMode @@ -8185,14 +8199,30 @@ func (s *GatewayService) RecordUsageWithLongContext(ctx context.Context, input * CacheCreation1hTokens: result.Usage.CacheCreation1hTokens, } var err error - // 渠道定价覆盖 - var chPricing2 *ChannelModelPricing - if s.channelService != nil && apiKey.Group != nil { - chPricing2 = s.channelService.GetChannelModelPricing(ctx, apiKey.Group.ID, billingModel) + // 优先尝试 Resolver + CalculateCostUnified(仅在有渠道定价时使用) + useUnified := false + if s.resolver != nil && apiKey.Group != nil { + gid := apiKey.Group.ID + resolved := s.resolver.Resolve(ctx, PricingInput{ + Model: billingModel, + GroupID: &gid, + }) + if resolved.Source == "channel" { + // 有渠道定价,渠道区间已包含上下文分层 + cost, err = s.billingService.CalculateCostUnified(CostInput{ + Ctx: ctx, + Model: billingModel, + GroupID: &gid, + Tokens: tokens, + RequestCount: 1, + RateMultiplier: multiplier, + Resolver: s.resolver, + }) + useUnified = true + } } - if chPricing2 != nil { - cost, err = s.billingService.CalculateCostWithChannel(billingModel, tokens, multiplier, chPricing2) - } else { + if !useUnified { + // 无渠道定价,保持原有长上下文双倍计费逻辑(如 Gemini 200K 阈值) cost, err = s.billingService.CalculateCostWithLongContext(billingModel, tokens, multiplier, input.LongContextThreshold, input.LongContextMultiplier) } if err != nil { @@ -8257,6 +8287,9 @@ func (s *GatewayService) RecordUsageWithLongContext(ctx context.Context, input * if result.ImageCount > 0 { billingMode := "image" usageLog.BillingMode = &billingMode + } else if cost != nil && cost.BillingMode != "" { + billingMode := cost.BillingMode + usageLog.BillingMode = &billingMode } else { billingMode := "token" usageLog.BillingMode = &billingMode diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go index 5aa4db8a5a..4c5122f6fa 100644 --- a/backend/internal/service/openai_gateway_record_usage_test.go +++ b/backend/internal/service/openai_gateway_record_usage_test.go @@ -145,6 +145,7 @@ func newOpenAIRecordUsageServiceForTest(usageRepo UsageLogRepository, userRepo U nil, &DeferredService{}, nil, + nil, ) svc.userGroupRateResolver = newUserGroupRateResolver( rateRepo, diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index bdc4362373..f7639f9713 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -322,6 +322,7 @@ type OpenAIGatewayService struct { openAITokenProvider *OpenAITokenProvider toolCorrector *CodexToolCorrector openaiWSResolver OpenAIWSProtocolResolver + resolver *ModelPricingResolver openaiWSPoolOnce sync.Once openaiWSStateStoreOnce sync.Once @@ -357,6 +358,7 @@ func NewOpenAIGatewayService( httpUpstream HTTPUpstream, deferredService *DeferredService, openAITokenProvider *OpenAITokenProvider, + resolver *ModelPricingResolver, ) *OpenAIGatewayService { svc := &OpenAIGatewayService{ accountRepo: accountRepo, @@ -384,6 +386,7 @@ func NewOpenAIGatewayService( openAITokenProvider: openAITokenProvider, toolCorrector: NewCodexToolCorrector(), openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), + resolver: resolver, responseHeaderFilter: compileResponseHeaderFilter(cfg), codexSnapshotThrottle: newAccountWriteThrottle(openAICodexSnapshotPersistMinInterval), } @@ -4152,6 +4155,8 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec multiplier = resolver.Resolve(ctx, user.ID, *apiKey.GroupID, apiKey.Group.RateMultiplier) } + var cost *CostBreakdown + var err error billingModel := forwardResultBillingModel(result.Model, result.UpstreamModel) if result.BillingModel != "" { billingModel = strings.TrimSpace(result.BillingModel) @@ -4160,7 +4165,21 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec if result.ServiceTier != nil { serviceTier = strings.TrimSpace(*result.ServiceTier) } - cost, err := s.billingService.CalculateCostWithServiceTier(billingModel, tokens, multiplier, serviceTier) + if s.resolver != nil && apiKey.Group != nil { + gid := apiKey.Group.ID + cost, err = s.billingService.CalculateCostUnified(CostInput{ + Ctx: ctx, + Model: billingModel, + GroupID: &gid, + Tokens: tokens, + RequestCount: 1, + RateMultiplier: multiplier, + ServiceTier: serviceTier, + Resolver: s.resolver, + }) + } else { + cost, err = s.billingService.CalculateCostWithServiceTier(billingModel, tokens, multiplier, serviceTier) + } if err != nil { cost = &CostBreakdown{ActualCost: 0} } @@ -4207,8 +4226,11 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec FirstTokenMs: result.FirstTokenMs, CreatedAt: time.Now(), } - // 设置计费模式(OpenAI 网关都是 token 计费) - { + // 设置计费模式 + if cost != nil && cost.BillingMode != "" { + billingMode := cost.BillingMode + usageLog.BillingMode = &billingMode + } else { billingMode := "token" usageLog.BillingMode = &billingMode } diff --git a/backend/internal/service/openai_ws_protocol_forward_test.go b/backend/internal/service/openai_ws_protocol_forward_test.go index 8c5c936844..e8d9f8f792 100644 --- a/backend/internal/service/openai_ws_protocol_forward_test.go +++ b/backend/internal/service/openai_ws_protocol_forward_test.go @@ -615,6 +615,7 @@ func TestNewOpenAIGatewayService_InitializesOpenAIWSResolver(t *testing.T) { nil, nil, nil, + nil, ) decision := svc.getOpenAIWSProtocolResolver().Resolve(nil)