mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
merge: feature/channel → release (v0.1.105.15) — 计费迁移到 CalculateCostUnified
This commit is contained in:
@@ -1 +1 @@
|
||||
0.1.105.14
|
||||
0.1.105.15
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 依赖。
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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"}}
|
||||
|
||||
@@ -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 区间计费
|
||||
|
||||
@@ -42,6 +42,7 @@ func newGatewayRecordUsageServiceForTest(usageRepo UsageLogRepository, userRepo
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -145,6 +145,7 @@ func newOpenAIRecordUsageServiceForTest(usageRepo UsageLogRepository, userRepo U
|
||||
nil,
|
||||
&DeferredService{},
|
||||
nil,
|
||||
nil,
|
||||
)
|
||||
svc.userGroupRateResolver = newUserGroupRateResolver(
|
||||
rateRepo,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -615,6 +615,7 @@ func TestNewOpenAIGatewayService_InitializesOpenAIWSResolver(t *testing.T) {
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
)
|
||||
|
||||
decision := svc.getOpenAIWSProtocolResolver().Resolve(nil)
|
||||
|
||||
Reference in New Issue
Block a user