package service import ( "context" "errors" "fmt" "io" "net/http" "net/http/httptest" "strings" "testing" "time" "github.com/Wei-Shaw/sub2api/internal/config" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" coderws "github.com/coder/websocket" "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" "github.com/tidwall/gjson" ) func TestResolveOpenAIWSClientFirstMessageTimeout(t *testing.T) { defaultTimeout := time.Duration(config.DefaultOpenAIWSClientFirstMessageTimeoutSeconds) * time.Second require.Equal(t, defaultTimeout, ResolveOpenAIWSClientFirstMessageTimeout(nil)) cfg := &config.Config{} require.Equal(t, defaultTimeout, ResolveOpenAIWSClientFirstMessageTimeout(cfg)) cfg.Gateway.OpenAIWS.ClientFirstMessageTimeoutSeconds = 120 require.Equal(t, 120*time.Second, ResolveOpenAIWSClientFirstMessageTimeout(cfg)) } func TestPrepareOpenAIWSHTTPBridgeBodyStripsWSFields(t *testing.T) { body, err := prepareOpenAIWSHTTPBridgeBody([]byte(`{"type":"response.create","generate":true,"model":"gpt-5","stream":false,"previous_response_id":"resp_prev","input":"hi"}`)) require.NoError(t, err) require.False(t, gjson.GetBytes(body, "type").Exists()) require.False(t, gjson.GetBytes(body, "generate").Exists()) require.False(t, gjson.GetBytes(body, "previous_response_id").Exists()) require.Equal(t, "gpt-5", gjson.GetBytes(body, "model").String()) require.True(t, gjson.GetBytes(body, "stream").Bool()) require.Equal(t, "hi", gjson.GetBytes(body, "input").String()) } func TestProxyOpenAIWSHTTPBridgeTurnAPIKeyAdaptsClientTools(t *testing.T) { gin.SetMode(gin.TestMode) sse := strings.Join([]string{ `data: {"type":"response.output_item.added","sequence_number":0,"output_index":0,"item":{"type":"function_call","id":"item_exec","call_id":"call_exec","name":"exec","status":"in_progress"}}`, ``, `data: {"type":"response.function_call_arguments.done","sequence_number":1,"output_index":0,"item_id":"item_exec","call_id":"call_exec","name":"exec","arguments":"{\"input\":\"pwd\"}"}`, ``, `data: {"type":"response.output_item.done","sequence_number":2,"output_index":0,"item":{"type":"function_call","id":"item_exec","call_id":"call_exec","name":"exec","arguments":"{\"input\":\"pwd\"}","status":"completed"}}`, ``, `data: {"type":"response.completed","sequence_number":3,"response":{"id":"resp_tools","status":"completed","output":[{"type":"function_call","id":"item_exec","call_id":"call_exec","name":"exec","arguments":"{\"input\":\"pwd\"}","status":"completed"}],"usage":{"input_tokens":1,"output_tokens":1}}}`, ``, }, "\n") upstream := &httpUpstreamRecorder{resp: &http.Response{ StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(sse)), }} svc := &OpenAIGatewayService{ cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}, httpUpstream: upstream, } account := &Account{ID: 5659, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Concurrency: 1} payload := []byte(`{ "type":"response.create","model":"gpt-5","stream":true, "tools":[{"type":"custom","name":"exec","description":"Run a command"}], "input":[ {"type":"custom_tool_call","id":"previous_item","call_id":"previous_call","name":"exec","input":"echo ready"}, {"type":"custom_tool_call_output","call_id":"previous_call","output":"ready"} ] }`) recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) var events [][]byte result, err := svc.proxyOpenAIWSHTTPBridgeTurn( context.Background(), c, account, "test-token", payload, len(payload), "gpt-5", "", "", "", "", 2, func(message []byte) error { events = append(events, append([]byte(nil), message...)) return nil }, ) require.NoError(t, err) require.NotNil(t, result) require.Equal(t, "function", gjson.GetBytes(upstream.lastBody, "tools.0.type").String()) require.Equal(t, "function_call", gjson.GetBytes(upstream.lastBody, "input.0.type").String()) require.JSONEq(t, `{"input":"echo ready"}`, gjson.GetBytes(upstream.lastBody, "input.0.arguments").String()) require.False(t, gjson.GetBytes(upstream.lastBody, "input.0.input").Exists()) require.Equal(t, "function_call_output", gjson.GetBytes(upstream.lastBody, "input.1.type").String()) var outputDone, completed []byte for _, event := range events { switch gjson.GetBytes(event, "type").String() { case "response.output_item.done": outputDone = event case "response.completed": completed = event } } require.NotEmpty(t, outputDone) require.Equal(t, "custom_tool_call", gjson.GetBytes(outputDone, "item.type").String()) require.Equal(t, "pwd", gjson.GetBytes(outputDone, "item.input").String()) require.False(t, gjson.GetBytes(outputDone, "item.arguments").Exists()) require.NotEmpty(t, completed) require.Equal(t, "custom_tool_call", gjson.GetBytes(completed, "response.output.0.type").String()) require.Equal(t, "pwd", gjson.GetBytes(completed, "response.output.0.input").String()) require.True(t, result.wsReplayInputExists) require.Len(t, result.wsReplayInput, 1) require.Equal(t, "custom_tool_call", gjson.GetBytes(result.wsReplayInput[0], "type").String()) require.Equal(t, "pwd", gjson.GetBytes(result.wsReplayInput[0], "input").String()) } func TestProxyOpenAIWSHTTPBridgeTurnAPIKeyRestoresClientToolsInResponseDone(t *testing.T) { gin.SetMode(gin.TestMode) sse := strings.Join([]string{ `data: {"type":"response.output_item.added","sequence_number":0,"output_index":0,"item":{"type":"function_call","id":"item_exec","call_id":"call_exec","name":"exec","status":"in_progress"}}`, ``, `data: {"type":"response.function_call_arguments.done","sequence_number":1,"output_index":0,"item_id":"item_exec","call_id":"call_exec","name":"exec","arguments":"{\"input\":\"pwd\"}"}`, ``, `data: {"type":"response.done","sequence_number":2,"response":{"id":"resp_tools","status":"completed","output":[{"type":"function_call","id":"item_exec","call_id":"call_exec","name":"exec","arguments":"{\"input\":\"pwd\"}","status":"completed"}],"usage":{"input_tokens":1,"output_tokens":1}}}`, ``, }, "\n") upstream := &httpUpstreamRecorder{resp: &http.Response{ StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(sse)), }} svc := &OpenAIGatewayService{ cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}, httpUpstream: upstream, } account := &Account{ID: 5764, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Concurrency: 1} payload := []byte(`{ "type":"response.create","model":"gpt-5","stream":true, "tools":[{"type":"custom","name":"exec","description":"Run a command"}], "input":"run pwd" }`) recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) var events [][]byte result, err := svc.proxyOpenAIWSHTTPBridgeTurn( context.Background(), c, account, "test-token", payload, len(payload), "gpt-5", "", "", "", "", 1, func(message []byte) error { events = append(events, append([]byte(nil), message...)) return nil }, ) require.NoError(t, err) require.NotNil(t, result) require.Len(t, events, 4) terminal := events[len(events)-1] require.Equal(t, "response.done", gjson.GetBytes(terminal, "type").String()) require.Equal(t, int64(3), gjson.GetBytes(terminal, "sequence_number").Int()) require.Equal(t, "custom_tool_call", gjson.GetBytes(terminal, "response.output.0.type").String()) require.Equal(t, "pwd", gjson.GetBytes(terminal, "response.output.0.input").String()) require.False(t, gjson.GetBytes(terminal, "response.output.0.arguments").Exists()) require.True(t, result.wsReplayInputExists) require.Len(t, result.wsReplayInput, 1) require.Equal(t, "custom_tool_call", gjson.GetBytes(result.wsReplayInput[0], "type").String()) } func TestProxyOpenAIWSHTTPBridgeTurnGrokPromotesDiscoveryAndRestoresNamespaceSSE(t *testing.T) { gin.SetMode(gin.TestMode) sse := strings.Join([]string{ `data: {"type":"response.output_item.added","sequence_number":0,"output_index":0,"item":{"type":"function_call","id":"item_spawn","call_id":"call_spawn","name":"multi_agent_v1__spawn_agent","status":"in_progress"}}`, "", `data: {"type":"response.function_call_arguments.done","sequence_number":1,"output_index":0,"item_id":"item_spawn","call_id":"call_spawn","name":"multi_agent_v1__spawn_agent","arguments":"{\"message\":\"work\"}"}`, "", `data: {"type":"response.output_item.done","sequence_number":2,"output_index":0,"item":{"type":"function_call","id":"item_spawn","call_id":"call_spawn","name":"multi_agent_v1__spawn_agent","arguments":"{\"message\":\"work\"}","status":"completed"}}`, "", `data: {"type":"response.completed","sequence_number":3,"response":{"id":"resp_spawn","status":"completed","output":[{"type":"function_call","id":"item_spawn","call_id":"call_spawn","name":"multi_agent_v1__spawn_agent","arguments":"{\"message\":\"work\"}","status":"completed"}],"usage":{"input_tokens":1,"output_tokens":1}}}`, "", }, "\n") upstream := &httpUpstreamRecorder{resp: &http.Response{ StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(sse)), }} svc := &OpenAIGatewayService{ cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}, httpUpstream: upstream, } account := &Account{ ID: 5765, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, Credentials: map[string]any{"base_url": xai.DefaultCLIBaseURL}, } payload := []byte(`{ "type":"response.create","model":"grok-4.5","stream":true, "tools":[{"type":"tool_search"}], "input":[ {"type":"tool_search_call","call_id":"call_search","arguments":{"query":"subagent"},"status":"completed"}, {"type":"tool_search_output","call_id":"call_search","execution":"client","status":"completed","tools":[ {"type":"namespace","name":"multi_agent_v1","tools":[{"type":"function","name":"spawn_agent","parameters":{"type":"object"}}]} ]} ] }`) recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) var events [][]byte result, err := svc.proxyOpenAIWSHTTPBridgeTurn( context.Background(), c, account, "access-token", payload, len(payload), "grok-4.5", "", "", "", "", 1, func(message []byte) error { events = append(events, append([]byte(nil), message...)) return nil }, ) require.NoError(t, err) require.NotNil(t, result) require.Equal(t, "multi_agent_v1__spawn_agent", gjson.GetBytes(upstream.lastBody, "tools.1.name").String()) require.Equal(t, "function_call_output", gjson.GetBytes(upstream.lastBody, "input.1.type").String()) require.False(t, gjson.GetBytes(upstream.lastBody, "input.1.tools").Exists()) require.False(t, gjson.GetBytes(upstream.lastBody, "input.1.status").Exists()) require.False(t, gjson.GetBytes(upstream.lastBody, "input.1.execution").Exists()) state, ok := openAIWSHTTPBridgeToolStateFromContext(c) require.True(t, ok) require.Equal(t, "multi_agent_v1", state.ClientMapping.NamespaceTools["multi_agent_v1__spawn_agent"].Namespace) require.Len(t, events, 4) require.Equal(t, "spawn_agent", gjson.GetBytes(events[0], "item.name").String()) require.Equal(t, "multi_agent_v1", gjson.GetBytes(events[0], "item.namespace").String()) require.Equal(t, "spawn_agent", gjson.GetBytes(events[2], "item.name").String()) require.Equal(t, "multi_agent_v1", gjson.GetBytes(events[2], "item.namespace").String()) require.Equal(t, "spawn_agent", gjson.GetBytes(events[3], "response.output.0.name").String()) require.Equal(t, "multi_agent_v1", gjson.GetBytes(events[3], "response.output.0.namespace").String()) } func TestProxyOpenAIWSHTTPBridgeTurnGrokInheritsToolSearchAndPromotesFollowupDiscovery(t *testing.T) { gin.SetMode(gin.TestMode) firstSSE := "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_first\",\"output\":[],\"usage\":{\"input_tokens\":1,\"output_tokens\":1}}}\n\n" secondSSE := "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_second\",\"output\":[{\"type\":\"function_call\",\"id\":\"item_spawn\",\"call_id\":\"call_spawn\",\"name\":\"multi_agent_v1__spawn_agent\",\"arguments\":\"{}\"}],\"usage\":{\"input_tokens\":1,\"output_tokens\":1}}}\n\n" upstream := &httpUpstreamRecorder{responses: []*http.Response{ {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(firstSSE))}, {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(secondSSE))}, }} svc := &OpenAIGatewayService{ cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}, httpUpstream: upstream, } account := &Account{ ID: 5766, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, Credentials: map[string]any{"base_url": xai.DefaultCLIBaseURL, "subscription_tier": "free"}, } recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) first := []byte(`{"type":"response.create","model":"grok-4.5","stream":true,"tools":[{"type":"tool_search"}],"input":"discover tools"}`) _, err := svc.proxyOpenAIWSHTTPBridgeTurn( context.Background(), c, account, "access-token", first, len(first), "grok-4.5", "", "", "", "grok-ws-cache", 1, func([]byte) error { return nil }, ) require.NoError(t, err) second := []byte(`{ "type":"response.create","model":"grok-4.5","stream":true, "input":[{"type":"tool_search_output","call_id":"call_search","status":"completed","tools":[ {"type":"namespace","name":"multi_agent_v1","tools":[{"type":"function","name":"spawn_agent","parameters":{"type":"object"}}]} ]}] }`) var events [][]byte _, err = svc.proxyOpenAIWSHTTPBridgeTurn( context.Background(), c, account, "access-token", second, len(second), "grok-4.5", "", "", "", "grok-ws-cache", 2, func(message []byte) error { events = append(events, append([]byte(nil), message...)) return nil }, ) require.NoError(t, err) require.Len(t, upstream.bodies, 2) require.Equal(t, "tool_search", gjson.GetBytes(upstream.bodies[1], "tools.0.name").String()) require.Equal(t, "multi_agent_v1__spawn_agent", gjson.GetBytes(upstream.bodies[1], "tools.1.name").String()) require.NotEqual(t, grokFreeCacheDisabledToolChoice, gjson.GetBytes(upstream.bodies[1], "tool_choice").String()) require.Equal(t, "grok-ws-cache", gjson.GetBytes(upstream.bodies[1], "prompt_cache_key").String()) require.Equal(t, "function_call_output", gjson.GetBytes(upstream.bodies[1], "input.0.type").String()) require.Len(t, events, 1) require.Equal(t, "spawn_agent", gjson.GetBytes(events[0], "response.output.0.name").String()) require.Equal(t, "multi_agent_v1", gjson.GetBytes(events[0], "response.output.0.namespace").String()) } func TestOpenAIWSHTTPBridgeAPIKeyReusesClientToolMappingWhenFollowupOmitsTools(t *testing.T) { gin.SetMode(gin.TestMode) firstSSEBody := strings.Join([]string{ `data: {"type":"response.completed","response":{"id":"resp_custom_first","model":"gpt-5.6-sol","output":[{"type":"function_call","id":"fc_custom_1","call_id":"call_custom_1","name":"exec","arguments":"{\"input\":\"pwd\"}"}],"usage":{"input_tokens":9,"output_tokens":1}}}`, "", }, "\n") secondSSEBody := strings.Join([]string{ `data: {"type":"response.completed","response":{"id":"resp_custom_second","model":"gpt-5.6-sol","output":[],"usage":{"input_tokens":1,"output_tokens":1}}}`, "", }, "\n") upstream := &httpUpstreamRecorder{responses: []*http.Response{ {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(firstSSEBody))}, {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(secondSSEBody))}, }} cfg := &config.Config{} cfg.Security.URLAllowlist.Enabled = false cfg.Security.URLAllowlist.AllowInsecureHTTP = true cfg.Gateway.OpenAIWS.Enabled = true cfg.Gateway.OpenAIWS.APIKeyEnabled = true cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = true cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes = 1 cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8 cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3 cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3 cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3 svc := &OpenAIGatewayService{ cfg: cfg, httpUpstream: upstream, cache: &stubGatewayCache{}, openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), toolCorrector: NewCodexToolCorrector(), } account := &Account{ ID: 9001, Name: "api-key-custom-followup", Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "sk-upstream"}, Extra: map[string]any{"responses_websockets_v2_enabled": true}, Concurrency: 1, Status: StatusActive, Schedulable: true, } errCh := make(chan error, 1) wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { conn, err := coderws.Accept(w, r, nil) if err != nil { errCh <- err return } defer func() { _ = conn.CloseNow() }() readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second) _, firstMessage, err := conn.Read(readCtx) cancelRead() if err != nil { errCh <- err return } rec := httptest.NewRecorder() ginCtx, _ := gin.CreateTestContext(rec) ginCtx.Request = r.Clone(r.Context()) errCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "sk-test", firstMessage, nil) })) defer wsServer.Close() dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second) clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil) cancelDial() require.NoError(t, err) defer func() { _ = clientConn.CloseNow() }() writeMessage := func(payload string) { writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second) defer cancelWrite() require.NoError(t, clientConn.Write(writeCtx, coderws.MessageText, []byte(payload))) } readMessage := func() []byte { readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second) defer cancelRead() messageType, event, readErr := clientConn.Read(readCtx) require.NoError(t, readErr) require.Equal(t, coderws.MessageText, messageType) return event } writeMessage(`{"type":"response.create","model":"gpt-5.6-sol","stream":true,"tools":[{"type":"custom","name":"exec"}],"input":"run pwd"}`) firstEvent := readMessage() require.Equal(t, "response.completed", gjson.GetBytes(firstEvent, "type").String()) require.Equal(t, "custom_tool_call", gjson.GetBytes(firstEvent, "response.output.0.type").String()) require.Equal(t, "pwd", gjson.GetBytes(firstEvent, "response.output.0.input").String()) writeMessage(`{"type":"response.create","model":"gpt-5.6-sol","stream":true,"previous_response_id":"resp_custom_first","input":[{"type":"custom_tool_call_output","id":"ctco_client_output_1","call_id":"call_custom_1","output":"ok"}]}`) secondEvent := readMessage() require.Equal(t, "response.completed", gjson.GetBytes(secondEvent, "type").String()) require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done")) select { case proxyErr := <-errCh: require.NoError(t, proxyErr) case <-time.After(5 * time.Second): t.Fatal("timed out waiting for websocket bridge proxy to finish") } require.Len(t, upstream.bodies, 2) firstTools := gjson.GetBytes(upstream.bodies[0], "tools").Array() require.Len(t, firstTools, 1) require.Equal(t, "function", firstTools[0].Get("type").String()) secondTools := gjson.GetBytes(upstream.bodies[1], "tools").Array() require.Len(t, secondTools, 1) require.Equal(t, "function", secondTools[0].Get("type").String()) require.Equal(t, "exec", secondTools[0].Get("name").String()) secondInput := gjson.GetBytes(upstream.bodies[1], "input").Array() require.Len(t, secondInput, 3) require.Equal(t, "run pwd", secondInput[0].String()) require.Equal(t, "function_call", secondInput[1].Get("type").String()) require.Equal(t, "fc_custom_1", secondInput[1].Get("id").String()) require.JSONEq(t, `{"input":"pwd"}`, secondInput[1].Get("arguments").String()) require.False(t, secondInput[1].Get("input").Exists()) require.Equal(t, "function_call_output", secondInput[2].Get("type").String()) require.False(t, secondInput[2].Get("id").Exists()) } func TestOpenAIWSHTTPBridgeDecisionKeepsSmallFramesOnWS(t *testing.T) { svc := &OpenAIGatewayService{ cfg: &config.Config{ Gateway: config.GatewayConfig{ OpenAIWS: config.GatewayOpenAIWSConfig{ HTTPBridgeEnabled: true, HTTPBridgeThresholdBytes: 100, }, }, }, } require.False(t, svc.shouldBridgeOpenAIWSHTTP(nil, 99, "")) require.True(t, svc.shouldBridgeOpenAIWSHTTP(nil, 100, "")) require.False(t, svc.shouldBridgeOpenAIWSHTTP(nil, 1000, "resp_existing")) svc.cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = false require.False(t, svc.shouldBridgeOpenAIWSHTTP(nil, 1000, "")) require.True(t, svc.shouldBridgeOpenAIWSHTTP(&Account{Platform: PlatformGrok}, 1, "resp_existing")) } func TestProxyOpenAIWSHTTPBridgeTurnTransportErrorFailoverSafety(t *testing.T) { gin.SetMode(gin.TestMode) tests := []struct { name string turn int wantFailover bool wantWrites int }{ {name: "first_turn_fails_over_before_downstream_event", turn: 1, wantFailover: true}, {name: "later_turn_does_not_replay_completed_turns", turn: 2, wantWrites: 1}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { upstream := &httpUpstreamRecorder{err: io.EOF} svc := &OpenAIGatewayService{ cfg: &config.Config{}, httpUpstream: upstream, } account := &Account{ ID: 8, Name: "api-key", Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Concurrency: 1, } recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) payload := []byte(`{"type":"response.create","model":"gpt-5","input":"hi"}`) var writes [][]byte result, err := svc.proxyOpenAIWSHTTPBridgeTurn( context.Background(), c, account, "sk-test", payload, len(payload), "gpt-5", "", "", "", "", tt.turn, func(message []byte) error { writes = append(writes, append([]byte(nil), message...)) return nil }, ) require.Nil(t, result) var failoverErr *UpstreamFailoverError if tt.wantFailover { require.ErrorAs(t, err, &failoverErr) require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode) require.JSONEq(t, string(openAITransportFailoverBody), string(failoverErr.ResponseBody)) } else { require.Error(t, err) require.False(t, errors.As(err, &failoverErr)) } require.Len(t, writes, tt.wantWrites) if tt.wantWrites > 0 { require.Equal(t, "error", gjson.GetBytes(writes[0], "type").String()) require.Equal(t, int64(http.StatusBadGateway), gjson.GetBytes(writes[0], "status").Int()) } }) } } func TestProxyOpenAIWSHTTPBridgeTurnHTTPStatusFailoverSafety(t *testing.T) { gin.SetMode(gin.TestMode) tests := []struct { name string turn int status int wantFailover bool wantWrites int }{ {name: "first_turn_401", turn: 1, status: http.StatusUnauthorized, wantFailover: true}, {name: "first_turn_429", turn: 1, status: http.StatusTooManyRequests, wantFailover: true}, {name: "first_turn_500", turn: 1, status: http.StatusInternalServerError, wantFailover: true}, {name: "later_turn_500_does_not_replay", turn: 2, status: http.StatusInternalServerError, wantWrites: 1}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { upstream := &httpUpstreamRecorder{resp: &http.Response{ StatusCode: tt.status, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(`{"error":{"type":"server_error","message":"temporary upstream failure"}}`)), }} svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} account := &Account{ID: 9, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Concurrency: 1} recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) payload := []byte(`{"type":"response.create","model":"gpt-5","input":"hi"}`) var writes [][]byte result, err := svc.proxyOpenAIWSHTTPBridgeTurn( context.Background(), c, account, "sk-test", payload, len(payload), "gpt-5", "", "", "", "", tt.turn, func(message []byte) error { writes = append(writes, append([]byte(nil), message...)) return nil }, ) require.Nil(t, result) var failoverErr *UpstreamFailoverError if tt.wantFailover { require.ErrorAs(t, err, &failoverErr) require.Equal(t, tt.status, failoverErr.StatusCode) } else { require.Error(t, err) require.False(t, errors.As(err, &failoverErr)) } require.Len(t, writes, tt.wantWrites) }) } } func TestProxyOpenAIWSHTTPBridgeTurnSSEErrorFailoverSafety(t *testing.T) { gin.SetMode(gin.TestMode) for _, turn := range []int{1, 2} { t.Run(fmt.Sprintf("turn_%d", turn), func(t *testing.T) { upstream := &httpUpstreamRecorder{resp: &http.Response{ StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader( "data: {\"type\":\"error\",\"error\":{\"type\":\"rate_limit_error\",\"code\":\"rate_limit_exceeded\",\"message\":\"limited\"}}\n\n", )), }} svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} account := &Account{ID: 10, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Concurrency: 1} recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) payload := []byte(`{"type":"response.create","model":"gpt-5","input":"hi"}`) var writes [][]byte result, err := svc.proxyOpenAIWSHTTPBridgeTurn( context.Background(), c, account, "sk-test", payload, len(payload), "gpt-5", "", "", "", "", turn, func(message []byte) error { writes = append(writes, append([]byte(nil), message...)) return nil }, ) var failoverErr *UpstreamFailoverError require.Nil(t, result) require.ErrorAs(t, err, &failoverErr) require.Equal(t, http.StatusTooManyRequests, failoverErr.StatusCode) require.Empty(t, writes) }) } } // 桥接转发 error / response.failed 给 WS 客户端前必须把容量降载码改写为可重试 // 的 server_error:Codex 对 server_is_overloaded/slow_down 判致命并终止会话。 // 账号状态判定使用改写前的原始事件,不受影响。 func TestProxyOpenAIWSHTTPBridgeTurnRewritesCapacityShedCodeForClient(t *testing.T) { gin.SetMode(gin.TestMode) tests := []struct { name string turn int body string wantErr bool }{ { name: "turn2_error_frame", turn: 2, body: "data: {\"type\":\"error\",\"error\":{\"type\":\"service_unavailable_error\",\"code\":\"server_is_overloaded\",\"message\":\"Our servers are currently overloaded. Please try again later.\"}}\n\n", wantErr: true, }, { // 后续 turn 不允许 replay,容量错误必须改写后交给客户端重试。 name: "turn2_bare_response_failed", turn: 2, body: "data: {\"type\":\"response.failed\",\"response\":{\"id\":\"resp_shed\",\"status\":\"failed\",\"error\":{\"code\":\"server_is_overloaded\",\"message\":\"Our servers are currently overloaded. Please try again later.\"}}}\n\n", }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { upstream := &httpUpstreamRecorder{resp: &http.Response{ StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(tt.body)), }} svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} account := &Account{ID: 11, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Concurrency: 1} recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) payload := []byte(`{"type":"response.create","model":"gpt-5","input":"hi"}`) var writes [][]byte _, err := svc.proxyOpenAIWSHTTPBridgeTurn( context.Background(), c, account, "sk-test", payload, len(payload), "gpt-5", "", "", "", "", tt.turn, func(message []byte) error { writes = append(writes, append([]byte(nil), message...)) return nil }, ) if tt.wantErr { require.Error(t, err) } else { require.NoError(t, err) } require.Len(t, writes, 1) require.Contains(t, string(writes[0]), `"code":"server_error"`) require.NotContains(t, string(writes[0]), "server_is_overloaded") require.Contains(t, string(writes[0]), "Our servers are currently overloaded") }) } } func TestProxyOpenAIWSHTTPBridgeTurnStagesMetadataBeforeCapacityFailover(t *testing.T) { gin.SetMode(gin.TestMode) body := strings.Join([]string{ `data: {"type":"response.created","response":{"id":"resp_shed"}}`, "", `data: {"type":"response.in_progress","response":{"id":"resp_shed"}}`, "", `data: {"type":"response.failed","response":{"id":"resp_shed","status":"failed","error":{"message":"Our servers are currently overloaded. Please try again later."}}}`, "", }, "\n") upstream := &httpUpstreamRecorder{resp: &http.Response{ StatusCode: http.StatusOK, Header: http.Header{"X-Request-Id": []string{"rid-ws-bridge-capacity"}}, Body: io.NopCloser(strings.NewReader(body)), }} svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} account := &Account{ID: 12, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Concurrency: 1} recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) payload := []byte(`{"type":"response.create","model":"gpt-5","input":"hi"}`) var writes [][]byte result, err := svc.proxyOpenAIWSHTTPBridgeTurn( context.Background(), c, account, "sk-test", payload, len(payload), "gpt-5", "", "", "", "", 1, func(message []byte) error { writes = append(writes, append([]byte(nil), message...)) return nil }, ) require.Nil(t, result) var failoverErr *UpstreamFailoverError require.ErrorAs(t, err, &failoverErr) require.True(t, failoverErr.RetryableOnSameAccount) require.True(t, failoverErr.RequestScopedTransient) require.Empty(t, writes) } func TestProxyOpenAIWSHTTPBridgeTurnDoesNotReplayCapacityAfterSemanticOutput(t *testing.T) { gin.SetMode(gin.TestMode) logSink, restore := captureStructuredLog(t) defer restore() body := strings.Join([]string{ `data: {"type":"response.created","response":{"id":"resp_partial"}}`, "", `data: {"type":"response.output_text.delta","delta":"partial"}`, "", `data: {"type":"response.failed","response":{"id":"resp_partial","status":"failed","error":{"code":"server_is_overloaded","message":"Our servers are currently overloaded. Please try again later."}}}`, "", }, "\n") upstream := &httpUpstreamRecorder{resp: &http.Response{ StatusCode: http.StatusOK, Header: http.Header{"X-Request-Id": []string{"rid-ws-bridge-post-output"}}, Body: io.NopCloser(strings.NewReader(body)), }} svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} account := &Account{ID: 13, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Concurrency: 1} recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) payload := []byte(`{"type":"response.create","model":"gpt-5","input":"hi"}`) var writes [][]byte result, err := svc.proxyOpenAIWSHTTPBridgeTurn( context.Background(), c, account, "sk-test", payload, len(payload), "gpt-5", "", "", "", "", 1, func(message []byte) error { writes = append(writes, append([]byte(nil), message...)) return nil }, ) require.NotNil(t, result) require.NoError(t, err) require.Len(t, writes, 3) var failoverErr *UpstreamFailoverError require.False(t, errors.As(err, &failoverErr)) require.Contains(t, string(writes[2]), `"code":"server_error"`) require.NotContains(t, string(writes[2]), "server_is_overloaded") require.True(t, logSink.ContainsMessage("gateway.failover_suppressed_after_semantic_output")) require.True(t, logSink.ContainsFieldValue("path", "ws_http_bridge")) } func TestProxyOpenAIWSHTTPBridgeTurnRequiresTerminalEvent(t *testing.T) { gin.SetMode(gin.TestMode) tests := []struct { name string body string wantFailover bool wantWrites int }{ {name: "done_without_events_fails_over", body: "data: [DONE]\n\n", wantFailover: true}, { name: "created_then_done_fails_over_before_semantic_output", body: "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_truncated\"}}\n\n" + "data: [DONE]\n\n", wantFailover: true, }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { upstream := &httpUpstreamRecorder{resp: &http.Response{ StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(tt.body)), }} svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} account := &Account{ID: 11, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Concurrency: 1} recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) payload := []byte(`{"type":"response.create","model":"gpt-5","input":"hi"}`) var writes [][]byte result, err := svc.proxyOpenAIWSHTTPBridgeTurn( context.Background(), c, account, "sk-test", payload, len(payload), "gpt-5", "", "", "", "", 1, func(message []byte) error { writes = append(writes, append([]byte(nil), message...)) return nil }, ) var failoverErr *UpstreamFailoverError if tt.wantFailover { require.Nil(t, result) require.ErrorAs(t, err, &failoverErr) } else { require.NotNil(t, result) require.Error(t, err) require.False(t, errors.As(err, &failoverErr)) } require.Len(t, writes, tt.wantWrites) }) } } func TestOpenAIWSHTTPBridgeRelaysSSEFramesAsWebSocketMessages(t *testing.T) { gin.SetMode(gin.TestMode) sseBody := strings.Join([]string{ `data: {"type":"response.created","response":{"id":"resp_bridge","model":"gpt-5"}}`, "", `data: {"type":"response.output_text.delta","response":{"id":"resp_bridge"},"delta":"ok"}`, "", `data: {"type":"response.completed","response":{"id":"resp_bridge","model":"gpt-5","usage":{"input_tokens":3,"output_tokens":2}}}`, "", }, "\n") upstream := &httpUpstreamRecorder{resp: &http.Response{ StatusCode: http.StatusOK, Header: http.Header{ "Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_bridge"}, }, Body: io.NopCloser(strings.NewReader(sseBody)), }} svc := &OpenAIGatewayService{ cfg: &config.Config{ Gateway: config.GatewayConfig{ MaxLineSize: defaultMaxLineSize, OpenAIWS: config.GatewayOpenAIWSConfig{ HTTPBridgeEnabled: true, HTTPBridgeThresholdBytes: 1, }, }, }, httpUpstream: upstream, toolCorrector: NewCodexToolCorrector(), } account := &Account{ ID: 7, Name: "api-key", Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Concurrency: 1, Status: StatusActive, } payload := []byte(`{"type":"response.create","generate":true,"model":"gpt-5","stream":true,"client_metadata":{"ws_request_header_x_openai_internal_codex_responses_lite":"true"},"input":"hi"}`) type bridgeResult struct { result *OpenAIForwardResult err error } resultCh := make(chan bridgeResult, 1) wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover}) if err != nil { resultCh <- bridgeResult{err: err} return } defer func() { _ = conn.CloseNow() }() rec := httptest.NewRecorder() ginCtx, _ := gin.CreateTestContext(rec) req := r.Clone(r.Context()) req.Header = req.Header.Clone() ginCtx.Request = req writeClient := func(message []byte) error { writeCtx, cancel := context.WithTimeout(r.Context(), 3*time.Second) defer cancel() return conn.Write(writeCtx, coderws.MessageText, message) } result, bridgeErr := svc.proxyOpenAIWSHTTPBridgeTurn( r.Context(), ginCtx, account, "sk-test", payload, len(payload), "gpt-5", "", "", "", "", 1, writeClient, ) resultCh <- bridgeResult{result: result, err: bridgeErr} })) defer wsServer.Close() dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second) clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil) cancelDial() require.NoError(t, err) defer func() { _ = clientConn.CloseNow() }() readEvent := func() []byte { readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second) msgType, event, readErr := clientConn.Read(readCtx) cancelRead() require.NoError(t, readErr) require.Equal(t, coderws.MessageText, msgType) return event } created := readEvent() delta := readEvent() completed := readEvent() require.Equal(t, "response.created", gjson.GetBytes(created, "type").String()) require.Equal(t, "response.output_text.delta", gjson.GetBytes(delta, "type").String()) require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String()) select { case bridge := <-resultCh: require.NoError(t, bridge.err) require.NotNil(t, bridge.result) require.Equal(t, "resp_bridge", bridge.result.RequestID) require.Equal(t, 3, bridge.result.Usage.InputTokens) require.Equal(t, 2, bridge.result.Usage.OutputTokens) require.True(t, bridge.result.OpenAIWSMode) case <-time.After(3 * time.Second): t.Fatal("timed out waiting for bridge result") } require.NotNil(t, upstream.lastReq) require.Equal(t, http.MethodPost, upstream.lastReq.Method) require.Equal(t, "true", upstream.lastReq.Header.Get(responsesLiteHeader)) require.False(t, gjson.GetBytes(upstream.lastBody, "type").Exists()) require.False(t, gjson.GetBytes(upstream.lastBody, "generate").Exists()) require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool()) } func TestProxyOpenAIWSHTTPBridgeTurnForGrokDefaultsEmptyModelTo45(t *testing.T) { gin.SetMode(gin.TestMode) upstream := &httpUpstreamRecorder{resp: &http.Response{ StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(strings.Join([]string{ `data: {"type":"response.created","response":{"id":"resp_grok_default","model":"grok-4.5"}}`, "", `data: {"type":"response.completed","response":{"id":"resp_grok_default","model":"grok-4.5","usage":{"input_tokens":1,"output_tokens":1}}}`, "", }, "\n"))), }} svc := &OpenAIGatewayService{ cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}, httpUpstream: upstream, } account := &Account{ ID: 72, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, Credentials: map[string]any{"base_url": xai.DefaultCLIBaseURL}, } payload := []byte(`{"type":"response.create","generate":true,"stream":true,"input":"hi"}`) recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) var events [][]byte result, err := svc.proxyOpenAIWSHTTPBridgeTurn( context.Background(), c, account, "access-token", payload, len(payload), "", "", "", "", "", 1, func(message []byte) error { events = append(events, append([]byte(nil), message...)) return nil }, ) require.NoError(t, err) require.NotNil(t, result) require.Equal(t, grokDefaultResponsesModel, gjson.GetBytes(upstream.lastBody, "model").String()) require.Len(t, events, 2) } func TestProxyOpenAIWSHTTPBridgeTurnPromotesCodexAdditionalToolsForMixedCache(t *testing.T) { gin.SetMode(gin.TestMode) upstream := &httpUpstreamRecorder{resp: &http.Response{ StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(strings.Join([]string{ `data: {"type":"response.created","response":{"id":"resp_grok_codex_lite","model":"grok-4.5"}}`, "", `data: {"type":"response.completed","response":{"id":"resp_grok_codex_lite","model":"grok-4.5","usage":{"input_tokens":4,"output_tokens":1}}}`, "", }, "\n"))), }} svc := &OpenAIGatewayService{ cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}, httpUpstream: upstream, } account := &Account{ ID: 73, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, Credentials: map[string]any{ "base_url": xai.DefaultCLIBaseURL, "subscription_tier": "free", }, } payload := []byte(`{ "type":"response.create","generate":true,"model":"grok","stream":true, "input":[ {"type":"additional_tools","role":"developer","tools":[ {"type":"function","name":"lookup","parameters":{"type":"object"}}, {"type":"function","name":"web_search","parameters":{"type":"object"}}, {"type":"custom","name":"apply_patch"}, {"type":"namespace","name":"collaboration"} ]}, {"type":"message","role":"user","content":[{"type":"input_text","text":"hello"}]} ] }`) recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) c.Request.Header.Set(grokClientToolCacheOptInHeader, "prefer-cache") var events [][]byte result, err := svc.proxyOpenAIWSHTTPBridgeTurn( context.Background(), c, account, "access-token", payload, len(payload), "grok", "", "", "", "isolated-ws-cache-id", 1, func(message []byte) error { events = append(events, append([]byte(nil), message...)) return nil }, ) require.NoError(t, err) require.NotNil(t, result) require.Len(t, events, 2) require.False(t, gjson.GetBytes(upstream.lastBody, `input.#(type=="additional_tools")`).Exists()) tools := gjson.GetBytes(upstream.lastBody, "tools").Array() require.Len(t, tools, 4) require.Equal(t, "function", tools[0].Get("type").String()) require.Equal(t, "lookup", tools[0].Get("name").String()) require.Equal(t, "web_search", tools[1].Get("type").String()) require.Equal(t, "function", tools[2].Get("type").String()) require.Equal(t, "apply_patch", tools[2].Get("name").String()) require.Equal(t, "x_search", tools[3].Get("type").String()) require.False(t, gjson.GetBytes(upstream.lastBody, "tool_choice").Exists()) require.False(t, gjson.GetBytes(upstream.lastBody, `tools.#(type=="custom")`).Exists()) require.False(t, gjson.GetBytes(upstream.lastBody, `tools.#(type=="namespace")`).Exists()) require.Equal(t, "isolated-ws-cache-id", gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String()) require.Equal(t, "isolated-ws-cache-id", upstream.lastReq.Header.Get(grokConversationIDHeader)) require.Empty(t, upstream.lastReq.Header.Get(grokClientToolCacheOptInHeader)) } func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridgeAndPreservesMappedModels(t *testing.T) { gin.SetMode(gin.TestMode) bridgeResponse := func(responseID, requestID string, cachedTokens int) *http.Response { sseBody := strings.Join([]string{ `data: {"type":"response.created","response":{"id":"` + responseID + `","model":"grok-4.3"}}`, "", `data: {"type":"response.output_text.delta","response":{"id":"` + responseID + `"},"delta":"ok"}`, "", `data: {"type":"response.completed","response":{"id":"` + responseID + `","model":"grok-4.3","usage":{"input_tokens":4,"output_tokens":2,"input_tokens_details":{"cached_tokens":` + fmt.Sprintf("%d", cachedTokens) + `}}}}`, "", }, "\n") return &http.Response{ StatusCode: http.StatusOK, Header: http.Header{ "Content-Type": []string{"text/event-stream"}, "Xai-Request-Id": []string{requestID}, }, Body: io.NopCloser(strings.NewReader(sseBody)), } } upstream := &httpUpstreamRecorder{responses: []*http.Response{ bridgeResponse("resp_grok_ws_1", "xai-ws-req-1", 0), bridgeResponse("resp_grok_ws_2", "xai-ws-req-2", 3), bridgeResponse("resp_grok_ws_3", "xai-ws-req-3", 0), }} svc := &OpenAIGatewayService{ cfg: &config.Config{ Gateway: config.GatewayConfig{ MaxLineSize: defaultMaxLineSize, }, }, httpUpstream: upstream, } account := &Account{ ID: 71, Name: "grok", Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, Status: StatusActive, Credentials: map[string]any{ "base_url": xai.DefaultCLIBaseURL, }, } errCh := make(chan error, 1) wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover}) if err != nil { errCh <- err return } defer func() { _ = conn.CloseNow() }() readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second) msgType, firstMessage, err := conn.Read(readCtx) cancelRead() if err != nil { errCh <- err return } if msgType != coderws.MessageText { errCh <- errors.New("first message was not text") return } rec := httptest.NewRecorder() ginCtx, _ := gin.CreateTestContext(rec) req := r.Clone(r.Context()) req.Header = req.Header.Clone() ginCtx.Request = req ginCtx.Set("api_key", &APIKey{ID: 7101}) 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() dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second) clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil) cancelDial() require.NoError(t, err) 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,"input":"hi","prompt_cache_retention":"24h"}`)) cancelWrite() require.NoError(t, err) readEvent := func() []byte { readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second) msgType, event, readErr := clientConn.Read(readCtx) cancelRead() require.NoError(t, readErr) require.Equal(t, coderws.MessageText, msgType) return event } created := readEvent() delta := readEvent() completed := readEvent() require.Equal(t, "response.created", gjson.GetBytes(created, "type").String()) 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_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":"channel-alias","stream":true,"previous_response_id":"resp_grok_ws_1","input":"second turn"}`)) cancelWrite() require.NoError(t, err) created = readEvent() delta = readEvent() completed = readEvent() require.Equal(t, "response.created", gjson.GetBytes(created, "type").String()) 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"}`)) cancelWrite() require.NoError(t, err) created = readEvent() delta = readEvent() completed = readEvent() require.Equal(t, "response.created", gjson.GetBytes(created, "type").String()) 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_3", gjson.GetBytes(completed, "response.id").String()) _ = clientConn.Close(coderws.StatusNormalClosure, "done") select { case proxyErr := <-errCh: require.NoError(t, proxyErr) case <-time.After(3 * time.Second): require.Fail(t, "proxy did not finish after client close") } require.Len(t, upstream.requests, 3) require.Len(t, upstream.bodies, 3) require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), 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.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)) require.Equal(t, "web_search", gjson.GetBytes(upstream.lastBody, "tools.0.type").String()) require.Equal(t, "x_search", gjson.GetBytes(upstream.lastBody, "tools.1.type").String()) require.Equal(t, "none", gjson.GetBytes(upstream.lastBody, "tool_choice").String()) firstIdentity := gjson.GetBytes(upstream.bodies[0], "prompt_cache_key").String() 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.NotEmpty(t, secondIdentity) require.NotEqual(t, firstIdentity, secondIdentity) require.NotEmpty(t, 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)) require.False(t, gjson.GetBytes(upstream.lastBody, "type").Exists()) require.False(t, gjson.GetBytes(upstream.lastBody, "generate").Exists()) require.False(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_retention").Exists()) } func TestOpenAIWSHTTPBridgeAcceptsFirstFrameAboveLegacy16MiB(t *testing.T) { gin.SetMode(gin.TestMode) sseBody := strings.Join([]string{ `data: {"type":"response.created","response":{"id":"resp_large_bridge","model":"gpt-5"}}`, "", `data: {"type":"response.completed","response":{"id":"resp_large_bridge","model":"gpt-5","usage":{"input_tokens":9,"output_tokens":1}}}`, "", }, "\n") upstream := &httpUpstreamRecorder{resp: &http.Response{ StatusCode: http.StatusOK, Header: http.Header{ "Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_large_bridge"}, }, Body: io.NopCloser(strings.NewReader(sseBody)), }} cfg := &config.Config{ Gateway: config.GatewayConfig{ MaxLineSize: defaultMaxLineSize, OpenAIWS: config.GatewayOpenAIWSConfig{ Enabled: true, APIKeyEnabled: true, ResponsesWebsocketsV2: true, ClientReadLimitBytes: 64 * 1024 * 1024, HTTPBridgeEnabled: true, HTTPBridgeThresholdBytes: 15 * 1024 * 1024, }, }, } svc := &OpenAIGatewayService{ cfg: cfg, httpUpstream: upstream, toolCorrector: NewCodexToolCorrector(), } account := &Account{ ID: 9, Name: "api-key", Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "sk-upstream"}, Extra: map[string]any{ "openai_apikey_responses_websockets_v2_enabled": true, }, Concurrency: 1, Status: StatusActive, } payload := []byte(`{"type":"response.create","generate":true,"model":"gpt-5","stream":true,"input":"` + strings.Repeat("x", 17*1024*1024) + `"}`) require.Greater(t, len(payload), 16*1024*1024) require.Less(t, int64(len(payload)), ResolveOpenAIWSClientReadLimitBytes(cfg)) errCh := make(chan error, 1) wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover}) if err != nil { errCh <- err return } defer func() { _ = conn.CloseNow() }() conn.SetReadLimit(ResolveOpenAIWSClientReadLimitBytes(cfg)) readCtx, cancelRead := context.WithTimeout(r.Context(), 10*time.Second) msgType, firstMessage, err := conn.Read(readCtx) cancelRead() if err != nil { errCh <- err return } if msgType != coderws.MessageText && msgType != coderws.MessageBinary { errCh <- NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "unexpected client websocket message type", nil) return } rec := httptest.NewRecorder() ginCtx, _ := gin.CreateTestContext(rec) req := r.Clone(r.Context()) req.Header = req.Header.Clone() req.Header.Set("User-Agent", "codex_cli_rs/0.135.0") ginCtx.Request = req proxyCtx, cancelProxy := context.WithTimeout(r.Context(), 20*time.Second) defer cancelProxy() errCh <- svc.ProxyResponsesWebSocketFromClient(proxyCtx, ginCtx, conn, account, "sk-test", firstMessage, nil) })) defer wsServer.Close() dialCtx, cancelDial := context.WithTimeout(context.Background(), 5*time.Second) clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil) cancelDial() require.NoError(t, err) defer func() { _ = clientConn.CloseNow() }() writeCtx, cancelWrite := context.WithTimeout(context.Background(), 20*time.Second) err = clientConn.Write(writeCtx, coderws.MessageText, payload) cancelWrite() require.NoError(t, err) var eventTypes []string for { readCtx, cancelRead := context.WithTimeout(context.Background(), 10*time.Second) msgType, event, readErr := clientConn.Read(readCtx) cancelRead() require.NoError(t, readErr) require.Equal(t, coderws.MessageText, msgType) eventType := gjson.GetBytes(event, "type").String() eventTypes = append(eventTypes, eventType) if eventType == "response.completed" { break } } require.Contains(t, eventTypes, "response.created") require.Contains(t, eventTypes, "response.completed") require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done")) select { case proxyErr := <-errCh: require.NoError(t, proxyErr) case <-time.After(10 * time.Second): t.Fatal("timed out waiting for websocket bridge proxy to finish") } require.NotNil(t, upstream.lastReq) require.Equal(t, http.MethodPost, upstream.lastReq.Method) require.Greater(t, len(upstream.lastBody), 16*1024*1024) require.False(t, gjson.GetBytes(upstream.lastBody, "type").Exists()) require.False(t, gjson.GetBytes(upstream.lastBody, "generate").Exists()) require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool()) require.Equal(t, "gpt-5", gjson.GetBytes(upstream.lastBody, "model").String()) } func TestOpenAIWSHTTPBridgeKeepsContinuationFramesOnHTTPWithoutPreviousResponseID(t *testing.T) { gin.SetMode(gin.TestMode) firstSSEBody := strings.Join([]string{ `data: {"type":"response.completed","response":{"id":"resp_bridge_first","model":"gpt-5.1","output":[{"type":"function_call","id":"fc_bridge_1","call_id":"call_bridge_1","name":"shell","arguments":"{}"}],"usage":{"input_tokens":9,"output_tokens":1}}}`, "", }, "\n") secondSSEBody := strings.Join([]string{ `data: {"type":"response.completed","response":{"id":"resp_bridge_second","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`, "", }, "\n") upstream := &httpUpstreamRecorder{responses: []*http.Response{ { StatusCode: http.StatusOK, Header: http.Header{ "Content-Type": []string{"text/event-stream"}, }, Body: io.NopCloser(strings.NewReader(firstSSEBody)), }, { StatusCode: http.StatusOK, Header: http.Header{ "Content-Type": []string{"text/event-stream"}, }, Body: io.NopCloser(strings.NewReader(secondSSEBody)), }, }} cfg := &config.Config{} cfg.Security.URLAllowlist.Enabled = false cfg.Security.URLAllowlist.AllowInsecureHTTP = true cfg.Gateway.OpenAIWS.Enabled = true cfg.Gateway.OpenAIWS.OAuthEnabled = true cfg.Gateway.OpenAIWS.APIKeyEnabled = true cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = true cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes = 1 cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1 cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8 cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3 cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3 cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3 captureConn := &openAIWSCaptureConn{} captureDialer := &openAIWSCaptureDialer{conn: captureConn} pool := newOpenAIWSConnPool(cfg) pool.setClientDialerForTest(captureDialer) svc := &OpenAIGatewayService{ cfg: cfg, httpUpstream: upstream, cache: &stubGatewayCache{}, openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), toolCorrector: NewCodexToolCorrector(), openaiWSPool: pool, } account := &Account{ ID: 19, Name: "api-key-bridge-handoff", Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "sk-upstream"}, Extra: map[string]any{ "responses_websockets_v2_enabled": true, }, Concurrency: 1, Status: StatusActive, Schedulable: true, } errCh := make(chan error, 1) wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover}) if err != nil { errCh <- err return } defer func() { _ = conn.CloseNow() }() readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second) msgType, firstMessage, err := conn.Read(readCtx) cancelRead() if err != nil { errCh <- err return } if msgType != coderws.MessageText && msgType != coderws.MessageBinary { errCh <- NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "unexpected client websocket message type", nil) return } rec := httptest.NewRecorder() ginCtx, _ := gin.CreateTestContext(rec) req := r.Clone(r.Context()) req.Header = req.Header.Clone() req.Header.Set("User-Agent", "codex_cli_rs/0.135.0") ginCtx.Request = req errCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "sk-test", firstMessage, nil) })) defer wsServer.Close() dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second) clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil) cancelDial() require.NoError(t, err) defer func() { _ = clientConn.CloseNow() }() writeMessage := func(payload string) { writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second) defer cancelWrite() require.NoError(t, clientConn.Write(writeCtx, coderws.MessageText, []byte(payload))) } readMessage := func() []byte { readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second) defer cancelRead() msgType, event, readErr := clientConn.Read(readCtx) require.NoError(t, readErr) require.Equal(t, coderws.MessageText, msgType) return event } writeMessage(`{"type":"response.create","model":"gpt-5.1","stream":true,"input":"first"}`) firstTurnEvent := readMessage() require.Equal(t, "response.completed", gjson.GetBytes(firstTurnEvent, "type").String()) require.Equal(t, "resp_bridge_first", gjson.GetBytes(firstTurnEvent, "response.id").String()) writeMessage(`{"type":"response.create","model":"gpt-5.1","stream":false,"previous_response_id":"resp_bridge_first","input":[{"type":"function_call_output","call_id":"call_bridge_1","output":"ok"}]}`) secondTurnEvent := readMessage() require.Equal(t, "response.completed", gjson.GetBytes(secondTurnEvent, "type").String()) require.Equal(t, "resp_bridge_second", gjson.GetBytes(secondTurnEvent, "response.id").String()) require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done")) select { case proxyErr := <-errCh: require.NoError(t, proxyErr) case <-time.After(5 * time.Second): t.Fatal("timed out waiting for websocket bridge proxy to finish") } require.Len(t, upstream.bodies, 2, "进入 HTTP bridge 后同一客户端 WS 连接内应保持 HTTP/SSE bridge") require.False(t, gjson.GetBytes(upstream.bodies[0], "previous_response_id").Exists()) require.False(t, gjson.GetBytes(upstream.bodies[1], "previous_response_id").Exists()) secondInput := gjson.GetBytes(upstream.bodies[1], "input").Array() require.Len(t, secondInput, 3) require.Equal(t, "first", secondInput[0].String()) require.Equal(t, "function_call", secondInput[1].Get("type").String()) require.Equal(t, "call_bridge_1", secondInput[1].Get("call_id").String()) require.Equal(t, "function_call_output", secondInput[2].Get("type").String()) require.Equal(t, "call_bridge_1", secondInput[2].Get("call_id").String()) require.Equal(t, 0, captureDialer.DialCount()) require.Empty(t, captureConn.writes) } func TestOpenAIWSHTTPBridge_IdleTimeoutClosesClientSession(t *testing.T) { gin.SetMode(gin.TestMode) sseBody := strings.Join([]string{ `data: {"type":"response.completed","response":{"id":"resp_bridge_idle","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`, "", }, "\n") upstream := &httpUpstreamRecorder{resp: &http.Response{ StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(sseBody)), }} cfg := &config.Config{} cfg.Security.URLAllowlist.Enabled = false cfg.Security.URLAllowlist.AllowInsecureHTTP = true cfg.Gateway.OpenAIWS.Enabled = true cfg.Gateway.OpenAIWS.APIKeyEnabled = true cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = true cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes = 1 cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds = 1 cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1 cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8 svc := &OpenAIGatewayService{ cfg: cfg, httpUpstream: upstream, cache: &stubGatewayCache{}, openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), toolCorrector: NewCodexToolCorrector(), } account := &Account{ ID: 20, Name: "api-key-bridge-idle-timeout", Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "sk-upstream"}, Extra: map[string]any{"responses_websockets_v2_enabled": true}, Concurrency: 1, Status: StatusActive, Schedulable: true, } errCh := make(chan error, 1) wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover}) if err != nil { errCh <- err return } defer func() { _ = conn.CloseNow() }() readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second) _, firstMessage, err := conn.Read(readCtx) cancelRead() if err != nil { errCh <- err return } rec := httptest.NewRecorder() ginCtx, _ := gin.CreateTestContext(rec) ginCtx.Request = r.Clone(r.Context()) errCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "sk-test", firstMessage, nil) })) defer wsServer.Close() dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second) clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil) cancelDial() require.NoError(t, err) defer func() { _ = clientConn.CloseNow() }() writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second) err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1","stream":false,"input":"hello"}`)) 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()) closeReadCtx, cancelCloseRead := context.WithTimeout(context.Background(), 3*time.Second) _, _, err = clientConn.Read(closeReadCtx) cancelCloseRead() var clientClose coderws.CloseError require.ErrorAs(t, err, &clientClose) require.Equal(t, coderws.StatusNormalClosure, clientClose.Code) require.Equal(t, "websocket idle timeout", clientClose.Reason) select { case proxyErr := <-errCh: var closeErr *OpenAIWSClientCloseError require.ErrorAs(t, proxyErr, &closeErr) require.Equal(t, coderws.StatusNormalClosure, closeErr.StatusCode()) require.Equal(t, "websocket idle timeout", closeErr.Reason()) case <-time.After(4 * time.Second): t.Fatal("timed out waiting for idle HTTP bridge session to close") } require.Len(t, upstream.bodies, 1, "an idle client must not leave a continuation request running") }