Files
sub2api/backend/internal/service/openai_gateway_usage.go
T
李建琦 6d655c9903
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
2026-08-21 18:30:13 +08:00

1089 lines
42 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package service
// 本文件由 openai_gateway_service.go 纯移动拆分而来:用量记录、计费成本计算与
// Codex 用量快照。仅做代码搬迁,无任何行为变更。
import (
"context"
"errors"
"fmt"
"net/http"
"strconv"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
"github.com/Wei-Shaw/sub2api/internal/pkg/timezone"
"go.uber.org/zap"
)
// OpenAIRecordUsageInput input for recording usage
type OpenAIRecordUsageInput struct {
Result *OpenAIForwardResult
APIKey *APIKey
User *User
Account *Account
Subscription *UserSubscription
InboundEndpoint string
UpstreamEndpoint string
UserAgent string // 请求的 User-Agent
IPAddress string // 请求的客户端 IP 地址
SessionID string // 客户端显式会话标识(session_id / X-Session-Id 等请求头),仅用于用量行会话关联
RequestPayloadHash string
APIKeyService APIKeyQuotaUpdater
QuotaPlatform string // user×platform quota platform resolved by the handler before async billing.
// PricingAt 是请求级定价时刻(请求开始捕获,与利润门的 D 同源):高峰因子
// 按该时刻计算,保证同一请求从准入到扣费不中途变价。零值回退记录时刻
//(既有行为),供未装配的路径(图片/异步/cyber 等)沿用。
PricingAt time.Time
// CyberBlocked 为 true 时把该用量行标记为 cyberrequest_type=cyber),计费逻辑不变。
CyberBlocked bool
ChannelUsageFields
}
// CyberPolicyUsageInput 是 cyber 拒绝、未走正常 RecordUsage 的请求记录用量的入参。
// 用量按上游真实 token 计费,与 WS cyber 及正常请求口径一致(InputTokens/OutputTokens
// 取自上游 response.failed 报告的 usage,即 mark.UpstreamInTok/OutTok)。
type CyberPolicyUsageInput struct {
APIKey *APIKey
Account *Account
Subscription *UserSubscription
RequestID string
Model string
Stream bool
InputTokens int
OutputTokens int
// 渠道归因与请求级 meta,使 cyber 计费行与正常 RecordUsage 行口径一致
// (否则 cyber 行 channel_id 等为空,渠道维度统计会遗漏 cyber 命中)。
InboundEndpoint string
UpstreamEndpoint string
UserAgent string
IPAddress string
SessionID string
RequestPayloadHash string
APIKeyService APIKeyQuotaUpdater
ChannelUsageFields
}
// RecordCyberPolicyUsageLog 为被上游 cyber_policy 拒绝、未走正常 RecordUsage 的请求
// HTTP forward 返回错误路径)记录用量并按上游真实 token 计费,使其与 WS cyber 路径、
// 与正常请求的计费口径统一(不再是 tokens=0 免费行)。token 取自上游 response.failed
// 报告的 usage(非流式直接拒通常为 0,cost 随之为 0)。复用 RecordUsage 完成成本计算、
// 扣费与用量行写入(request_type=cyber 由 CyberBlocked 置位)。仅 forward 返回错误的
// 路径由 handler 调用,避免与成功路径的正常 RecordUsage 重复。
func (s *OpenAIGatewayService) RecordCyberPolicyUsageLog(ctx context.Context, in CyberPolicyUsageInput) {
if s == nil || in.APIKey == nil || in.APIKey.User == nil || in.Account == nil || strings.TrimSpace(in.Model) == "" {
return
}
result := &OpenAIForwardResult{
RequestID: in.RequestID,
Model: in.Model,
Stream: in.Stream,
Usage: OpenAIUsage{
InputTokens: in.InputTokens,
OutputTokens: in.OutputTokens,
},
}
if err := s.RecordUsage(ctx, &OpenAIRecordUsageInput{
Result: result,
APIKey: in.APIKey,
User: in.APIKey.User,
Account: in.Account,
Subscription: in.Subscription,
InboundEndpoint: in.InboundEndpoint,
UpstreamEndpoint: in.UpstreamEndpoint,
UserAgent: in.UserAgent,
IPAddress: in.IPAddress,
SessionID: in.SessionID,
RequestPayloadHash: in.RequestPayloadHash,
APIKeyService: in.APIKeyService,
ChannelUsageFields: in.ChannelUsageFields,
CyberBlocked: true,
}); err != nil {
logger.LegacyPrintf("service.openai_gateway", "cyber usage record failed: request_id=%s err=%v", in.RequestID, err)
}
}
// ResolveUserGroupRateMultiplier resolves the same cached multiplier used by OpenAI usage billing.
func (s *OpenAIGatewayService) ResolveUserGroupRateMultiplier(ctx context.Context, userID, groupID int64, groupDefaultMultiplier float64) float64 {
if s == nil {
return groupDefaultMultiplier
}
resolver := s.userGroupRateResolver
if resolver == nil {
resolver = newUserGroupRateResolver(nil, nil, resolveUserGroupRateCacheTTL(s.cfg), nil, "service.openai_gateway")
}
return resolver.Resolve(ctx, userID, groupID, groupDefaultMultiplier)
}
// openAIUsagePricingAt 返回本次用量记录使用的定价时刻:优先请求级 PricingAt
// (与利润门 D 同源同刻),未装配时回退记录时刻(既有行为)。
func openAIUsagePricingAt(input *OpenAIRecordUsageInput) time.Time {
if input != nil && !input.PricingAt.IsZero() {
return input.PricingAt
}
return timezone.Now()
}
// RecordUsage records usage and deducts balance
func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRecordUsageInput) error {
if input == nil {
return errors.New("openai usage input is nil")
}
result := input.Result
if result == nil {
return errors.New("openai usage result is nil")
}
if s.rateLimitService != nil && input.Account != nil && input.Account.Platform == PlatformOpenAI {
s.rateLimitService.ResetOpenAI403Counter(ctx, input.Account.ID)
}
apiKey := input.APIKey
user := input.User
account := input.Account
subscription := input.Subscription
if !isGrokVideoUsageResult(result, nil) {
ApplyOpenAIImageBillingResolution(result)
}
// OpenAI input_tokens 是总输入,包含缓存读取和缓存写入明细。
// 将三类 token 拆成互斥桶,避免缓存写入同时按普通输入和 cache_write 重复计费。
actualInputTokens := result.Usage.InputTokens - result.Usage.CacheReadInputTokens - result.Usage.CacheCreationInputTokens
if actualInputTokens < 0 {
actualInputTokens = 0
}
// Calculate cost
tokens := UsageTokens{
InputTokens: actualInputTokens,
ImageInputTokens: result.Usage.ImageInputTokens,
OutputTokens: result.Usage.OutputTokens,
CacheCreationTokens: result.Usage.CacheCreationInputTokens,
CacheReadTokens: result.Usage.CacheReadInputTokens,
ImageOutputTokens: result.Usage.ImageOutputTokens,
}
// Get rate multiplier
multiplier := 1.0
if s.cfg != nil {
multiplier = s.cfg.Default.RateMultiplier
}
if apiKey.GroupID != nil && apiKey.Group != nil {
multiplier = s.ResolveUserGroupRateMultiplier(ctx, user.ID, *apiKey.GroupID, apiKey.Group.RateMultiplier)
}
// token 倍率叠加高峰因子(token 计费含图片 token,图片按次倍率不受影响)。
// 高峰因子按请求级 PricingAt 现算(与利润门 D 同源同刻,跨峰谷请求不中途
// 变价);未装配 PricingAt 的路径回退记录时刻,保持既有行为。不并入上面的
// Resolve,以免污染 user:group 倍率缓存。
baseMultiplier := multiplier
pricingAt := openAIUsagePricingAt(input)
multiplier, imageMultiplier := computePeakAwareMultipliers(apiKey, baseMultiplier, pricingAt)
videoMultiplier := resolveVideoRateMultiplier(apiKey, baseMultiplier)
var cost *CostBreakdown
var err error
billingModel := forwardResultBillingModel(result.Model, result.UpstreamModel)
if result.BillingModel != "" {
billingModel = strings.TrimSpace(result.BillingModel)
}
if input.BillingModelSource == BillingModelSourceChannelMapped && input.ChannelMappedModel != "" && input.ChannelMappedModel != input.OriginalModel {
billingModel = input.ChannelMappedModel
}
if input.BillingModelSource == BillingModelSourceRequested && input.OriginalModel != "" {
billingModel = input.OriginalModel
}
billingModels := usageBillingModelCandidates(
billingModel,
result.BillingModel,
input.ChannelMappedModel,
input.OriginalModel,
result.UpstreamModel,
result.Model,
)
billingModels = s.filterCNProviderBillingModelCandidates(ctx, account, apiKey, billingModels)
serviceTier := ""
if result.ServiceTier != nil {
serviceTier = strings.TrimSpace(*result.ServiceTier)
}
billingAccount := account
if account.IsShadow() {
billingAccount, err = resolveCredentialAccount(ctx, s.accountRepo, account)
if err != nil {
return err
}
}
longContextBillingGate := openAILongContextBillingGate(billingAccount)
cost, err = s.calculateOpenAIRecordUsageCost(
ctx,
result,
apiKey,
billingModels,
multiplier,
imageMultiplier,
videoMultiplier,
baseMultiplier,
tokens,
serviceTier,
longContextBillingGate,
pricingAt,
)
if err != nil {
if !isUsagePricingUnavailableError(err) {
return err
}
logger.L().With(
zap.String("component", "service.openai_gateway"),
zap.Strings("billing_models", billingModels),
zap.String("requested_model", input.OriginalModel),
zap.String("mapped_model", input.ChannelMappedModel),
zap.String("upstream_model", result.UpstreamModel),
zap.Int64("api_key_id", apiKey.ID),
zap.Int64("account_id", account.ID),
).Warn("openai_usage.pricing_missing_record_zero_cost", zap.Error(err))
cost = &CostBreakdown{BillingMode: string(BillingModeToken)}
}
// response_model:按上游成功响应自报的模型计费(渠道显式开启才生效)。
// 采纳条件见 responseModelBillingDeclaration + hasIdentifiedOpenAIResponsePricing
// + responseModelBillingAdoptable。任一条件不满足都静默回落基线,即开启本模式前的
// 既有行为。响应模型与基线同名时直接跳过:重算必然同价,白跑一次定价解析。
baselineBillingModel := firstUsageBillingModel(billingModels)
if responseModel := responseModelBillingDeclaration(
input.BillingModelSource,
result.UpstreamResponseModel,
result.UpstreamResponseModelConflict,
result.ImageCount > 0 || result.VideoCount > 0 || result.WebSearchCalls > 0 ||
result.AudioUsage != nil || result.SearchCount > 0,
); responseModel != "" && !strings.EqualFold(responseModel, baselineBillingModel) {
if identified, responseChannelPriced := s.hasIdentifiedOpenAIResponsePricing(ctx, responseModel, apiKey); identified {
responseModels := s.filterCNProviderBillingModelCandidates(ctx, account, apiKey, usageBillingModelCandidates(responseModel))
responseCost, responseErr := s.calculateOpenAIRecordUsageCost(
ctx, result, apiKey, responseModels, multiplier, imageMultiplier,
videoMultiplier, baseMultiplier, tokens, serviceTier, longContextBillingGate, pricingAt,
)
// 基线定价源以 baselineBillingModel 为准:它正是 calculateOpenAIRecordUsageCost
// 内部做渠道定价判断时使用的模型,且"首候选有渠道价"必然意味着首候选就是实际
// 定价基准(有渠道价就一定能算出价,循环不会落到后续候选)。
baselineChannelPriced := s.resolveOpenAIChannelPricing(ctx, baselineBillingModel, apiKey) != nil
if responseErr == nil && responseModelBillingAdoptable(cost, responseCost, baselineChannelPriced, responseChannelPriced) {
logResponseModelBillingApplied("service.openai_gateway", account, result.RequestID,
baselineBillingModel, responseModel, cost, responseCost)
billingModels = responseModels
cost = responseCost
}
}
}
// Determine billing type
isSubscriptionBilling := subscription != nil && apiKey.Group != nil && apiKey.Group.IsSubscriptionType()
billingType := BillingTypeBalance
if isSubscriptionBilling {
billingType = BillingTypeSubscription
}
// Create usage log
durationMs := int(result.Duration.Milliseconds())
accountRateMultiplier := account.BillingRateMultiplier()
requestID := resolveUsageBillingRequestID(ctx, result.RequestID)
if result.OpenAIWSMode {
if upstreamRequestID := strings.TrimSpace(result.RequestID); upstreamRequestID != "" {
requestID = upstreamRequestID
}
}
// Async Grok video: always use the stable task id for dedup (status + content polls
// share one bill). Context-local client/local IDs would otherwise create a new row
// per poll if Redis claim is lost.
if result.VideoCount > 0 {
if stable := StableGrokVideoBillingRequestID(firstNonEmpty(
strings.TrimPrefix(strings.TrimSpace(result.RequestID), "grok-video:"),
strings.TrimSpace(result.ResponseID),
strings.TrimPrefix(strings.TrimSpace(requestID), "grok-video:"),
)); stable != "" {
requestID = stable
}
}
// 确定 RequestedModel(渠道映射前的原始模型)
requestedModel := result.Model
if input.OriginalModel != "" {
requestedModel = input.OriginalModel
}
sentModel := upstreamSentModel(result.Model, result.UpstreamModel)
if result.UpstreamResponseModelConflict {
logger.L().Warn("upstream_response_model_conflict",
zap.String("platform", account.Platform),
zap.Int64("account_id", account.ID),
zap.String("request_id", requestID),
zap.String("sent_model", sentModel),
zap.String("selected_response_model", strings.TrimSpace(result.UpstreamResponseModel)),
)
}
usageLog := &UsageLog{
UserID: user.ID,
APIKeyID: apiKey.ID,
AccountID: account.ID,
RequestID: requestID,
Model: result.Model,
RequestedModel: requestedModel,
UpstreamModel: optionalTrimmedStringPtr(result.UpstreamModel),
UpstreamResponseModel: optionalTrimmedStringPtr(result.UpstreamResponseModel),
UpstreamModelMismatch: upstreamModelMismatch(sentModel, result.UpstreamResponseModel),
ServiceTier: result.ServiceTier,
ReasoningEffort: result.ReasoningEffort,
InboundEndpoint: optionalTrimmedStringPtr(input.InboundEndpoint),
UpstreamEndpoint: optionalTrimmedStringPtr(input.UpstreamEndpoint),
InputTokens: actualInputTokens,
OutputTokens: result.Usage.OutputTokens,
CacheCreationTokens: result.Usage.CacheCreationInputTokens,
CacheReadTokens: result.Usage.CacheReadInputTokens,
ImageInputTokens: result.Usage.ImageInputTokens,
ImageOutputTokens: result.Usage.ImageOutputTokens,
ImageCount: result.ImageCount,
ImageSize: optionalTrimmedStringPtr(result.ImageSize),
ImageInputSize: optionalTrimmedStringPtr(result.ImageInputSize),
ImageOutputSize: optionalTrimmedStringPtr(result.ImageOutputSize),
ImageSizeSource: optionalTrimmedStringPtr(result.ImageSizeSource),
ImageSizeBreakdown: result.ImageSizeBreakdown,
}
isVideoUsage := isGrokVideoUsageResult(result, billingModels)
if isVideoUsage {
usageLog.VideoCount = result.VideoCount
usageLog.VideoResolution = optionalTrimmedStringPtr(NormalizeVideoBillingResolutionOrDefault(result.VideoResolution))
videoDurationSeconds := NormalizeVideoBillingDurationSecondsOrDefault(result.VideoDurationSeconds)
usageLog.VideoDurationSeconds = &videoDurationSeconds
}
if cost != nil {
usageLog.InputCost = cost.InputCost
usageLog.ImageInputCost = cost.ImageInputCost
usageLog.OutputCost = cost.OutputCost
usageLog.ImageOutputCost = cost.ImageOutputCost
usageLog.CacheCreationCost = cost.CacheCreationCost
usageLog.CacheReadCost = cost.CacheReadCost
usageLog.TotalCost = cost.TotalCost
usageLog.ActualCost = cost.ActualCost
usageLog.LongContextBillingApplied = cost.LongContextBillingApplied
}
if isVideoUsage && (cost == nil || cost.BillingMode != string(BillingModeToken)) {
usageLog.RateMultiplier = videoMultiplier
} else if result.ImageCount > 0 && (cost == nil || cost.BillingMode != string(BillingModeToken)) {
usageLog.RateMultiplier = imageMultiplier
} else {
usageLog.RateMultiplier = multiplier
}
usageLog.AccountRateMultiplier = &accountRateMultiplier
usageLog.BillingType = billingType
usageLog.Stream = result.Stream
if input.CyberBlocked {
usageLog.RequestType = RequestTypeCyberBlocked
}
usageLog.OpenAIWSMode = result.OpenAIWSMode
usageLog.DurationMs = &durationMs
usageLog.FirstTokenMs = result.FirstTokenMs
usageLog.CreatedAt = time.Now()
// 设置渠道信息
usageLog.ChannelID = optionalInt64Ptr(input.ChannelID)
usageLog.ModelMappingChain = optionalTrimmedStringPtr(input.ModelMappingChain)
// 设置计费模式
if cost != nil && cost.BillingMode != "" {
billingMode := cost.BillingMode
usageLog.BillingMode = &billingMode
} else if isVideoUsage {
billingMode := string(BillingModeVideo)
usageLog.BillingMode = &billingMode
} else if result.ImageCount > 0 {
billingMode := string(BillingModeImage)
usageLog.BillingMode = &billingMode
} else {
billingMode := string(BillingModeToken)
usageLog.BillingMode = &billingMode
}
// 添加 UserAgent
if input.UserAgent != "" {
usageLog.UserAgent = &input.UserAgent
}
// 添加 IPAddress
if input.IPAddress != "" {
usageLog.IPAddress = &input.IPAddress
}
// 添加 SessionID(客户端显式会话标识;缺失/无效时保持 nil)
usageLog.SessionID = optionalTrimmedStringPtr(input.SessionID)
if apiKey.GroupID != nil {
usageLog.GroupID = apiKey.GroupID
}
if subscription != nil {
usageLog.SubscriptionID = &subscription.ID
}
// 计算账号统计定价费用(使用最终上游模型匹配自定义规则)
if apiKey.GroupID != nil {
applyAccountStatsCost(ctx, usageLog, s.channelService, s.billingService,
account.ID, *apiKey.GroupID, result.UpstreamModel, result.Model,
tokens, cost.TotalCost,
)
}
if s.cfg != nil && s.cfg.RunMode == config.RunModeSimple {
writeUsageLogBestEffort(ctx, s.usageLogRepo, usageLog, "service.openai_gateway")
logger.LegacyPrintf("service.openai_gateway", "[SIMPLE MODE] Usage recorded (not billed): user=%d, tokens=%d", usageLog.UserID, usageLog.TotalTokens())
s.deferredService.ScheduleLastUsedUpdate(account.ID)
return nil
}
// Async usage billing runs outside the original request context, so it
// cannot recover ForcePlatform there. Fall back for internal/test callers.
quotaPlatform := input.QuotaPlatform
if quotaPlatform == "" {
quotaPlatform = PlatformFromAPIKey(apiKey)
}
billingErr := func() error {
_, err := applyUsageBilling(ctx, requestID, usageLog, &postUsageBillingParams{
Cost: cost,
User: user,
APIKey: apiKey,
Account: account,
Subscription: subscription,
RequestPayloadHash: resolveUsageBillingPayloadFingerprint(ctx, input.RequestPayloadHash),
IsSubscriptionBill: isSubscriptionBilling,
AccountRateMultiplier: accountRateMultiplier,
APIKeyService: input.APIKeyService,
Platform: quotaPlatform,
}, s.billingDeps(), s.usageBillingRepo)
return err
}()
if billingErr != nil {
usageLog.ActualCost = 0
writeUsageLogBestEffort(ctx, s.usageLogRepo, usageLog, "service.openai_gateway")
return billingErr
}
writeUsageLogBestEffort(ctx, s.usageLogRepo, usageLog, "service.openai_gateway")
return nil
}
// hasIdentifiedOpenAIResponsePricing 判断上游自报的响应模型是否可以作为计费基准,
// 并回传它是否解析到了渠道级定价(供 responseModelBillingAdoptable 的跨定价源守卫使用,
// 避免为此再解析一次)。
// 只接受管理员为该模型显式配置的渠道定价,或价格表中能被确定性识别的条目;
// 刻意不接受按子串猜出来的系列兜底价,否则上游随便编一个含 "haiku" 的名字就能把
// 计费拉到最便宜的系列价上。详见 responseModelBillingDeclaration。
func (s *OpenAIGatewayService) hasIdentifiedOpenAIResponsePricing(ctx context.Context, model string, apiKey *APIKey) (identified bool, channelPriced bool) {
model = strings.TrimSpace(model)
if model == "" {
return false, false
}
if s.resolveOpenAIChannelPricing(ctx, model, apiKey) != nil {
return true, true
}
return s.billingService.HasIdentifiedTokenPricing(model), false
}
// openAILongContextBillingGate returns the per-account long-context opt-in.
// The flag is an OpenAI-only account setting, so other platforms (Grok) return
// nil — "no per-account gate" — and are governed by the group toggle alone.
// Returning a hardcoded false for them would veto the official model ladders
// (e.g. the Grok >=200k 2x card) that no account setting can ever re-enable.
func openAILongContextBillingGate(account *Account) *bool {
if account == nil || !account.IsOpenAI() {
return nil
}
enabled := account.IsOpenAILongContextBillingEnabled()
return &enabled
}
func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost(
ctx context.Context,
result *OpenAIForwardResult,
apiKey *APIKey,
billingModels []string,
multiplier float64,
imageMultiplier float64,
videoMultiplier float64,
webSearchMultiplier float64,
tokens UsageTokens,
serviceTier string,
longContextBillingGate *bool,
pricingAt time.Time,
) (*CostBreakdown, error) {
billingModel := firstUsageBillingModel(billingModels)
if result != nil && result.WebSearchCalls > 0 {
// Codex alpha/search 网页搜索按次计费:上游不返回 usage/token 字段,单价只取
// 分组覆盖价(nil 时默认 0.01 = 官方 $10/1000 次),不参与渠道级模型定价。
// 倍率与 image/video 按次口径一致:使用不含高峰因子的基础倍率
//(用户专属 > 分组 rate_multiplier > 系统默认),与分组表单的价格预览承诺一致。
return s.billingService.CalculateWebSearchCost(result.WebSearchCalls, webSearchPricePerCallFromAPIKey(apiKey), webSearchMultiplier), nil
}
if isGrokVideoUsageResult(result, billingModels) {
if resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey); resolved == nil || resolved.Mode != BillingModeToken {
return s.calculateOpenAIVideoCost(ctx, billingModel, apiKey, result, videoMultiplier), nil
}
}
if result != nil && result.AudioUsage != nil {
if resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey); resolved != nil &&
(resolved.Mode == BillingModePerRequest) {
gid := apiKey.Group.ID
return s.billingService.CalculateCostUnified(CostInput{
Ctx: ctx, Model: billingModel, GroupID: &gid, Group: apiKey.Group,
UsageUnits: result.AudioUsage.DurationOrUnits, SizeTier: result.AudioUsage.Mode,
RateMultiplier: webSearchMultiplier, Resolver: s.resolver, Resolved: resolved,
})
}
cfg := groupAudioPriceConfigFromAPIKey(apiKey)
return s.billingService.CalculateAudioCost(result.AudioUsage.Mode, result.AudioUsage.DurationOrUnits, cfg, webSearchMultiplier), nil
}
if result != nil && result.ImageCount > 0 {
// 渠道定价为 token 计费时走 token 路径,否则走图片计费
if resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey); resolved == nil || resolved.Mode != BillingModeToken {
return s.calculateOpenAIImageCost(ctx, billingModel, apiKey, result, imageMultiplier), nil
}
}
// Token path (optional search surcharge is additive — never replaces token cost).
var tokenCost *CostBreakdown
var lastErr error
if len(billingModels) > 0 && billingModel != "" {
for _, candidate := range billingModels {
candidate = strings.TrimSpace(candidate)
if candidate == "" {
continue
}
cost, err := s.calculateOpenAIRecordUsageTokenCost(
ctx,
apiKey,
candidate,
multiplier,
pricingAt,
tokens,
serviceTier,
longContextBillingGate,
)
if err == nil {
tokenCost = cost
break
}
lastErr = err
}
}
// Search surcharge is additive. Never let a zero/default search cost mask a
// real token-pricing failure for requests that attempted token billing.
searchCost := (*CostBreakdown)(nil)
if result != nil && result.SearchCount > 0 {
price := groupSearchPricePer1kFromAPIKey(apiKey)
if price != nil && *price == 0 {
logger.L().Info("openai_usage.search_price_per_1k_explicit_free",
zap.Int("search_count", result.SearchCount),
zap.String("model", billingModel),
zap.Int64("api_key_id", apiKey.ID),
zap.Any("group_id", apiKey.GroupID),
)
}
searchCost = s.billingService.CalculateSearchCost(result.SearchCount, price, webSearchMultiplier)
}
tokenBillingAttempted := len(billingModels) > 0 && billingModel != ""
if tokenCost == nil {
if tokenBillingAttempted {
if lastErr == nil {
lastErr = fmt.Errorf("%w: no non-empty billing model candidates", ErrModelPricingUnavailable)
}
return nil, fmt.Errorf("calculate OpenAI usage cost failed for billing models %s: %w", strings.Join(billingModels, ","), lastErr)
}
// Search-only (no model / pure tool path): allow search billing alone.
if searchCost != nil {
return searchCost, nil
}
// 空候选按「无价可循」处理并携带 ErrModelPricingUnavailable:上层据此走
// 零成本+告警落账,而不是丢弃整条 usage 记录。CN 账号的 claude-* 候选被
// filterCNProviderBillingModelCandidates 全数过滤后即落到这里。
if lastErr == nil {
lastErr = fmt.Errorf("%w: openai usage billing model is empty", ErrModelPricingUnavailable)
}
return nil, fmt.Errorf("calculate OpenAI usage cost failed for billing models %s: %w", strings.Join(billingModels, ","), lastErr)
}
if searchCost == nil || (searchCost.TotalCost == 0 && searchCost.ActualCost == 0) {
return tokenCost, nil
}
// Additive: tokens + search surcharge.
tokenCost.TotalCost += searchCost.TotalCost
tokenCost.ActualCost += searchCost.ActualCost
return tokenCost, nil
}
func isGrokVideoBillingModel(model string) bool {
return strings.HasPrefix(strings.ToLower(strings.TrimSpace(model)), "grok-imagine-video")
}
func isGrokVideoUsageResult(result *OpenAIForwardResult, billingModels []string) bool {
if result == nil || result.VideoCount <= 0 {
return false
}
// VideoCount alone is authoritative for async video completion billing.
// Prefer model-family match when present; never drop video mode on rename/mapping.
candidates := append([]string{}, billingModels...)
candidates = append(candidates, result.BillingModel, result.Model, result.UpstreamModel)
for _, candidate := range candidates {
if isGrokVideoBillingModel(candidate) {
return true
}
}
return true
}
func isUsagePricingUnavailableError(err error) bool {
if err == nil {
return false
}
if errors.Is(err, ErrModelPricingUnavailable) {
return true
}
msg := strings.ToLower(err.Error())
return strings.Contains(msg, "no pricing available") || strings.Contains(msg, "pricing not found")
}
func (s *OpenAIGatewayService) calculateOpenAIRecordUsageTokenCost(
ctx context.Context,
apiKey *APIKey,
billingModel string,
multiplier float64,
pricingAt time.Time,
tokens UsageTokens,
serviceTier string,
longContextBillingGate *bool,
) (*CostBreakdown, error) {
if s.resolver != nil && apiKey.Group != nil {
gid := apiKey.Group.ID
return s.billingService.CalculateCostUnified(CostInput{
Ctx: ctx, Model: billingModel, GroupID: &gid, Group: apiKey.Group,
Tokens: tokens, RequestCount: 1, RateMultiplier: multiplier, PricingAt: pricingAt,
ServiceTier: serviceTier, Resolver: s.resolver,
LongContextBillingEnabled: longContextBillingGate,
})
}
return s.billingService.calculateCostWithServiceTierPolicy(
billingModel,
tokens,
multiplier,
serviceTier,
longContextBillingGate == nil || *longContextBillingGate,
)
}
func (s *OpenAIGatewayService) calculateOpenAIImageCost(
ctx context.Context,
billingModel string,
apiKey *APIKey,
result *OpenAIForwardResult,
multiplier float64,
) *CostBreakdown {
sizeTier := NormalizeImageBillingTierOrDefault(result.ImageSize)
resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey)
if resolved != nil && resolved.Source == PricingSourceGroup &&
(resolved.Mode == BillingModePerRequest || resolved.Mode == BillingModeImage) {
gid := apiKey.Group.ID
cost, err := s.billingService.CalculateCostUnified(CostInput{
Ctx: ctx, Model: billingModel, GroupID: &gid, Group: apiKey.Group,
RequestCount: result.ImageCount, SizeTier: sizeTier,
RateMultiplier: multiplier, Resolver: s.resolver, Resolved: resolved,
})
if err == nil {
return cost
}
}
groupConfig := imagePriceConfigFromAPIKey(apiKey)
if apiKeyHasConfiguredImagePrice(apiKey, sizeTier) {
return s.billingService.CalculateImageCost(billingModel, sizeTier, result.ImageCount, groupConfig, multiplier)
}
if refreshed := s.apiKeyWithFreshGroupMediaPricing(ctx, apiKey); refreshed != apiKey {
apiKey = refreshed
groupConfig = imagePriceConfigFromAPIKey(apiKey)
if apiKeyHasConfiguredImagePrice(apiKey, sizeTier) {
return s.billingService.CalculateImageCost(billingModel, sizeTier, result.ImageCount, groupConfig, multiplier)
}
}
if resolved != nil && resolved.Source == PricingSourceChannel &&
(resolved.Mode == BillingModePerRequest || resolved.Mode == BillingModeImage) {
gid := apiKey.Group.ID
cost, err := s.billingService.CalculateCostUnified(CostInput{
Ctx: ctx,
Model: billingModel,
GroupID: &gid,
Group: apiKey.Group,
RequestCount: result.ImageCount,
SizeTier: sizeTier,
RateMultiplier: multiplier,
Resolver: s.resolver,
Resolved: resolved,
})
if err == nil {
return cost
}
logger.LegacyPrintf("service.openai_gateway", "Calculate image channel cost failed: %v", err)
}
return s.billingService.CalculateImageCost(billingModel, sizeTier, result.ImageCount, groupConfig, multiplier)
}
func (s *OpenAIGatewayService) calculateOpenAIVideoCost(
ctx context.Context,
billingModel string,
apiKey *APIKey,
result *OpenAIForwardResult,
multiplier float64,
) *CostBreakdown {
videoCount := result.VideoCount
if videoCount <= 0 {
videoCount = 1
}
resolution := NormalizeVideoBillingResolutionOrDefault(result.VideoResolution)
durationSeconds := NormalizeVideoBillingDurationSecondsOrDefault(result.VideoDurationSeconds)
resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey)
if resolved != nil && resolved.Source == PricingSourceGroup && resolved.Mode == BillingModeVideo {
gid := apiKey.Group.ID
cost, err := s.billingService.CalculateCostUnified(CostInput{
Ctx: ctx, Model: billingModel, GroupID: &gid, Group: apiKey.Group,
UsageUnits: float64(videoCount * durationSeconds), SizeTier: resolution,
RateMultiplier: multiplier, Resolver: s.resolver, Resolved: resolved,
})
if err == nil {
return cost
}
}
groupConfig := videoPriceConfigFromAPIKey(apiKey)
if apiKeyHasConfiguredVideoPrice(apiKey, billingModel, resolution) {
return s.billingService.CalculateVideoCost(billingModel, resolution, videoCount, durationSeconds, groupConfig, multiplier)
}
if refreshed := s.apiKeyWithFreshGroupMediaPricing(ctx, apiKey); refreshed != apiKey {
apiKey = refreshed
groupConfig = videoPriceConfigFromAPIKey(apiKey)
if apiKeyHasConfiguredVideoPrice(apiKey, billingModel, resolution) {
return s.billingService.CalculateVideoCost(billingModel, resolution, videoCount, durationSeconds, groupConfig, multiplier)
}
}
if resolved != nil && resolved.Source == PricingSourceChannel &&
(resolved.Mode == BillingModePerRequest || resolved.Mode == BillingModeImage || resolved.Mode == BillingModeVideo) {
// 渠道 per_request/image 定价保持"按请求次数"口径(价格由管理员按次配置),不乘视频时长。
gid := apiKey.Group.ID
units := float64(videoCount)
if resolved.Mode == BillingModeVideo {
units = float64(videoCount * durationSeconds)
}
cost, err := s.billingService.CalculateCostUnified(CostInput{
Ctx: ctx,
Model: billingModel,
GroupID: &gid,
Group: apiKey.Group,
RequestCount: videoCount,
UsageUnits: units,
SizeTier: resolution,
RateMultiplier: multiplier,
Resolver: s.resolver,
Resolved: resolved,
})
if err == nil {
cost.BillingMode = string(BillingModeVideo)
return cost
}
logger.LegacyPrintf("service.openai_gateway", "Calculate video channel cost failed: %v", err)
}
return s.billingService.CalculateVideoCost(billingModel, resolution, videoCount, durationSeconds, groupConfig, multiplier)
}
func (s *OpenAIGatewayService) apiKeyWithFreshGroupMediaPricing(ctx context.Context, apiKey *APIKey) *APIKey {
if apiKey == nil || apiKey.GroupID == nil || *apiKey.GroupID <= 0 {
return apiKey
}
if !groupMediaPricingLooksIncomplete(apiKey.Group) {
return apiKey
}
if s == nil || s.channelService == nil || s.channelService.groupRepo == nil {
return apiKey
}
group, err := s.channelService.groupRepo.GetByIDLite(ctx, *apiKey.GroupID)
if err != nil || group == nil {
return apiKey
}
clone := *apiKey
clone.Group = group
return &clone
}
// groupMediaPricingLooksIncomplete 判断分组对象是否可能缺失媒体/搜索/语音计费字段
// (例如由不含这些字段的旧快照或手工构造的上下文对象生成)。image/video 独立倍率在
// 数据库中的默认值均为 1.0;正常加载的分组不可能两个倍率同时为 0 且未开启独立倍率、
// 全部媒体/搜索/语音价为 nil——只有这种情况才回源查库,避免对未配置覆盖价的分组每条
// 用量都多打一次 DB 查询。
//
// 注意:apiKeyAuthSnapshotVersion 升级会强制刷新存量快照;本函数是热路径上的二次兜底,
// 不能仅凭 legacy video_price_* 判定完整而跳过 VideoModelPrices/search/audio 的回源。
func groupMediaPricingLooksIncomplete(group *Group) bool {
if group == nil {
return true
}
if group.ImageRateIndependent || group.VideoRateIndependent {
return false
}
if group.ImageRateMultiplier != 0 || group.VideoRateMultiplier != 0 {
return false
}
// Any first-class pricing field present means the projection is not a blank shell.
if len(group.VideoModelPrices) > 0 {
return false
}
if len(group.ModelPricing) > 0 || group.LongContextPricingEnabled {
return false
}
if group.SearchPricePer1k != nil ||
group.AudioRealtimePricePerMin != nil ||
group.AudioTTSPricePerMillionChars != nil ||
group.AudioSTTPricePerHour != nil ||
group.WebSearchPricePerCall != nil {
return false
}
return group.ImagePrice1K == nil && group.ImagePrice2K == nil && group.ImagePrice4K == nil &&
group.VideoPrice480P == nil && group.VideoPrice720P == nil && group.VideoPrice1080P == nil
}
// filterCNProviderBillingModelCandidates 过滤国产供应商(kimi/zhipu/deepseek
// 账号的计费候选模型名:claude-* 候选仅在运营者显式配置了分组/渠道定价时保留。
//
// 背景:候选链的兜底候选含客户端请求的原始模型名。CN 上游的 Anthropic 兼容端点
// 接受 claude-* 模型名但从不真正服务 Claude 模型;若放行,目录里的 Claude 价卡
// 与 getFallbackPricing 的 "claude"→Sonnet 统一兜底会把 CN 流量按 Claude 原价
// (数倍~数十倍)静默误计,且 usage 日志显示的正是 claude-* 名,无从察觉。
// 候选全部落空时走既有的零成本+告警路径(openai_usage.pricing_missing_record_
// zero_cost),与定价层「未知型号不回退以避免误计价」的既有设计意图一致;
// 运营者的修复手段是配置账号级 model_mapping(映射到已定价的 CN 模型)或
// 分组/渠道显式定价。
func (s *OpenAIGatewayService) filterCNProviderBillingModelCandidates(ctx context.Context, account *Account, apiKey *APIKey, candidates []string) []string {
if account == nil || !account.IsCNProvider() {
return candidates
}
out := make([]string, 0, len(candidates))
for _, candidate := range candidates {
trimmed := strings.TrimSpace(candidate)
if trimmed == "" {
continue
}
if strings.Contains(strings.ToLower(trimmed), "claude") &&
s.resolveOpenAIChannelPricing(ctx, trimmed, apiKey) == nil {
continue
}
out = append(out, candidate)
}
return out
}
func (s *OpenAIGatewayService) resolveOpenAIChannelPricing(ctx context.Context, billingModel string, apiKey *APIKey) *ResolvedPricing {
if s.resolver == nil || apiKey == nil || apiKey.Group == nil {
return nil
}
gid := apiKey.Group.ID
resolved := s.resolver.Resolve(ctx, PricingInput{Model: billingModel, GroupID: &gid, Group: apiKey.Group})
if resolved.Source == PricingSourceGroup || resolved.Source == PricingSourceChannel {
return resolved
}
return nil
}
// ParseCodexRateLimitHeaders extracts Codex usage limits from response headers.
// Exported for use in ratelimit_service when handling OpenAI 429 responses.
func ParseCodexRateLimitHeaders(headers http.Header) *OpenAICodexUsageSnapshot {
snapshot := &OpenAICodexUsageSnapshot{}
hasData := false
// Helper to parse float64 from header
parseFloat := func(key string) *float64 {
if v := headers.Get(key); v != "" {
if f, err := strconv.ParseFloat(v, 64); err == nil {
return &f
}
}
return nil
}
// Helper to parse int from header
parseInt := func(key string) *int {
if v := headers.Get(key); v != "" {
if i, err := strconv.Atoi(v); err == nil {
return &i
}
}
return nil
}
// Primary (weekly) limits
if v := parseFloat("x-codex-primary-used-percent"); v != nil {
snapshot.PrimaryUsedPercent = v
hasData = true
}
if v := parseInt("x-codex-primary-reset-after-seconds"); v != nil {
snapshot.PrimaryResetAfterSeconds = v
hasData = true
}
if v := parseInt("x-codex-primary-window-minutes"); v != nil {
snapshot.PrimaryWindowMinutes = v
hasData = true
}
// Secondary (5h) limits
if v := parseFloat("x-codex-secondary-used-percent"); v != nil {
snapshot.SecondaryUsedPercent = v
hasData = true
}
if v := parseInt("x-codex-secondary-reset-after-seconds"); v != nil {
snapshot.SecondaryResetAfterSeconds = v
hasData = true
}
if v := parseInt("x-codex-secondary-window-minutes"); v != nil {
snapshot.SecondaryWindowMinutes = v
hasData = true
}
// Overflow ratio
if v := parseFloat("x-codex-primary-over-secondary-limit-percent"); v != nil {
snapshot.PrimaryOverSecondaryPercent = v
hasData = true
}
if !hasData {
return nil
}
snapshot.UpdatedAt = time.Now().Format(time.RFC3339)
return snapshot
}
func codexSnapshotBaseTime(snapshot *OpenAICodexUsageSnapshot, fallback time.Time) time.Time {
if snapshot == nil {
return fallback
}
if snapshot.UpdatedAt == "" {
return fallback
}
base, err := time.Parse(time.RFC3339, snapshot.UpdatedAt)
if err != nil {
return fallback
}
return base
}
func codexResetAtRFC3339(base time.Time, resetAfterSeconds *int) *string {
if resetAfterSeconds == nil {
return nil
}
sec := *resetAfterSeconds
if sec < 0 {
sec = 0
}
resetAt := base.Add(time.Duration(sec) * time.Second).Format(time.RFC3339)
return &resetAt
}
func buildCodexUsageExtraUpdates(snapshot *OpenAICodexUsageSnapshot, fallbackNow time.Time) map[string]any {
if snapshot == nil {
return nil
}
baseTime := codexSnapshotBaseTime(snapshot, fallbackNow)
updates := make(map[string]any)
// 保存原始 primary/secondary 字段,便于排查问题
if snapshot.PrimaryUsedPercent != nil {
updates["codex_primary_used_percent"] = *snapshot.PrimaryUsedPercent
}
if snapshot.PrimaryResetAfterSeconds != nil {
updates["codex_primary_reset_after_seconds"] = *snapshot.PrimaryResetAfterSeconds
}
if snapshot.PrimaryWindowMinutes != nil {
updates["codex_primary_window_minutes"] = *snapshot.PrimaryWindowMinutes
}
if snapshot.SecondaryUsedPercent != nil {
updates["codex_secondary_used_percent"] = *snapshot.SecondaryUsedPercent
}
if snapshot.SecondaryResetAfterSeconds != nil {
updates["codex_secondary_reset_after_seconds"] = *snapshot.SecondaryResetAfterSeconds
}
if snapshot.SecondaryWindowMinutes != nil {
updates["codex_secondary_window_minutes"] = *snapshot.SecondaryWindowMinutes
}
if snapshot.PrimaryOverSecondaryPercent != nil {
updates["codex_primary_over_secondary_percent"] = *snapshot.PrimaryOverSecondaryPercent
}
updates["codex_usage_updated_at"] = baseTime.Format(time.RFC3339)
// 归一化到 5h/7d 规范字段
if normalized := snapshot.Normalize(); normalized != nil {
if normalized.Used5hPercent != nil {
updates["codex_5h_used_percent"] = *normalized.Used5hPercent
}
if normalized.Reset5hSeconds != nil {
updates["codex_5h_reset_after_seconds"] = *normalized.Reset5hSeconds
}
if normalized.Window5hMinutes != nil {
updates["codex_5h_window_minutes"] = *normalized.Window5hMinutes
}
if normalized.Used7dPercent != nil {
updates["codex_7d_used_percent"] = *normalized.Used7dPercent
}
if normalized.Reset7dSeconds != nil {
updates["codex_7d_reset_after_seconds"] = *normalized.Reset7dSeconds
}
if normalized.Window7dMinutes != nil {
updates["codex_7d_window_minutes"] = *normalized.Window7dMinutes
}
if reset5hAt := codexResetAtRFC3339(baseTime, normalized.Reset5hSeconds); reset5hAt != nil {
updates["codex_5h_reset_at"] = *reset5hAt
}
if reset7dAt := codexResetAtRFC3339(baseTime, normalized.Reset7dSeconds); reset7dAt != nil {
updates["codex_7d_reset_at"] = *reset7dAt
}
}
return updates
}
// updateCodexUsageSnapshot saves the Codex usage snapshot to account's Extra field
// updateCodexUsageSnapshot 把 /responses 的 x-codex-* 全局头快照写入账号 codex_* Extra。
// ⚠️ 调用方必须排除 spark 影子账号(account.IsShadow()):影子的 codex_* 仅由 QueryUsage
// (/wham/usage bengalfox 道)更新,不能被全局头口径污染(外审第7轮 P1)。本函数仅持 accountID,
// 无法在此自检影子,故守卫前置到各调用点。
func (s *OpenAIGatewayService) updateCodexUsageSnapshot(ctx context.Context, accountID int64, snapshot *OpenAICodexUsageSnapshot) {
if snapshot == nil {
return
}
if s == nil || s.accountRepo == nil {
return
}
now := time.Now()
updates := buildCodexUsageExtraUpdates(snapshot, now)
if len(updates) == 0 {
return
}
if !s.getCodexSnapshotThrottle().Allow(accountID, now) {
return
}
go func() {
updateCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
_ = s.accountRepo.UpdateExtra(updateCtx, accountID, updates)
}()
}
func (s *OpenAIGatewayService) UpdateCodexUsageSnapshotFromHeaders(ctx context.Context, accountID int64, headers http.Header) {
if accountID <= 0 || headers == nil {
return
}
if snapshot := ParseCodexRateLimitHeaders(headers); snapshot != nil {
s.updateCodexUsageSnapshot(ctx, accountID, snapshot)
}
}