mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 16:18:29 +08:00
Merge pull request #4787 from KtzeAbyss/fix/4760-ws-turn-model-billing
fix(openai): track WebSocket models per turn
This commit is contained in:
@@ -9,6 +9,7 @@ import (
|
||||
"runtime/debug"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
@@ -43,6 +44,49 @@ type OpenAIGatewayHandler struct {
|
||||
cfg *config.Config
|
||||
}
|
||||
|
||||
type openAIWSTurnChannelMappingSnapshot struct {
|
||||
turn int
|
||||
mapping service.ChannelMappingResult
|
||||
}
|
||||
|
||||
var errOpenAIWSUnsupportedModelSwitch = errors.New("selected account does not support websocket model switch")
|
||||
|
||||
func newOpenAIWSUnsupportedModelSwitchError(model string) error {
|
||||
cause := fmt.Errorf("%w: model %q", errOpenAIWSUnsupportedModelSwitch, strings.TrimSpace(model))
|
||||
return service.NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "model switch requires reconnect", cause)
|
||||
}
|
||||
|
||||
func shouldReportOpenAIWSProxyAccountFailure(err error) bool {
|
||||
return err != nil && !errors.Is(err, errOpenAIWSUnsupportedModelSwitch)
|
||||
}
|
||||
|
||||
func openAIWSTurnBillingModel(result *service.OpenAIForwardResult, mapping service.ChannelMappingResult, requestedModel, upstreamModel string) string {
|
||||
billingModel := ""
|
||||
if result != nil {
|
||||
billingModel = strings.TrimSpace(result.BillingModel)
|
||||
}
|
||||
if billingModel == "" {
|
||||
billingModel = strings.TrimSpace(upstreamModel)
|
||||
}
|
||||
if billingModel == "" {
|
||||
billingModel = strings.TrimSpace(requestedModel)
|
||||
}
|
||||
|
||||
requestedModel = strings.TrimSpace(requestedModel)
|
||||
switch mapping.BillingModelSource {
|
||||
case service.BillingModelSourceRequested:
|
||||
if requestedModel != "" {
|
||||
billingModel = requestedModel
|
||||
}
|
||||
case service.BillingModelSourceChannelMapped:
|
||||
mappedModel := strings.TrimSpace(mapping.MappedModel)
|
||||
if mappedModel != "" && mappedModel != requestedModel {
|
||||
billingModel = mappedModel
|
||||
}
|
||||
}
|
||||
return billingModel
|
||||
}
|
||||
|
||||
type grokMediaEligibilityProber interface {
|
||||
ProbeMediaEligibility(ctx context.Context, accountID int64) (bool, string, error)
|
||||
}
|
||||
@@ -1519,7 +1563,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
previousResponseCanMove := !firstMessageToolCoverage.HasFunctionCallOutput || firstMessageToolCoverage.ContextCoversAllCallIDs
|
||||
reqLog = reqLog.With(
|
||||
zap.Bool("ws_ingress", true),
|
||||
zap.String("model", reqModel),
|
||||
zap.String("session_initial_model", reqModel),
|
||||
zap.Bool("has_previous_response_id", previousResponseID != ""),
|
||||
zap.String("previous_response_id_kind", previousResponseIDKind),
|
||||
)
|
||||
@@ -1771,6 +1815,10 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
reasoningEffortMappings = apiKey.Group.ReasoningEffortMappings
|
||||
}
|
||||
var requestPayloadHash string
|
||||
// Passthrough rejects overlapping response.create frames, so one immutable
|
||||
// turn-tagged slot preserves the exact mapping used for the in-flight request.
|
||||
var turnChannelMapping atomic.Pointer[openAIWSTurnChannelMappingSnapshot]
|
||||
turnChannelMapping.Store(&openAIWSTurnChannelMappingSnapshot{turn: 1, mapping: channelMappingWS})
|
||||
hooks := &service.OpenAIWSIngressHooks{
|
||||
InitialRequestModel: reqModel,
|
||||
MaxReasoningEffort: maxReasoningEffort,
|
||||
@@ -1795,6 +1843,22 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
}
|
||||
return nil
|
||||
},
|
||||
MapRequestModel: func(turn int, originalModel string) (string, error) {
|
||||
model := strings.TrimSpace(originalModel)
|
||||
if model == "" {
|
||||
model = reqModel
|
||||
}
|
||||
mapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(ctx, apiKey.GroupID, model)
|
||||
mappedModelUnchanged := false
|
||||
if previous := turnChannelMapping.Load(); previous != nil && previous.turn < turn {
|
||||
mappedModelUnchanged = strings.TrimSpace(previous.mapping.MappedModel) == strings.TrimSpace(mapping.MappedModel)
|
||||
}
|
||||
if turn > 1 && !mappedModelUnchanged && !account.IsModelSupported(model) && !account.IsModelSupported(mapping.MappedModel) {
|
||||
return "", newOpenAIWSUnsupportedModelSwitchError(mapping.MappedModel)
|
||||
}
|
||||
turnChannelMapping.Store(&openAIWSTurnChannelMappingSnapshot{turn: turn, mapping: mapping})
|
||||
return mapping.MappedModel, nil
|
||||
},
|
||||
BeforeTurn: func(turn int) error {
|
||||
// turn==1 的会话屏蔽已由握手层检查覆盖;连接内 flag 只拦截后续 turn。
|
||||
if cyberBlockedThisConn {
|
||||
@@ -1836,7 +1900,27 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
// 届时 defer 已清除标记)。
|
||||
defer clearCyberPolicyTurnState(c)
|
||||
releaseTurnSlots()
|
||||
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, turnErr != nil, cyberBlockKey, clientRequestedUsageFields(c, channelMappingWS, reqModel, ""), requestPayloadHash)
|
||||
turnRequestedModel := reqModel
|
||||
turnUpstreamModel := ""
|
||||
if result != nil && turn > 1 {
|
||||
if model := strings.TrimSpace(result.Model); model != "" {
|
||||
turnRequestedModel = model
|
||||
}
|
||||
}
|
||||
if result != nil {
|
||||
turnUpstreamModel = strings.TrimSpace(result.UpstreamModel)
|
||||
}
|
||||
var turnMapping service.ChannelMappingResult
|
||||
if snapshot := turnChannelMapping.Load(); snapshot != nil && snapshot.turn == turn {
|
||||
turnMapping = snapshot.mapping
|
||||
} else {
|
||||
turnMapping, _ = h.gatewayService.ResolveChannelMappingAndRestrict(ctx, apiKey.GroupID, turnRequestedModel)
|
||||
}
|
||||
if turnUpstreamModel == "" {
|
||||
turnUpstreamModel = turnRequestedModel
|
||||
}
|
||||
turnUsageFields := turnMapping.ToUsageFields(turnRequestedModel, turnUpstreamModel)
|
||||
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, turnRequestedModel, turnErr != nil, cyberBlockKey, turnUsageFields, requestPayloadHash)
|
||||
if service.GetOpsCyberPolicy(c) != nil {
|
||||
cyberBlockedThisConn = true
|
||||
}
|
||||
@@ -1858,11 +1942,22 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
if result == nil {
|
||||
return
|
||||
}
|
||||
result.BillingModel = openAIWSTurnBillingModel(result, turnMapping, turnRequestedModel, turnUpstreamModel)
|
||||
reqLog.Debug("openai.websocket_turn_billing",
|
||||
zap.Int("turn", turn),
|
||||
zap.String("turn_requested_model", turnRequestedModel),
|
||||
zap.String("turn_upstream_model", turnUpstreamModel),
|
||||
zap.String("billing_model", result.BillingModel),
|
||||
)
|
||||
// 排除 spark 影子:其 codex_* 仅由 QueryUsage(/wham/usage bengalfox)更新(外审第7轮 P1)。
|
||||
if account.Type == service.AccountTypeOAuth && !account.IsShadow() {
|
||||
h.gatewayService.UpdateCodexUsageSnapshotFromHeaders(ctx, account.ID, result.ResponseHeaders)
|
||||
}
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), openAIForwardSucceededForScheduling(result), result.FirstTokenMs)
|
||||
scheduleModel := turnUpstreamModel
|
||||
if scheduleModel == "" {
|
||||
scheduleModel = turnRequestedModel
|
||||
}
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, scheduleModel, openAIForwardSucceededForScheduling(result), result.FirstTokenMs)
|
||||
inboundEndpoint := GetInboundEndpoint(c)
|
||||
upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result)
|
||||
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
|
||||
@@ -1883,7 +1978,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
APIKeyService: h.apiKeyService,
|
||||
QuotaPlatform: quotaPlatform,
|
||||
SessionID: sessionID,
|
||||
ChannelUsageFields: clientRequestedUsageFields(c, channelMappingWS, reqModel, result.UpstreamModel),
|
||||
ChannelUsageFields: turnUsageFields,
|
||||
CyberBlocked: cyberBlocked,
|
||||
}); err != nil {
|
||||
reqLog.Error("openai.websocket_record_usage_failed",
|
||||
@@ -1896,11 +1991,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
},
|
||||
}
|
||||
|
||||
// 应用渠道模型映射到 WebSocket 首条消息
|
||||
wsFirstMessage := firstMessage
|
||||
if channelMappingWS.Mapped {
|
||||
wsFirstMessage = h.gatewayService.ReplaceModelInBody(firstMessage, channelMappingWS.MappedModel)
|
||||
}
|
||||
// 切组/会话失配防护:previous_response_id 未在当前分组命中粘连账号(StickyPreviousHit=false),
|
||||
// 说明该会话链不属于本次调度到的账号,原样转发会触发上游会话链鉴权失败(“鉴权失败,请检查 API Key”)。
|
||||
// 故剥离首包里的 previous_response_id,改用首包内 input 重建上下文;带 function_call_output 的
|
||||
@@ -1944,7 +2035,9 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
|
||||
if shouldReportOpenAIWSProxyAccountFailure(err) {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
|
||||
}
|
||||
closeStatus, closeReason := summarizeWSCloseErrorForLog(err)
|
||||
proxyFailedFields := []zap.Field{
|
||||
zap.Int64("account_id", account.ID),
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -1369,6 +1370,229 @@ func TestOpenAIResponsesWebSocket_PassthroughUsageLogLeavesUserAgentNilWhenMissi
|
||||
require.Equal(t, "medium", *got.log.ReasoningEffort)
|
||||
}
|
||||
|
||||
func TestOpenAIResponsesWebSocket_PassthroughTracksModelPerTurn(t *testing.T) {
|
||||
got := runOpenAIResponsesWebSocketUsageLogCase(t, openAIResponsesWSUsageLogCase{
|
||||
firstPayload: `{"type":"response.create","model":"sol","stream":false}`,
|
||||
secondPayload: `{"type":"response.create","model":"terra","stream":false}`,
|
||||
channelMapping: map[string]string{
|
||||
"sol": "sol-channel",
|
||||
"terra": "terra-channel",
|
||||
},
|
||||
accountModelMapping: map[string]any{
|
||||
"sol": "gpt-5.6-sol",
|
||||
"terra": "gpt-5.6-terra",
|
||||
"sol-channel": "gpt-5.6-sol",
|
||||
"terra-channel": "gpt-5.6-terra",
|
||||
},
|
||||
})
|
||||
|
||||
require.Len(t, got.upstreamPayloads, 2)
|
||||
require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(got.upstreamPayloads[0], "model").String())
|
||||
require.Equal(t, "gpt-5.6-terra", gjson.GetBytes(got.upstreamPayloads[1], "model").String())
|
||||
require.Len(t, got.clientEvents, 2)
|
||||
require.Equal(t, "sol", gjson.GetBytes(got.clientEvents[0], "response.model").String())
|
||||
require.Equal(t, "terra", gjson.GetBytes(got.clientEvents[1], "response.model").String())
|
||||
|
||||
require.Len(t, got.logs, 2)
|
||||
require.Equal(t, "sol", got.logs[0].Model)
|
||||
require.Equal(t, "sol", got.logs[0].RequestedModel)
|
||||
require.NotNil(t, got.logs[0].UpstreamModel)
|
||||
require.Equal(t, "gpt-5.6-sol", *got.logs[0].UpstreamModel)
|
||||
require.NotNil(t, got.logs[0].ModelMappingChain)
|
||||
require.Equal(t, "sol→sol-channel→gpt-5.6-sol", *got.logs[0].ModelMappingChain)
|
||||
|
||||
require.Equal(t, "terra", got.logs[1].Model)
|
||||
require.Equal(t, "terra", got.logs[1].RequestedModel)
|
||||
require.NotNil(t, got.logs[1].UpstreamModel)
|
||||
require.Equal(t, "gpt-5.6-terra", *got.logs[1].UpstreamModel)
|
||||
require.NotNil(t, got.logs[1].ModelMappingChain)
|
||||
require.Equal(t, "terra→terra-channel→gpt-5.6-terra", *got.logs[1].ModelMappingChain)
|
||||
require.InDelta(t, got.logs[1].TotalCost*2, got.logs[0].TotalCost, 1e-12,
|
||||
"each turn must be billed with its own channel-mapped model")
|
||||
}
|
||||
|
||||
func TestOpenAIResponsesWebSocket_UnchangedChannelTargetOutsideAccountMappingKeysRemainsValid(t *testing.T) {
|
||||
got := runOpenAIResponsesWebSocketUsageLogCase(t, openAIResponsesWSUsageLogCase{
|
||||
firstPayload: `{"type":"response.create","model":"public-alias","stream":false}`,
|
||||
secondPayload: `{"type":"response.create","stream":false}`,
|
||||
channelMapping: map[string]string{
|
||||
"public-alias": "gpt-5.6-sol",
|
||||
},
|
||||
accountModelMapping: map[string]any{
|
||||
"public-alias": "gpt-5.6-terra",
|
||||
},
|
||||
})
|
||||
|
||||
require.Len(t, got.upstreamPayloads, 2)
|
||||
require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(got.upstreamPayloads[0], "model").String())
|
||||
require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(got.upstreamPayloads[1], "model").String())
|
||||
require.Len(t, got.clientEvents, 2)
|
||||
require.Equal(t, "public-alias", gjson.GetBytes(got.clientEvents[0], "response.model").String())
|
||||
require.Equal(t, "public-alias", gjson.GetBytes(got.clientEvents[1], "response.model").String())
|
||||
require.Len(t, got.logs, 2)
|
||||
for _, usageLog := range got.logs {
|
||||
require.Equal(t, "public-alias", usageLog.RequestedModel)
|
||||
require.NotNil(t, usageLog.UpstreamModel)
|
||||
require.Equal(t, "gpt-5.6-sol", *usageLog.UpstreamModel)
|
||||
require.NotNil(t, usageLog.ModelMappingChain)
|
||||
require.Equal(t, "public-alias→gpt-5.6-sol", *usageLog.ModelMappingChain)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIResponsesWebSocket_PassthroughKeepsTurnMappingSnapshot(t *testing.T) {
|
||||
got := runOpenAIResponsesWebSocketUsageLogCase(t, openAIResponsesWSUsageLogCase{
|
||||
firstPayload: `{"type":"response.create","model":"sol","stream":false}`,
|
||||
secondPayload: `{"type":"response.create","model":"sol","stream":false}`,
|
||||
channelMapping: map[string]string{
|
||||
"sol": "gpt-5.6-sol",
|
||||
},
|
||||
afterFirstUpstreamRequest: func(channelSvc *service.ChannelService) error {
|
||||
if channelSvc == nil {
|
||||
return errors.New("channel service is nil")
|
||||
}
|
||||
_, err := channelSvc.Update(context.Background(), 7701, &service.UpdateChannelInput{
|
||||
ModelMapping: map[string]map[string]string{
|
||||
service.PlatformOpenAI: {"sol": "gpt-5.6-terra"},
|
||||
},
|
||||
})
|
||||
return err
|
||||
},
|
||||
})
|
||||
|
||||
require.Len(t, got.upstreamPayloads, 2)
|
||||
require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(got.upstreamPayloads[0], "model").String())
|
||||
require.Equal(t, "gpt-5.6-terra", gjson.GetBytes(got.upstreamPayloads[1], "model").String())
|
||||
|
||||
require.Len(t, got.logs, 2)
|
||||
require.Equal(t, "sol", got.logs[0].Model)
|
||||
require.NotNil(t, got.logs[0].UpstreamModel)
|
||||
require.Equal(t, "gpt-5.6-sol", *got.logs[0].UpstreamModel)
|
||||
require.NotNil(t, got.logs[0].ModelMappingChain)
|
||||
require.Equal(t, "sol→gpt-5.6-sol", *got.logs[0].ModelMappingChain)
|
||||
require.InDelta(t, 40e-6, got.logs[0].TotalCost, 1e-12,
|
||||
"the in-flight turn must retain the channel-mapped billing model used when it was sent")
|
||||
|
||||
require.Equal(t, "sol", got.logs[1].Model)
|
||||
require.NotNil(t, got.logs[1].UpstreamModel)
|
||||
require.Equal(t, "gpt-5.6-terra", *got.logs[1].UpstreamModel)
|
||||
require.NotNil(t, got.logs[1].ModelMappingChain)
|
||||
require.Equal(t, "sol→gpt-5.6-terra", *got.logs[1].ModelMappingChain)
|
||||
require.InDelta(t, got.logs[1].TotalCost*2, got.logs[0].TotalCost, 1e-12,
|
||||
"the next turn must use the updated channel mapping")
|
||||
}
|
||||
|
||||
func TestOpenAIResponsesWebSocket_CtxPoolAppliesPerTurnMappingAndPreservesRequestedModel(t *testing.T) {
|
||||
got := runOpenAIResponsesWebSocketUsageLogCase(t, openAIResponsesWSUsageLogCase{
|
||||
firstPayload: `{"type":"response.create","model":"gpt-5.6-sol","stream":false}`,
|
||||
secondPayload: `{"type":"response.create","model":"gpt-5.6-terra","stream":false}`,
|
||||
ingressMode: service.OpenAIWSIngressModeCtxPool,
|
||||
billingModelSource: service.BillingModelSourceRequested,
|
||||
channelMapping: map[string]string{
|
||||
"gpt-5.6-terra": "gpt-5.6-sol",
|
||||
},
|
||||
})
|
||||
|
||||
require.Len(t, got.upstreamPayloads, 2)
|
||||
require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(got.upstreamPayloads[0], "model").String())
|
||||
require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(got.upstreamPayloads[1], "model").String())
|
||||
require.Len(t, got.clientEvents, 2)
|
||||
require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(got.clientEvents[0], "response.model").String())
|
||||
require.Equal(t, "gpt-5.6-terra", gjson.GetBytes(got.clientEvents[1], "response.model").String())
|
||||
|
||||
require.Len(t, got.logs, 2)
|
||||
require.Equal(t, "gpt-5.6-sol", got.logs[0].RequestedModel)
|
||||
require.Nil(t, got.logs[0].ModelMappingChain)
|
||||
require.InDelta(t, 40e-6, got.logs[0].TotalCost, 1e-12)
|
||||
require.Equal(t, "gpt-5.6-terra", got.logs[1].RequestedModel)
|
||||
require.NotNil(t, got.logs[1].ModelMappingChain)
|
||||
require.Equal(t, "gpt-5.6-terra→gpt-5.6-sol", *got.logs[1].ModelMappingChain)
|
||||
require.InDelta(t, 20e-6, got.logs[1].TotalCost, 1e-12,
|
||||
"BillingModelSourceRequested must use the client model before channel mapping")
|
||||
}
|
||||
|
||||
func TestOpenAIWSTurnBillingModelPreservesImagePricingModel(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
resultModel string
|
||||
mapping service.ChannelMappingResult
|
||||
requestedModel string
|
||||
upstreamModel string
|
||||
wantBillingModel string
|
||||
}{
|
||||
{
|
||||
name: "upstream billing preserves image model",
|
||||
resultModel: "gpt-image-2",
|
||||
mapping: service.ChannelMappingResult{BillingModelSource: service.BillingModelSourceUpstream},
|
||||
requestedModel: "gpt-5.6-sol",
|
||||
upstreamModel: "gpt-5.6-sol",
|
||||
wantBillingModel: "gpt-image-2",
|
||||
},
|
||||
{
|
||||
name: "unmapped channel preserves image model",
|
||||
resultModel: "gpt-image-2",
|
||||
mapping: service.ChannelMappingResult{MappedModel: "gpt-5.6-sol", BillingModelSource: service.BillingModelSourceChannelMapped},
|
||||
requestedModel: "gpt-5.6-sol",
|
||||
upstreamModel: "gpt-5.6-sol",
|
||||
wantBillingModel: "gpt-image-2",
|
||||
},
|
||||
{
|
||||
name: "requested source overrides image model",
|
||||
resultModel: "gpt-image-2",
|
||||
mapping: service.ChannelMappingResult{BillingModelSource: service.BillingModelSourceRequested},
|
||||
requestedModel: "public-image-alias",
|
||||
upstreamModel: "gpt-5.6-sol",
|
||||
wantBillingModel: "public-image-alias",
|
||||
},
|
||||
{
|
||||
name: "mapped channel source overrides image model",
|
||||
resultModel: "gpt-image-2",
|
||||
mapping: service.ChannelMappingResult{MappedModel: "priced-channel-model", BillingModelSource: service.BillingModelSourceChannelMapped},
|
||||
requestedModel: "public-image-alias",
|
||||
upstreamModel: "gpt-5.6-sol",
|
||||
wantBillingModel: "priced-channel-model",
|
||||
},
|
||||
{
|
||||
name: "text turn falls back to upstream model",
|
||||
mapping: service.ChannelMappingResult{BillingModelSource: service.BillingModelSourceUpstream},
|
||||
requestedModel: "public-alias",
|
||||
upstreamModel: "gpt-5.6-sol",
|
||||
wantBillingModel: "gpt-5.6-sol",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := &service.OpenAIForwardResult{BillingModel: tt.resultModel}
|
||||
require.Equal(t, tt.wantBillingModel, openAIWSTurnBillingModel(result, tt.mapping, tt.requestedModel, tt.upstreamModel))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestShouldReportOpenAIWSProxyAccountFailure(t *testing.T) {
|
||||
t.Run("unsupported client model switch does not penalize account", func(t *testing.T) {
|
||||
err := fmt.Errorf("wrapped ingress turn: %w", newOpenAIWSUnsupportedModelSwitchError("gpt-unsupported"))
|
||||
require.False(t, shouldReportOpenAIWSProxyAccountFailure(err))
|
||||
|
||||
var closeErr *service.OpenAIWSClientCloseError
|
||||
require.ErrorAs(t, err, &closeErr)
|
||||
require.Equal(t, coderws.StatusPolicyViolation, closeErr.StatusCode())
|
||||
require.Equal(t, "model switch requires reconnect", closeErr.Reason())
|
||||
})
|
||||
|
||||
t.Run("upstream policy violation still penalizes account", func(t *testing.T) {
|
||||
err := service.NewOpenAIWSClientCloseError(
|
||||
coderws.StatusPolicyViolation,
|
||||
"upstream websocket authentication failed",
|
||||
errors.New("upstream rejected credentials"),
|
||||
)
|
||||
require.True(t, shouldReportOpenAIWSProxyAccountFailure(err))
|
||||
})
|
||||
|
||||
t.Run("generic proxy failure still penalizes account", func(t *testing.T) {
|
||||
require.True(t, shouldReportOpenAIWSProxyAccountFailure(errors.New("upstream websocket read failed")))
|
||||
})
|
||||
}
|
||||
|
||||
func TestSetOpenAIClientTransportHTTP(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
@@ -1516,14 +1740,22 @@ func newOpenAIWSHandlerTestServer(t *testing.T, h *OpenAIGatewayHandler, subject
|
||||
}
|
||||
|
||||
type openAIResponsesWSUsageLogCase struct {
|
||||
firstPayload string
|
||||
userAgent *string
|
||||
channelMapping map[string]string
|
||||
firstPayload string
|
||||
secondPayload string
|
||||
userAgent *string
|
||||
ingressMode string
|
||||
channelMapping map[string]string
|
||||
billingModelSource string
|
||||
accountModelMapping map[string]any
|
||||
afterFirstUpstreamRequest func(channelSvc *service.ChannelService) error
|
||||
}
|
||||
|
||||
type openAIResponsesWSUsageLogResult struct {
|
||||
log *service.UsageLog
|
||||
logs []*service.UsageLog
|
||||
upstreamFirstPayload []byte
|
||||
upstreamPayloads [][]byte
|
||||
clientEvents [][]byte
|
||||
}
|
||||
|
||||
type openAIWSUsageHandlerAccountRepoStub struct {
|
||||
@@ -1633,12 +1865,50 @@ func (s *openAIWSUsageHandlerUsageLogRepoStub) Create(ctx context.Context, log *
|
||||
|
||||
type openAIWSUsageHandlerChannelRepoStub struct {
|
||||
service.ChannelRepository
|
||||
mu sync.Mutex
|
||||
channels []service.Channel
|
||||
groupPlatforms map[int64]string
|
||||
}
|
||||
|
||||
func (s *openAIWSUsageHandlerChannelRepoStub) ListAll(ctx context.Context) ([]service.Channel, error) {
|
||||
return s.channels, nil
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
out := make([]service.Channel, 0, len(s.channels))
|
||||
for i := range s.channels {
|
||||
out = append(out, *s.channels[i].Clone())
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (s *openAIWSUsageHandlerChannelRepoStub) GetByID(ctx context.Context, id int64) (*service.Channel, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
for i := range s.channels {
|
||||
if s.channels[i].ID == id {
|
||||
return s.channels[i].Clone(), nil
|
||||
}
|
||||
}
|
||||
return nil, service.ErrChannelNotFound
|
||||
}
|
||||
|
||||
func (s *openAIWSUsageHandlerChannelRepoStub) Update(ctx context.Context, channel *service.Channel) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
for i := range s.channels {
|
||||
if s.channels[i].ID == channel.ID {
|
||||
s.channels[i] = *channel.Clone()
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return service.ErrChannelNotFound
|
||||
}
|
||||
|
||||
func (s *openAIWSUsageHandlerChannelRepoStub) GetGroupIDs(ctx context.Context, channelID int64) ([]int64, error) {
|
||||
channel, err := s.GetByID(ctx, channelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return append([]int64(nil), channel.GroupIDs...), nil
|
||||
}
|
||||
|
||||
func (s *openAIWSUsageHandlerChannelRepoStub) GetGroupPlatforms(ctx context.Context, groupIDs []int64) (map[int64]string, error) {
|
||||
@@ -2147,8 +2417,13 @@ func runOpenAIResponsesWebSocketUsageLogCase(t *testing.T, tc openAIResponsesWSU
|
||||
t.Helper()
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
upstreamPayloadCh := make(chan []byte, 1)
|
||||
turnCount := 1
|
||||
if strings.TrimSpace(tc.secondPayload) != "" {
|
||||
turnCount = 2
|
||||
}
|
||||
upstreamPayloadCh := make(chan []byte, turnCount)
|
||||
upstreamErrCh := make(chan error, 1)
|
||||
var channelSvc *service.ChannelService
|
||||
upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{
|
||||
CompressionMode: coderws.CompressionContextTakeover,
|
||||
@@ -2161,29 +2436,39 @@ func runOpenAIResponsesWebSocketUsageLogCase(t *testing.T, tc openAIResponsesWSU
|
||||
_ = conn.CloseNow()
|
||||
}()
|
||||
|
||||
readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
|
||||
msgType, payload, readErr := conn.Read(readCtx)
|
||||
cancelRead()
|
||||
if readErr != nil {
|
||||
upstreamErrCh <- readErr
|
||||
return
|
||||
}
|
||||
if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
|
||||
upstreamErrCh <- errors.New("unexpected upstream websocket message type")
|
||||
return
|
||||
}
|
||||
upstreamPayloadCh <- payload
|
||||
for turn := 1; turn <= turnCount; turn++ {
|
||||
readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
|
||||
msgType, payload, readErr := conn.Read(readCtx)
|
||||
cancelRead()
|
||||
if readErr != nil {
|
||||
upstreamErrCh <- readErr
|
||||
return
|
||||
}
|
||||
if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
|
||||
upstreamErrCh <- errors.New("unexpected upstream websocket message type")
|
||||
return
|
||||
}
|
||||
upstreamPayloadCh <- payload
|
||||
if turn == 1 && tc.afterFirstUpstreamRequest != nil {
|
||||
if callbackErr := tc.afterFirstUpstreamRequest(channelSvc); callbackErr != nil {
|
||||
upstreamErrCh <- callbackErr
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
writeCtx, cancelWrite := context.WithTimeout(r.Context(), 3*time.Second)
|
||||
writeErr := conn.Write(writeCtx, coderws.MessageText, []byte(
|
||||
`{"type":"response.completed","response":{"id":"resp_usage_e2e","model":"gpt-5.4","usage":{"input_tokens":2,"output_tokens":1}}}`,
|
||||
))
|
||||
cancelWrite()
|
||||
if writeErr != nil {
|
||||
upstreamErrCh <- writeErr
|
||||
return
|
||||
response := fmt.Sprintf(
|
||||
`{"type":"response.completed","response":{"id":"resp_usage_e2e_%d","model":%q,"usage":{"input_tokens":2,"output_tokens":1}}}`,
|
||||
turn,
|
||||
gjson.GetBytes(payload, "model").String(),
|
||||
)
|
||||
writeCtx, cancelWrite := context.WithTimeout(r.Context(), 3*time.Second)
|
||||
writeErr := conn.Write(writeCtx, coderws.MessageText, []byte(response))
|
||||
cancelWrite()
|
||||
if writeErr != nil {
|
||||
upstreamErrCh <- writeErr
|
||||
return
|
||||
}
|
||||
}
|
||||
_ = conn.Close(coderws.StatusNormalClosure, "done")
|
||||
upstreamErrCh <- nil
|
||||
}))
|
||||
defer upstreamServer.Close()
|
||||
@@ -2198,14 +2483,18 @@ func runOpenAIResponsesWebSocketUsageLogCase(t *testing.T, tc openAIResponsesWSU
|
||||
Schedulable: true,
|
||||
Concurrency: 1,
|
||||
Credentials: map[string]any{
|
||||
"api_key": "sk-test",
|
||||
"base_url": upstreamServer.URL,
|
||||
"api_key": "sk-test",
|
||||
"base_url": upstreamServer.URL,
|
||||
"model_mapping": tc.accountModelMapping,
|
||||
},
|
||||
Extra: map[string]any{
|
||||
"openai_apikey_responses_websockets_v2_enabled": true,
|
||||
"openai_apikey_responses_websockets_v2_mode": service.OpenAIWSIngressModePassthrough,
|
||||
},
|
||||
}
|
||||
if strings.TrimSpace(tc.ingressMode) != "" {
|
||||
account.Extra["openai_apikey_responses_websockets_v2_mode"] = tc.ingressMode
|
||||
}
|
||||
|
||||
cfg := &config.Config{}
|
||||
cfg.RunMode = config.RunModeSimple
|
||||
@@ -2221,17 +2510,17 @@ func runOpenAIResponsesWebSocketUsageLogCase(t *testing.T, tc openAIResponsesWSU
|
||||
cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
|
||||
|
||||
accountRepo := &openAIWSUsageHandlerAccountRepoStub{account: account}
|
||||
usageRepo := &openAIWSUsageHandlerUsageLogRepoStub{created: make(chan *service.UsageLog, 1)}
|
||||
usageRepo := &openAIWSUsageHandlerUsageLogRepoStub{created: make(chan *service.UsageLog, turnCount)}
|
||||
|
||||
var channelSvc *service.ChannelService
|
||||
if len(tc.channelMapping) > 0 {
|
||||
channelSvc = service.NewChannelService(&openAIWSUsageHandlerChannelRepoStub{
|
||||
channels: []service.Channel{{
|
||||
ID: 7701,
|
||||
Name: "openai-ws-e2e-channel",
|
||||
Status: service.StatusActive,
|
||||
GroupIDs: []int64{groupID},
|
||||
ModelMapping: map[string]map[string]string{service.PlatformOpenAI: tc.channelMapping},
|
||||
ID: 7701,
|
||||
Name: "openai-ws-e2e-channel",
|
||||
Status: service.StatusActive,
|
||||
GroupIDs: []int64{groupID},
|
||||
ModelMapping: map[string]map[string]string{service.PlatformOpenAI: tc.channelMapping},
|
||||
BillingModelSource: tc.billingModelSource,
|
||||
}},
|
||||
groupPlatforms: map[int64]string{groupID: service.PlatformOpenAI},
|
||||
}, nil, nil, nil)
|
||||
@@ -2314,26 +2603,44 @@ func runOpenAIResponsesWebSocketUsageLogCase(t *testing.T, tc openAIResponsesWSU
|
||||
cancelWrite()
|
||||
require.NoError(t, err)
|
||||
|
||||
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
_, event, err := clientConn.Read(readCtx)
|
||||
cancelRead()
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String())
|
||||
clientEvents := make([][]byte, 0, turnCount)
|
||||
readCompleted := func() {
|
||||
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
_, event, readErr := clientConn.Read(readCtx)
|
||||
cancelRead()
|
||||
require.NoError(t, readErr)
|
||||
require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String())
|
||||
clientEvents = append(clientEvents, append([]byte(nil), event...))
|
||||
}
|
||||
readCompleted()
|
||||
if turnCount == 2 {
|
||||
writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second)
|
||||
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(tc.secondPayload))
|
||||
cancelWrite()
|
||||
require.NoError(t, err)
|
||||
readCompleted()
|
||||
}
|
||||
_ = clientConn.Close(coderws.StatusNormalClosure, "done")
|
||||
|
||||
var usageLog *service.UsageLog
|
||||
select {
|
||||
case usageLog = <-usageRepo.created:
|
||||
require.NotNil(t, usageLog)
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("等待 WebSocket usage log 写入超时")
|
||||
usageLogs := make([]*service.UsageLog, 0, turnCount)
|
||||
for len(usageLogs) < turnCount {
|
||||
select {
|
||||
case usageLog := <-usageRepo.created:
|
||||
require.NotNil(t, usageLog)
|
||||
usageLogs = append(usageLogs, usageLog)
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("等待 WebSocket usage log 写入超时")
|
||||
}
|
||||
}
|
||||
|
||||
var upstreamFirstPayload []byte
|
||||
select {
|
||||
case upstreamFirstPayload = <-upstreamPayloadCh:
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("等待上游 WebSocket 首帧超时")
|
||||
upstreamPayloads := make([][]byte, 0, turnCount)
|
||||
for len(upstreamPayloads) < turnCount {
|
||||
select {
|
||||
case payload := <-upstreamPayloadCh:
|
||||
upstreamPayloads = append(upstreamPayloads, payload)
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("等待上游 WebSocket 请求帧超时")
|
||||
}
|
||||
}
|
||||
|
||||
select {
|
||||
@@ -2344,8 +2651,11 @@ func runOpenAIResponsesWebSocketUsageLogCase(t *testing.T, tc openAIResponsesWSU
|
||||
}
|
||||
|
||||
return openAIResponsesWSUsageLogResult{
|
||||
log: usageLog,
|
||||
upstreamFirstPayload: upstreamFirstPayload,
|
||||
log: usageLogs[0],
|
||||
logs: usageLogs,
|
||||
upstreamFirstPayload: upstreamPayloads[0],
|
||||
upstreamPayloads: upstreamPayloads,
|
||||
clientEvents: clientEvents,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -207,8 +207,9 @@ func (e *OpenAIWSClientCloseError) Reason() string {
|
||||
|
||||
// OpenAIWSIngressHooks 定义入站 WS 每个 turn 的生命周期回调。
|
||||
type OpenAIWSIngressHooks struct {
|
||||
// InitialRequestModel 是首帧渠道映射前的请求模型,只用于 usage metadata
|
||||
// 的 reasoning effort 后缀推导,禁止用于上游请求或计费模型。
|
||||
// InitialRequestModel is the client-facing model from the first frame,
|
||||
// before channel or account mapping. Ingress modes preserve it for usage
|
||||
// attribution while MapRequestModel determines the upstream model.
|
||||
InitialRequestModel string
|
||||
// MaxReasoningEffort limits explicit reasoning effort values for this WS session.
|
||||
MaxReasoningEffort string
|
||||
@@ -216,7 +217,10 @@ type OpenAIWSIngressHooks struct {
|
||||
ReasoningEffortMappings []ReasoningEffortMapping
|
||||
BeforeTurn func(turn int) error
|
||||
BeforeRequest func(turn int, payload []byte, originalModel string) error
|
||||
AfterTurn func(turn int, result *OpenAIForwardResult, turnErr error)
|
||||
// MapRequestModel resolves the current turn's client model to the model
|
||||
// that must be written into the upstream response.create frame.
|
||||
MapRequestModel func(turn int, originalModel string) (string, error)
|
||||
AfterTurn func(turn int, result *OpenAIForwardResult, turnErr error)
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) getOpenAIWSConnPool() *openAIWSConnPool {
|
||||
|
||||
@@ -163,7 +163,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
return rebuilt, nil
|
||||
}
|
||||
|
||||
parseClientPayload := func(raw []byte) (openAIWSClientPayload, error) {
|
||||
parseClientPayload := func(turn int, raw []byte) (openAIWSClientPayload, error) {
|
||||
trimmed := bytes.TrimSpace(raw)
|
||||
if len(trimmed) == 0 {
|
||||
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "empty websocket request payload", nil)
|
||||
@@ -288,7 +288,17 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
normalized = rebuilt
|
||||
}
|
||||
}
|
||||
upstreamModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel))
|
||||
requestModel := originalModel
|
||||
if hooks != nil && hooks.MapRequestModel != nil {
|
||||
mappedModel, mapErr := hooks.MapRequestModel(turn, originalModel)
|
||||
if mapErr != nil {
|
||||
return openAIWSClientPayload{}, mapErr
|
||||
}
|
||||
if mappedModel = strings.TrimSpace(mappedModel); mappedModel != "" {
|
||||
requestModel = mappedModel
|
||||
}
|
||||
}
|
||||
upstreamModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(requestModel))
|
||||
if modelMissing || upstreamModel != originalModel {
|
||||
next, setErr := applyPayloadMutation(normalized, "model", upstreamModel)
|
||||
if setErr != nil {
|
||||
@@ -412,7 +422,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
firstPayload, err := parseClientPayload(firstClientMessage)
|
||||
firstPayload, err := parseClientPayload(1, firstClientMessage)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -513,7 +523,13 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
}
|
||||
grokCacheIdentity := ""
|
||||
if account.Platform == PlatformGrok {
|
||||
grokCacheIdentity, err = resolveGrokWSCacheIdentity(c, account, grokCacheSeedPayload, currentBridgePayload.originalModel)
|
||||
grokCacheIdentity, err = resolveGrokWSCacheIdentity(
|
||||
c,
|
||||
account,
|
||||
grokCacheSeedPayload,
|
||||
currentBridgePayload.payloadRaw,
|
||||
currentBridgePayload.originalModel,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("resolve Grok websocket cache identity: %w", err)
|
||||
}
|
||||
@@ -573,7 +589,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
}
|
||||
return fmt.Errorf("read client websocket request: %w", readErr)
|
||||
}
|
||||
nextPayload, parseErr := parseClientPayload(nextClientMessage)
|
||||
nextPayload, parseErr := parseClientPayload(turn+1, nextClientMessage)
|
||||
if parseErr != nil {
|
||||
return parseErr
|
||||
}
|
||||
@@ -792,7 +808,10 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
mappedModel := ""
|
||||
var mappedModelBytes []byte
|
||||
if originalModel != "" {
|
||||
mappedModel = normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel))
|
||||
mappedModel = strings.TrimSpace(gjson.GetBytes(payload, "model").String())
|
||||
if mappedModel == "" {
|
||||
mappedModel = normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel))
|
||||
}
|
||||
needModelReplace = mappedModel != "" && mappedModel != originalModel
|
||||
if needModelReplace {
|
||||
mappedModelBytes = []byte(mappedModel)
|
||||
@@ -1589,7 +1608,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
return fmt.Errorf("read client websocket request: %w", readErr)
|
||||
}
|
||||
|
||||
nextPayload, parseErr := parseClientPayload(nextClientMessage)
|
||||
nextPayload, parseErr := parseClientPayload(turn+1, nextClientMessage)
|
||||
if parseErr != nil {
|
||||
return parseErr
|
||||
}
|
||||
|
||||
@@ -286,7 +286,10 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
|
||||
needModelReplace := false
|
||||
var mappedModelBytes []byte
|
||||
if originalModel != "" {
|
||||
mappedModel = normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel))
|
||||
mappedModel = strings.TrimSpace(gjson.GetBytes(body, "model").String())
|
||||
if mappedModel == "" {
|
||||
mappedModel = normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel))
|
||||
}
|
||||
needModelReplace = mappedModel != "" && mappedModel != originalModel
|
||||
if needModelReplace {
|
||||
mappedModelBytes = []byte(mappedModel)
|
||||
@@ -487,12 +490,12 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
|
||||
return resultWithUsage(), terminalErr
|
||||
}
|
||||
|
||||
func resolveGrokWSCacheIdentity(c *gin.Context, account *Account, payload []byte, originalModel string) (string, error) {
|
||||
body, err := prepareOpenAIWSHTTPBridgeBody(payload)
|
||||
func resolveGrokWSCacheIdentity(c *gin.Context, account *Account, seedPayload, currentPayload []byte, originalModel string) (string, error) {
|
||||
body, err := prepareOpenAIWSHTTPBridgeBody(seedPayload)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
upstreamModel := resolveGrokWSUpstreamModel(account, body, originalModel)
|
||||
upstreamModel := resolveGrokWSUpstreamModel(account, currentPayload, originalModel)
|
||||
body, err = patchGrokResponsesBody(body, upstreamModel)
|
||||
if err != nil {
|
||||
return "", err
|
||||
@@ -502,7 +505,11 @@ func resolveGrokWSCacheIdentity(c *gin.Context, account *Account, payload []byte
|
||||
|
||||
func resolveGrokWSUpstreamModel(account *Account, body []byte, originalModel string) string {
|
||||
upstreamModel := strings.TrimSpace(gjson.GetBytes(body, "model").String())
|
||||
if account != nil && originalModel != "" {
|
||||
originalModel = strings.TrimSpace(originalModel)
|
||||
// Shared ingress has already applied channel and account mappings when the
|
||||
// body model differs from the client-facing model. Only resolve from the
|
||||
// original model when the body still carries that original value.
|
||||
if account != nil && originalModel != "" && (upstreamModel == "" || upstreamModel == originalModel) {
|
||||
if mappedModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel)); mappedModel != "" {
|
||||
upstreamModel = mappedModel
|
||||
}
|
||||
|
||||
@@ -520,7 +520,7 @@ func TestProxyOpenAIWSHTTPBridgeTurnPromotesCodexAdditionalToolsForMixedCache(t
|
||||
require.Empty(t, upstream.lastReq.Header.Get(grokClientToolCacheOptInHeader))
|
||||
}
|
||||
|
||||
func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T) {
|
||||
func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridgeAndPreservesMappedModels(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
bridgeResponse := func(responseID, requestID string, cachedTokens int) *http.Response {
|
||||
@@ -594,7 +594,14 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T)
|
||||
ginCtx.Request = req
|
||||
ginCtx.Set("api_key", &APIKey{ID: 7101})
|
||||
|
||||
errCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "access-token", firstMessage, nil)
|
||||
errCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "access-token", firstMessage, &OpenAIWSIngressHooks{
|
||||
MapRequestModel: func(_ int, originalModel string) (string, error) {
|
||||
if originalModel == "channel-alias" {
|
||||
return "grok-4.3", nil
|
||||
}
|
||||
return originalModel, nil
|
||||
},
|
||||
})
|
||||
}))
|
||||
defer wsServer.Close()
|
||||
|
||||
@@ -626,7 +633,7 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T)
|
||||
require.Equal(t, "resp_grok_ws_1", gjson.GetBytes(completed, "response.id").String())
|
||||
|
||||
writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second)
|
||||
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","generate":true,"model":"grok","stream":true,"previous_response_id":"resp_grok_ws_1","input":"second turn"}`))
|
||||
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","generate":true,"model":"channel-alias","stream":true,"previous_response_id":"resp_grok_ws_1","input":"second turn"}`))
|
||||
cancelWrite()
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -637,6 +644,7 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T)
|
||||
require.Equal(t, "response.output_text.delta", gjson.GetBytes(delta, "type").String())
|
||||
require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String())
|
||||
require.Equal(t, "resp_grok_ws_2", gjson.GetBytes(completed, "response.id").String())
|
||||
require.Equal(t, "channel-alias", gjson.GetBytes(completed, "response.model").String())
|
||||
|
||||
writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second)
|
||||
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","generate":true,"model":"grok-4.3","stream":true,"previous_response_id":"resp_grok_ws_2","input":"third turn with a different model"}`))
|
||||
@@ -666,7 +674,7 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T)
|
||||
require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent"))
|
||||
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
|
||||
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.bodies[0], "model").String())
|
||||
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.bodies[1], "model").String())
|
||||
require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.bodies[1], "model").String())
|
||||
require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.bodies[2], "model").String())
|
||||
require.NotEmpty(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String())
|
||||
require.Equal(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String(), upstream.lastReq.Header.Get(grokConversationIDHeader))
|
||||
@@ -677,9 +685,10 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T)
|
||||
secondIdentity := gjson.GetBytes(upstream.bodies[1], "prompt_cache_key").String()
|
||||
thirdIdentity := gjson.GetBytes(upstream.bodies[2], "prompt_cache_key").String()
|
||||
require.NotEmpty(t, firstIdentity)
|
||||
require.Equal(t, firstIdentity, secondIdentity)
|
||||
require.NotEmpty(t, secondIdentity)
|
||||
require.NotEqual(t, firstIdentity, secondIdentity)
|
||||
require.NotEmpty(t, thirdIdentity)
|
||||
require.NotEqual(t, firstIdentity, thirdIdentity)
|
||||
require.Equal(t, secondIdentity, thirdIdentity)
|
||||
require.Equal(t, firstIdentity, upstream.requests[0].Header.Get(grokConversationIDHeader))
|
||||
require.Equal(t, secondIdentity, upstream.requests[1].Header.Get(grokConversationIDHeader))
|
||||
require.Equal(t, thirdIdentity, upstream.requests[2].Header.Get(grokConversationIDHeader))
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"net"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
@@ -86,6 +87,7 @@ type RelayTraceEvent struct {
|
||||
|
||||
type relayState struct {
|
||||
usage Usage
|
||||
requestModelMu sync.RWMutex
|
||||
requestModel string
|
||||
lastResponseID string
|
||||
terminalEventType string
|
||||
@@ -164,6 +166,12 @@ func Relay(
|
||||
defer cancel()
|
||||
return upstreamConn.WriteFrame(writeCtx, msgType, payload)
|
||||
}
|
||||
writeClientFrameUpstream := func(msgType coderws.MessageType, payload []byte) error {
|
||||
if msgType == coderws.MessageText && strings.TrimSpace(gjson.GetBytes(payload, "type").String()) == "response.create" {
|
||||
state.setRequestModel(strings.TrimSpace(gjson.GetBytes(payload, "model").String()))
|
||||
}
|
||||
return writeUpstream(msgType, payload)
|
||||
}
|
||||
writeClient := func(msgType coderws.MessageType, payload []byte) error {
|
||||
writeCtx, cancel := context.WithTimeout(relayCtx, writeTimeout)
|
||||
defer cancel()
|
||||
@@ -215,7 +223,7 @@ func Relay(
|
||||
if !clientReaderStarted.CompareAndSwap(false, true) {
|
||||
return
|
||||
}
|
||||
go runClientToUpstream(relayCtx, clientConn, options.ReadClientFrame, writeUpstream, markActivity, clientToUpstreamFrames, onTrace, exitCh)
|
||||
go runClientToUpstream(relayCtx, clientConn, options.ReadClientFrame, writeClientFrameUpstream, markActivity, clientToUpstreamFrames, onTrace, exitCh)
|
||||
}
|
||||
if !options.StartClientAfterFirstDownstream {
|
||||
startClientReader()
|
||||
@@ -712,7 +720,7 @@ func emitTurnComplete(
|
||||
}
|
||||
requestModel := ""
|
||||
if state != nil {
|
||||
requestModel = state.requestModel
|
||||
requestModel = state.currentRequestModel()
|
||||
}
|
||||
onTurnComplete(RelayTurnResult{
|
||||
RequestModel: requestModel,
|
||||
@@ -873,13 +881,31 @@ func enrichResult(result *RelayResult, state *relayState, duration time.Duration
|
||||
if state == nil {
|
||||
return
|
||||
}
|
||||
result.RequestModel = state.requestModel
|
||||
result.RequestModel = state.currentRequestModel()
|
||||
result.Usage = state.usage
|
||||
result.RequestID = state.lastResponseID
|
||||
result.TerminalEventType = state.terminalEventType
|
||||
result.FirstTokenMs = state.firstTokenMs
|
||||
}
|
||||
|
||||
func (s *relayState) setRequestModel(model string) {
|
||||
if s == nil || model == "" {
|
||||
return
|
||||
}
|
||||
s.requestModelMu.Lock()
|
||||
s.requestModel = model
|
||||
s.requestModelMu.Unlock()
|
||||
}
|
||||
|
||||
func (s *relayState) currentRequestModel() string {
|
||||
if s == nil {
|
||||
return ""
|
||||
}
|
||||
s.requestModelMu.RLock()
|
||||
defer s.requestModelMu.RUnlock()
|
||||
return s.requestModel
|
||||
}
|
||||
|
||||
func isDisconnectError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
|
||||
@@ -472,6 +472,56 @@ func TestRelay_OnTurnComplete_PerTerminalEvent(t *testing.T) {
|
||||
require.Equal(t, 5, result.Usage.OutputTokens)
|
||||
}
|
||||
|
||||
func TestRelay_OnTurnComplete_UsesCurrentResponseCreateModel(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
clientConn := newPassthroughTestFrameConn(nil, false)
|
||||
upstreamConn := newPassthroughTestFrameConn(nil, false)
|
||||
firstPayload := []byte(`{"type":"response.create","model":"gpt-5.6-sol","input":[]}`)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
turns := make(chan RelayTurnResult, 2)
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
defer close(done)
|
||||
_, _ = Relay(ctx, clientConn, upstreamConn, firstPayload, RelayOptions{
|
||||
OnTurnComplete: func(turn RelayTurnResult) {
|
||||
turns <- turn
|
||||
},
|
||||
})
|
||||
}()
|
||||
|
||||
upstreamConn.readCh <- passthroughTestFrame{
|
||||
msgType: coderws.MessageText,
|
||||
payload: []byte(`{"type":"response.completed","response":{"id":"resp_sol","usage":{"input_tokens":2,"output_tokens":1}}}`),
|
||||
}
|
||||
firstTurn := <-turns
|
||||
require.Equal(t, "gpt-5.6-sol", firstTurn.RequestModel)
|
||||
|
||||
clientConn.readCh <- passthroughTestFrame{
|
||||
msgType: coderws.MessageText,
|
||||
payload: []byte(`{"type":"response.create","model":"gpt-5.6-terra","input":[]}`),
|
||||
}
|
||||
require.Eventually(t, func() bool {
|
||||
return len(upstreamConn.Writes()) == 2
|
||||
}, time.Second, 10*time.Millisecond)
|
||||
upstreamConn.readCh <- passthroughTestFrame{
|
||||
msgType: coderws.MessageText,
|
||||
payload: []byte(`{"type":"response.completed","response":{"id":"resp_terra","usage":{"input_tokens":3,"output_tokens":1}}}`),
|
||||
}
|
||||
secondTurn := <-turns
|
||||
require.Equal(t, "gpt-5.6-terra", secondTurn.RequestModel)
|
||||
|
||||
_ = clientConn.Close()
|
||||
_ = upstreamConn.Close()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("relay did not stop after client close")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRelay_OnTurnComplete_ProvidesTurnMetrics(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -25,6 +25,9 @@ type openAIWSClientFrameConn struct {
|
||||
interTurnIdleTimeout time.Duration
|
||||
interTurnStarted chan struct{}
|
||||
waitingForNextTurn atomic.Bool
|
||||
// The relay observes upstream payloads, while clients must keep seeing the
|
||||
// model identifier they supplied for the current turn.
|
||||
restoreResponseModel func([]byte) []byte
|
||||
}
|
||||
|
||||
// openAIWSPolicyEnforcingFrameConn wraps a client-side FrameConn and runs
|
||||
@@ -133,6 +136,8 @@ func openAIWSPassthroughPolicyModelFromSessionFrame(account *Account, payload []
|
||||
type openAIWSPassthroughUsageMeta struct {
|
||||
serviceTier atomic.Pointer[string]
|
||||
reasoningEffort atomic.Pointer[string]
|
||||
requestModel atomic.Pointer[string]
|
||||
upstreamModel atomic.Pointer[string]
|
||||
|
||||
// 仅在 client->upstream filter goroutine 中读写;Load 侧通过上方原子指针同步。
|
||||
sessionRequestModel string
|
||||
@@ -154,6 +159,7 @@ func (m *openAIWSPassthroughUsageMeta) initFromFirstFrame(policyOutput []byte, m
|
||||
}
|
||||
m.serviceTier.Store(extractOpenAIServiceTierFromBody(policyOutput))
|
||||
m.reasoningEffort.Store(extractOpenAIReasoningEffortFromBody(policyOutput, mappedModel, m.sessionRequestModel))
|
||||
m.storeTurnModels(m.sessionRequestModel, policyOutput)
|
||||
}
|
||||
|
||||
func (m *openAIWSPassthroughUsageMeta) updateSessionRequestModel(payload []byte) {
|
||||
@@ -181,6 +187,51 @@ func (m *openAIWSPassthroughUsageMeta) updateFromResponseCreate(policyOutput []b
|
||||
}
|
||||
m.serviceTier.Store(extractOpenAIServiceTierFromBody(policyOutput))
|
||||
m.reasoningEffort.Store(extractOpenAIReasoningEffortFromBody(policyOutput, mappedModel, requestModelForFrame))
|
||||
m.storeTurnModels(requestModelForFrame, policyOutput)
|
||||
}
|
||||
|
||||
func (m *openAIWSPassthroughUsageMeta) storeTurnModels(requestModel string, upstreamPayload []byte) {
|
||||
if m == nil {
|
||||
return
|
||||
}
|
||||
requestModel = strings.TrimSpace(requestModel)
|
||||
upstreamModel := strings.TrimSpace(gjson.GetBytes(upstreamPayload, "model").String())
|
||||
if upstreamModel == "" {
|
||||
upstreamModel = requestModel
|
||||
}
|
||||
m.requestModel.Store(openAIWSTrimmedStringPtr(requestModel))
|
||||
m.upstreamModel.Store(openAIWSTrimmedStringPtr(upstreamModel))
|
||||
}
|
||||
|
||||
func (m *openAIWSPassthroughUsageMeta) turnModels(fallback string) (string, string) {
|
||||
requestModel := strings.TrimSpace(fallback)
|
||||
upstreamModel := requestModel
|
||||
if m == nil {
|
||||
return requestModel, upstreamModel
|
||||
}
|
||||
if current := m.requestModel.Load(); current != nil && strings.TrimSpace(*current) != "" {
|
||||
requestModel = strings.TrimSpace(*current)
|
||||
}
|
||||
if current := m.upstreamModel.Load(); current != nil && strings.TrimSpace(*current) != "" {
|
||||
upstreamModel = strings.TrimSpace(*current)
|
||||
}
|
||||
return requestModel, upstreamModel
|
||||
}
|
||||
|
||||
func openAIWSTrimmedStringPtr(value string) *string {
|
||||
value = strings.TrimSpace(value)
|
||||
if value == "" {
|
||||
return nil
|
||||
}
|
||||
return &value
|
||||
}
|
||||
|
||||
func openAIWSDifferentModel(requestModel, upstreamModel string) string {
|
||||
upstreamModel = strings.TrimSpace(upstreamModel)
|
||||
if upstreamModel == "" || upstreamModel == strings.TrimSpace(requestModel) {
|
||||
return ""
|
||||
}
|
||||
return upstreamModel
|
||||
}
|
||||
|
||||
func openAIWSPassthroughRequestModelForFrame(payload []byte) string {
|
||||
@@ -585,6 +636,9 @@ func (c *openAIWSClientFrameConn) WriteFrame(ctx context.Context, msgType coderw
|
||||
if normalized, changed := normalizeCompletedImageGenerationStatus(payload); changed {
|
||||
payload = normalized
|
||||
}
|
||||
if c.restoreResponseModel != nil {
|
||||
payload = c.restoreResponseModel(payload)
|
||||
}
|
||||
}
|
||||
return c.conn.Write(ctx, msgType, payload)
|
||||
}
|
||||
@@ -655,10 +709,25 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
|
||||
// negotiated at session.update time. Without this fallback, an empty
|
||||
// model would miss any admin-configured model whitelist and be silently
|
||||
// passed through, defeating that policy on every frame after the first.
|
||||
capturedSessionModel := openAIWSPassthroughPolicyModelForFrame(account, firstClientMessage)
|
||||
initialRequestModel := ""
|
||||
if hooks != nil {
|
||||
initialRequestModel = hooks.InitialRequestModel
|
||||
initialRequestModel = strings.TrimSpace(hooks.InitialRequestModel)
|
||||
}
|
||||
if initialRequestModel == "" {
|
||||
initialRequestModel = openAIWSPassthroughRequestModelForFrame(firstClientMessage)
|
||||
}
|
||||
if hooks != nil && hooks.MapRequestModel != nil {
|
||||
mappedModel, mapErr := hooks.MapRequestModel(1, initialRequestModel)
|
||||
if mapErr != nil {
|
||||
return mapErr
|
||||
}
|
||||
if mappedModel = strings.TrimSpace(mappedModel); mappedModel != "" {
|
||||
firstClientMessage = s.ReplaceModelInBody(firstClientMessage, mappedModel)
|
||||
}
|
||||
}
|
||||
capturedSessionModel := openAIWSPassthroughPolicyModelForFrame(account, firstClientMessage)
|
||||
if capturedSessionModel != "" && capturedSessionModel != strings.TrimSpace(gjson.GetBytes(firstClientMessage, "model").String()) {
|
||||
firstClientMessage = s.ReplaceModelInBody(firstClientMessage, capturedSessionModel)
|
||||
}
|
||||
usageMeta := newOpenAIWSPassthroughUsageMeta(initialRequestModel, firstClientMessage)
|
||||
updatedFirst, blocked, policyErr := s.applyOpenAIFastPolicyToWSResponseCreate(ctx, account, capturedSessionModel, firstClientMessage)
|
||||
@@ -839,6 +908,14 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
|
||||
controlCtx: ctx,
|
||||
interTurnIdleTimeout: s.openAIWSIngressInterTurnIdleTimeout(),
|
||||
interTurnStarted: make(chan struct{}, 1),
|
||||
restoreResponseModel: func(payload []byte) []byte {
|
||||
eventType := strings.TrimSpace(gjson.GetBytes(payload, "type").String())
|
||||
if !openAIWSEventMayContainModel(eventType) {
|
||||
return payload
|
||||
}
|
||||
requestModel, upstreamModel := usageMeta.turnModels("")
|
||||
return replaceOpenAIWSMessageModel(payload, upstreamModel, requestModel)
|
||||
},
|
||||
}
|
||||
policyClientConn := &openAIWSPolicyEnforcingFrameConn{
|
||||
inner: clientFrameConn,
|
||||
@@ -878,17 +955,29 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
|
||||
}
|
||||
}
|
||||
}
|
||||
if isResponseCreate && hooks != nil && hooks.BeforeRequest != nil {
|
||||
turnNo := int(completedTurns.Load()) + 1
|
||||
if turnNo < 2 {
|
||||
turnNo = 2
|
||||
turnNo := int(completedTurns.Load()) + 1
|
||||
if turnNo < 2 {
|
||||
turnNo = 2
|
||||
}
|
||||
requestModelForThisFrame := ""
|
||||
if isResponseCreate {
|
||||
requestModelForThisFrame = usageMeta.requestModelForFrame(payload)
|
||||
if requestModelForThisFrame == "" {
|
||||
requestModelForThisFrame = capturedSessionModel
|
||||
}
|
||||
requestModel := usageMeta.requestModelForFrame(payload)
|
||||
if requestModel == "" {
|
||||
requestModel = capturedSessionModel
|
||||
if hooks != nil && hooks.BeforeRequest != nil {
|
||||
if err := hooks.BeforeRequest(turnNo, payload, requestModelForThisFrame); err != nil {
|
||||
return payload, nil, err
|
||||
}
|
||||
}
|
||||
if err := hooks.BeforeRequest(turnNo, payload, requestModel); err != nil {
|
||||
return payload, nil, err
|
||||
if hooks != nil && hooks.MapRequestModel != nil {
|
||||
upstreamModel, err := hooks.MapRequestModel(turnNo, requestModelForThisFrame)
|
||||
if err != nil {
|
||||
return payload, nil, err
|
||||
}
|
||||
if upstreamModel = strings.TrimSpace(upstreamModel); upstreamModel != "" {
|
||||
payload = s.ReplaceModelInBody(payload, upstreamModel)
|
||||
}
|
||||
}
|
||||
}
|
||||
// 在评估策略前先刷新 capturedSessionModel:客户端可能通过
|
||||
@@ -902,7 +991,9 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
|
||||
capturedSessionModel = updated
|
||||
}
|
||||
usageMeta.updateSessionRequestModel(payload)
|
||||
requestModelForThisFrame := usageMeta.requestModelForFrame(payload)
|
||||
if requestModelForThisFrame == "" {
|
||||
requestModelForThisFrame = usageMeta.requestModelForFrame(payload)
|
||||
}
|
||||
// Per-frame model first; if the client omits "model" on a
|
||||
// follow-up frame (legal in Realtime), fall back to the
|
||||
// session-level model captured from the first frame so the
|
||||
@@ -912,6 +1003,9 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
|
||||
if model == "" {
|
||||
model = capturedSessionModel
|
||||
}
|
||||
if isResponseCreate && model != "" && model != strings.TrimSpace(gjson.GetBytes(payload, "model").String()) {
|
||||
payload = s.ReplaceModelInBody(payload, model)
|
||||
}
|
||||
out, blocked, policyErr := s.applyOpenAIFastPolicyToWSResponseCreate(ctx, account, model, payload)
|
||||
// 多轮 passthrough usage:仅在成功(non-block / non-err)
|
||||
// 的 response.create 帧上更新 usageMeta,使用
|
||||
@@ -999,6 +1093,7 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
|
||||
},
|
||||
OnTurnComplete: func(turn openaiwsv2.RelayTurnResult) {
|
||||
turnNo := int(completedTurns.Add(1))
|
||||
turnRequestModel, turnUpstreamModel := usageMeta.turnModels(turn.RequestModel)
|
||||
turnResult := &OpenAIForwardResult{
|
||||
RequestID: turn.RequestID,
|
||||
Usage: OpenAIUsage{
|
||||
@@ -1008,7 +1103,8 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
|
||||
CacheReadInputTokens: turn.Usage.CacheReadInputTokens,
|
||||
ImageOutputTokens: turn.Usage.ImageOutputTokens,
|
||||
},
|
||||
Model: turn.RequestModel,
|
||||
Model: turnRequestModel,
|
||||
UpstreamModel: openAIWSDifferentModel(turnRequestModel, turnUpstreamModel),
|
||||
ServiceTier: usageMeta.serviceTier.Load(),
|
||||
ReasoningEffort: usageMeta.reasoningEffort.Load(),
|
||||
Stream: true,
|
||||
@@ -1019,11 +1115,13 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
|
||||
FirstTokenMs: turn.FirstTokenMs,
|
||||
}
|
||||
logOpenAIWSV2Passthrough(
|
||||
"relay_turn_completed account_id=%d turn=%d request_id=%s terminal_event=%s duration_ms=%d first_token_ms=%d input_tokens=%d output_tokens=%d cache_read_tokens=%d",
|
||||
"relay_turn_completed account_id=%d turn=%d request_id=%s terminal_event=%s turn_requested_model=%s turn_upstream_model=%s duration_ms=%d first_token_ms=%d input_tokens=%d output_tokens=%d cache_read_tokens=%d",
|
||||
account.ID,
|
||||
turnNo,
|
||||
truncateOpenAIWSLogValue(turnResult.RequestID, openAIWSIDValueMaxLen),
|
||||
truncateOpenAIWSLogValue(turn.TerminalEventType, openAIWSLogValueMaxLen),
|
||||
truncateOpenAIWSLogValue(turnRequestModel, openAIWSLogValueMaxLen),
|
||||
truncateOpenAIWSLogValue(turnUpstreamModel, openAIWSLogValueMaxLen),
|
||||
turnResult.Duration.Milliseconds(),
|
||||
openAIWSFirstTokenMsForLog(turnResult.FirstTokenMs),
|
||||
turnResult.Usage.InputTokens,
|
||||
@@ -1114,6 +1212,7 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
|
||||
return NewOpenAIWSClientCloseError(status, reason, cause)
|
||||
}
|
||||
|
||||
resultRequestModel, resultUpstreamModel := usageMeta.turnModels(relayResult.RequestModel)
|
||||
result := &OpenAIForwardResult{
|
||||
RequestID: relayResult.RequestID,
|
||||
Usage: OpenAIUsage{
|
||||
@@ -1123,7 +1222,8 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
|
||||
CacheReadInputTokens: relayResult.Usage.CacheReadInputTokens,
|
||||
ImageOutputTokens: relayResult.Usage.ImageOutputTokens,
|
||||
},
|
||||
Model: relayResult.RequestModel,
|
||||
Model: resultRequestModel,
|
||||
UpstreamModel: openAIWSDifferentModel(resultRequestModel, resultUpstreamModel),
|
||||
ServiceTier: usageMeta.serviceTier.Load(),
|
||||
ReasoningEffort: usageMeta.reasoningEffort.Load(),
|
||||
Stream: true,
|
||||
|
||||
Reference in New Issue
Block a user