refactor(gateway): snapshot usage worker inputs
This commit is contained in:
parent
619e5ae619
commit
2caee9d884
@ -510,11 +510,12 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
|
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
|
||||||
|
// ForceCacheBilling 提前拍成标量,避免 worker 闭包保活 failover 状态里的响应体。
|
||||||
|
forceCacheBilling := fs.ForceCacheBilling
|
||||||
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
|
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
|
||||||
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
|
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
|
||||||
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
|
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
|
||||||
Result: result,
|
Result: result,
|
||||||
ParsedRequest: parsedReq,
|
|
||||||
QuotaPlatform: quotaPlatform,
|
QuotaPlatform: quotaPlatform,
|
||||||
APIKey: apiKey,
|
APIKey: apiKey,
|
||||||
User: apiKey.User,
|
User: apiKey.User,
|
||||||
@ -525,7 +526,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
|
|||||||
UserAgent: userAgent,
|
UserAgent: userAgent,
|
||||||
IPAddress: clientIP,
|
IPAddress: clientIP,
|
||||||
RequestPayloadHash: requestPayloadHash,
|
RequestPayloadHash: requestPayloadHash,
|
||||||
ForceCacheBilling: fs.ForceCacheBilling,
|
ForceCacheBilling: forceCacheBilling,
|
||||||
APIKeyService: h.apiKeyService,
|
APIKeyService: h.apiKeyService,
|
||||||
ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
|
ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
@ -918,11 +919,12 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
|
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
|
||||||
|
// ForceCacheBilling 提前拍成标量,避免 worker 闭包保活 failover 状态里的响应体。
|
||||||
|
forceCacheBilling := fs.ForceCacheBilling
|
||||||
quotaPlatform := service.QuotaPlatform(c.Request.Context(), currentAPIKey)
|
quotaPlatform := service.QuotaPlatform(c.Request.Context(), currentAPIKey)
|
||||||
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
|
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
|
||||||
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
|
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
|
||||||
Result: result,
|
Result: result,
|
||||||
ParsedRequest: attemptParsedReq,
|
|
||||||
QuotaPlatform: quotaPlatform,
|
QuotaPlatform: quotaPlatform,
|
||||||
APIKey: currentAPIKey,
|
APIKey: currentAPIKey,
|
||||||
User: currentAPIKey.User,
|
User: currentAPIKey.User,
|
||||||
@ -933,7 +935,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
|
|||||||
UserAgent: userAgent,
|
UserAgent: userAgent,
|
||||||
IPAddress: clientIP,
|
IPAddress: clientIP,
|
||||||
RequestPayloadHash: requestPayloadHash,
|
RequestPayloadHash: requestPayloadHash,
|
||||||
ForceCacheBilling: fs.ForceCacheBilling,
|
ForceCacheBilling: forceCacheBilling,
|
||||||
APIKeyService: h.apiKeyService,
|
APIKeyService: h.apiKeyService,
|
||||||
ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
|
ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
|
|||||||
@ -527,6 +527,8 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
|
|||||||
requestPayloadHash := service.HashUsageRequestPayload(body)
|
requestPayloadHash := service.HashUsageRequestPayload(body)
|
||||||
inboundEndpoint := GetInboundEndpoint(c)
|
inboundEndpoint := GetInboundEndpoint(c)
|
||||||
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
|
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
|
||||||
|
// ForceCacheBilling 提前拍成标量,避免 worker 闭包保活 failover 状态里的响应体。
|
||||||
|
forceCacheBilling := fs.ForceCacheBilling
|
||||||
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
|
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
|
||||||
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
|
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
|
||||||
if err := h.gatewayService.RecordUsageWithLongContext(ctx, &service.RecordUsageLongContextInput{
|
if err := h.gatewayService.RecordUsageWithLongContext(ctx, &service.RecordUsageLongContextInput{
|
||||||
@ -543,7 +545,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
|
|||||||
RequestPayloadHash: requestPayloadHash,
|
RequestPayloadHash: requestPayloadHash,
|
||||||
LongContextThreshold: 200000, // Gemini 200K 阈值
|
LongContextThreshold: 200000, // Gemini 200K 阈值
|
||||||
LongContextMultiplier: 2.0, // 超出部分双倍计费
|
LongContextMultiplier: 2.0, // 超出部分双倍计费
|
||||||
ForceCacheBilling: fs.ForceCacheBilling,
|
ForceCacheBilling: forceCacheBilling,
|
||||||
APIKeyService: h.apiKeyService,
|
APIKeyService: h.apiKeyService,
|
||||||
ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
|
ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
|
|||||||
@ -1379,6 +1379,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
|||||||
zap.Int("candidate_count", scheduleDecision.CandidateCount),
|
zap.Int("candidate_count", scheduleDecision.CandidateCount),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
var requestPayloadHash string
|
||||||
hooks := &service.OpenAIWSIngressHooks{
|
hooks := &service.OpenAIWSIngressHooks{
|
||||||
InitialRequestModel: reqModel,
|
InitialRequestModel: reqModel,
|
||||||
BeforeRequest: func(turn int, payload []byte, originalModel string) error {
|
BeforeRequest: func(turn int, payload []byte, originalModel string) error {
|
||||||
@ -1464,7 +1465,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
|||||||
UpstreamEndpoint: upstreamEndpoint,
|
UpstreamEndpoint: upstreamEndpoint,
|
||||||
UserAgent: userAgent,
|
UserAgent: userAgent,
|
||||||
IPAddress: clientIP,
|
IPAddress: clientIP,
|
||||||
RequestPayloadHash: service.HashUsageRequestPayload(firstMessage),
|
RequestPayloadHash: requestPayloadHash,
|
||||||
APIKeyService: h.apiKeyService,
|
APIKeyService: h.apiKeyService,
|
||||||
ChannelUsageFields: channelMappingWS.ToUsageFields(reqModel, result.UpstreamModel),
|
ChannelUsageFields: channelMappingWS.ToUsageFields(reqModel, result.UpstreamModel),
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
@ -1484,6 +1485,9 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
|||||||
wsFirstMessage = h.gatewayService.ReplaceModelInBody(firstMessage, channelMappingWS.MappedModel)
|
wsFirstMessage = h.gatewayService.ReplaceModelInBody(firstMessage, channelMappingWS.MappedModel)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// WebSocket 首包可能很大,hash 必须在 hooks 外算成字符串,避免 AfterTurn 闭包保活请求体。
|
||||||
|
requestPayloadHash = service.HashUsageRequestPayload(wsFirstMessage)
|
||||||
|
|
||||||
if err := h.gatewayService.ProxyResponsesWebSocketFromClient(ctx, c, wsConn, account, token, wsFirstMessage, hooks); err != nil {
|
if err := h.gatewayService.ProxyResponsesWebSocketFromClient(ctx, c, wsConn, account, token, wsFirstMessage, hooks); err != nil {
|
||||||
var failoverErr *service.UpstreamFailoverError
|
var failoverErr *service.UpstreamFailoverError
|
||||||
if errors.As(err, &failoverErr) {
|
if errors.As(err, &failoverErr) {
|
||||||
|
|||||||
@ -73,9 +73,10 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
|
|||||||
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", err.Error())
|
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
requestModel := parsed.Model
|
||||||
|
|
||||||
reqLog = reqLog.With(
|
reqLog = reqLog.With(
|
||||||
zap.String("model", parsed.Model),
|
zap.String("model", requestModel),
|
||||||
zap.Bool("stream", parsed.Stream),
|
zap.Bool("stream", parsed.Stream),
|
||||||
zap.Bool("multipart", parsed.Multipart),
|
zap.Bool("multipart", parsed.Multipart),
|
||||||
zap.String("capability", string(parsed.RequiredCapability)),
|
zap.String("capability", string(parsed.RequiredCapability)),
|
||||||
@ -85,7 +86,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
|
|||||||
h.errorResponse(c, http.StatusForbidden, "permission_error", service.ImageGenerationPermissionMessage())
|
h.errorResponse(c, http.StatusForbidden, "permission_error", service.ImageGenerationPermissionMessage())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIImages, parsed.Model, parsed.ModerationBody()); decision != nil && decision.Blocked {
|
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIImages, requestModel, parsed.ModerationBody()); decision != nil && decision.Blocked {
|
||||||
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
|
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@ -98,13 +99,13 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if parsed.Multipart {
|
if parsed.Multipart {
|
||||||
setOpsRequestContext(c, parsed.Model, parsed.Stream)
|
setOpsRequestContext(c, requestModel, parsed.Stream)
|
||||||
} else {
|
} else {
|
||||||
setOpsRequestContext(c, parsed.Model, parsed.Stream)
|
setOpsRequestContext(c, requestModel, parsed.Stream)
|
||||||
}
|
}
|
||||||
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(parsed.Stream, false)))
|
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(parsed.Stream, false)))
|
||||||
|
|
||||||
channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, parsed.Model)
|
channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, requestModel)
|
||||||
|
|
||||||
if h.errorPassthroughService != nil {
|
if h.errorPassthroughService != nil {
|
||||||
service.BindErrorPassthroughService(c, h.errorPassthroughService)
|
service.BindErrorPassthroughService(c, h.errorPassthroughService)
|
||||||
@ -147,7 +148,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
|
|||||||
c.Request.Context(),
|
c.Request.Context(),
|
||||||
apiKey.GroupID,
|
apiKey.GroupID,
|
||||||
sessionHash,
|
sessionHash,
|
||||||
parsed.Model,
|
requestModel,
|
||||||
failedAccountIDs,
|
failedAccountIDs,
|
||||||
parsed.RequiredCapability,
|
parsed.RequiredCapability,
|
||||||
)
|
)
|
||||||
@ -324,14 +325,14 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
|
|||||||
IPAddress: clientIP,
|
IPAddress: clientIP,
|
||||||
RequestPayloadHash: requestPayloadHash,
|
RequestPayloadHash: requestPayloadHash,
|
||||||
APIKeyService: h.apiKeyService,
|
APIKeyService: h.apiKeyService,
|
||||||
ChannelUsageFields: channelMapping.ToUsageFields(parsed.Model, upstreamModel),
|
ChannelUsageFields: channelMapping.ToUsageFields(requestModel, upstreamModel),
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
logger.L().With(
|
logger.L().With(
|
||||||
zap.String("component", "handler.openai_gateway.images"),
|
zap.String("component", "handler.openai_gateway.images"),
|
||||||
zap.Int64("user_id", subject.UserID),
|
zap.Int64("user_id", subject.UserID),
|
||||||
zap.Int64("api_key_id", apiKey.ID),
|
zap.Int64("api_key_id", apiKey.ID),
|
||||||
zap.Any("group_id", apiKey.GroupID),
|
zap.Any("group_id", apiKey.GroupID),
|
||||||
zap.String("model", parsed.Model),
|
zap.String("model", requestModel),
|
||||||
zap.Int64("account_id", account.ID),
|
zap.Int64("account_id", account.ID),
|
||||||
).Error("openai.images.record_usage_failed", zap.Error(err))
|
).Error("openai.images.record_usage_failed", zap.Error(err))
|
||||||
}
|
}
|
||||||
|
|||||||
@ -4729,6 +4729,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
|
|||||||
if retryErr == nil {
|
if retryErr == nil {
|
||||||
if retryResp.StatusCode < 400 {
|
if retryResp.StatusCode < 400 {
|
||||||
// 重试请求被上游接受后同步 ParsedRequest,保证 usage/日志看到真实请求体。
|
// 重试请求被上游接受后同步 ParsedRequest,保证 usage/日志看到真实请求体。
|
||||||
|
lastWireBody = retryWireBody
|
||||||
if err := replaceBody(retryWireBody); err != nil {
|
if err := replaceBody(retryWireBody); err != nil {
|
||||||
_ = retryResp.Body.Close()
|
_ = retryResp.Body.Close()
|
||||||
return nil, err
|
return nil, err
|
||||||
@ -4769,6 +4770,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
|
|||||||
if retryErr2 == nil {
|
if retryErr2 == nil {
|
||||||
if retryResp2.StatusCode < 400 {
|
if retryResp2.StatusCode < 400 {
|
||||||
// 二阶段工具块降级成功时也必须更新当前 body。
|
// 二阶段工具块降级成功时也必须更新当前 body。
|
||||||
|
lastWireBody = retryWireBody2
|
||||||
if err := replaceBody(retryWireBody2); err != nil {
|
if err := replaceBody(retryWireBody2); err != nil {
|
||||||
_ = retryResp2.Body.Close()
|
_ = retryResp2.Body.Close()
|
||||||
return nil, err
|
return nil, err
|
||||||
@ -4847,6 +4849,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
|
|||||||
if retryErr == nil {
|
if retryErr == nil {
|
||||||
if budgetRetryResp.StatusCode < 400 {
|
if budgetRetryResp.StatusCode < 400 {
|
||||||
// budget 修正请求成功后,ParsedRequest 也要描述被接受的修正版。
|
// budget 修正请求成功后,ParsedRequest 也要描述被接受的修正版。
|
||||||
|
lastWireBody = budgetWireBody
|
||||||
if err := replaceBody(budgetWireBody); err != nil {
|
if err := replaceBody(budgetWireBody); err != nil {
|
||||||
_ = budgetRetryResp.Body.Close()
|
_ = budgetRetryResp.Body.Close()
|
||||||
return nil, err
|
return nil, err
|
||||||
@ -8228,10 +8231,10 @@ func (s *GatewayService) getUserGroupRateMultiplier(ctx context.Context, userID,
|
|||||||
return resolver.Resolve(ctx, userID, groupID, groupDefaultMultiplier)
|
return resolver.Resolve(ctx, userID, groupID, groupDefaultMultiplier)
|
||||||
}
|
}
|
||||||
|
|
||||||
// RecordUsageInput 记录使用量的输入参数
|
// RecordUsageInput 记录使用量的输入参数。
|
||||||
|
// 异步 worker 只接收计费所需快照,不能持有 ParsedRequest/RequestBodyRef 这类大请求体引用。
|
||||||
type RecordUsageInput struct {
|
type RecordUsageInput struct {
|
||||||
Result *ForwardResult
|
Result *ForwardResult
|
||||||
ParsedRequest *ParsedRequest
|
|
||||||
APIKey *APIKey
|
APIKey *APIKey
|
||||||
User *User
|
User *User
|
||||||
Account *Account
|
Account *Account
|
||||||
@ -8709,15 +8712,8 @@ func writeUsageLogBestEffort(ctx context.Context, repo UsageLogRepository, usage
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// recordUsageOpts 内部选项,参数化 RecordUsage 与 RecordUsageWithLongContext 的差异点。
|
// recordUsageOpts 内部选项,参数化普通计费与长上下文计费的差异点。
|
||||||
type recordUsageOpts struct {
|
type recordUsageOpts struct {
|
||||||
// Claude Max 策略所需的 ParsedRequest(可选,仅 Claude 路径传入)
|
|
||||||
ParsedRequest *ParsedRequest
|
|
||||||
|
|
||||||
// EnableClaudePath 启用 Claude 路径特有逻辑:
|
|
||||||
// - Claude Max 缓存计费策略
|
|
||||||
EnableClaudePath bool
|
|
||||||
|
|
||||||
// 长上下文计费(仅 Gemini 路径需要)
|
// 长上下文计费(仅 Gemini 路径需要)
|
||||||
LongContextThreshold int
|
LongContextThreshold int
|
||||||
LongContextMultiplier float64
|
LongContextMultiplier float64
|
||||||
@ -8740,9 +8736,7 @@ func (s *GatewayService) RecordUsage(ctx context.Context, input *RecordUsageInpu
|
|||||||
APIKeyService: input.APIKeyService,
|
APIKeyService: input.APIKeyService,
|
||||||
QuotaPlatform: input.QuotaPlatform,
|
QuotaPlatform: input.QuotaPlatform,
|
||||||
ChannelUsageFields: input.ChannelUsageFields,
|
ChannelUsageFields: input.ChannelUsageFields,
|
||||||
}, &recordUsageOpts{
|
}, &recordUsageOpts{})
|
||||||
EnableClaudePath: true,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// RecordUsageLongContextInput 记录使用量的输入参数(支持长上下文双倍计费)
|
// RecordUsageLongContextInput 记录使用量的输入参数(支持长上下文双倍计费)
|
||||||
@ -8808,9 +8802,7 @@ type recordUsageCoreInput struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// recordUsageCore 是 RecordUsage 和 RecordUsageWithLongContext 的统一实现。
|
// recordUsageCore 是 RecordUsage 和 RecordUsageWithLongContext 的统一实现。
|
||||||
// opts 中的字段控制两者之间的差异行为:
|
// LongContextThreshold > 0 时 Token 计费回退走 CalculateCostWithLongContext。
|
||||||
// - ParsedRequest != nil → 启用 Claude Max 缓存计费策略
|
|
||||||
// - LongContextThreshold > 0 → Token 计费回退走 CalculateCostWithLongContext
|
|
||||||
func (s *GatewayService) recordUsageCore(ctx context.Context, input *recordUsageCoreInput, opts *recordUsageOpts) error {
|
func (s *GatewayService) recordUsageCore(ctx context.Context, input *recordUsageCoreInput, opts *recordUsageOpts) error {
|
||||||
result := input.Result
|
result := input.Result
|
||||||
apiKey := input.APIKey
|
apiKey := input.APIKey
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user