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

1605 lines
65 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package service
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_errorCodex 对 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")
}