Merge pull request #4787 from KtzeAbyss/fix/4760-ws-turn-model-billing

fix(openai): track WebSocket models per turn
This commit is contained in:
Wesley Liddick
2026-07-27 10:25:57 +08:00
committed by GitHub
9 changed files with 719 additions and 101 deletions
@@ -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,