Merge pull request #5864 from wucm667/fix/issue-5850-http-bridge-replay
fix(openai): avoid duplicate HTTP bridge replay
This commit is contained in:
@@ -235,16 +235,16 @@ func AnalyzeToolCallOutputContextCoverageBytes(body []byte) ToolCallOutputContex
|
||||
return coverage
|
||||
}
|
||||
input := parseRawJSONView(body).Get("input")
|
||||
if !input.IsArray() {
|
||||
if !input.IsArray() && !input.IsObject() {
|
||||
return coverage
|
||||
}
|
||||
|
||||
missingCallID := false
|
||||
var outputCallIDs map[string]struct{}
|
||||
var contextIDs map[string]struct{}
|
||||
input.ForEach(func(_, item gjson.Result) bool {
|
||||
analyzeItem := func(item gjson.Result) {
|
||||
if !item.IsObject() {
|
||||
return true
|
||||
return
|
||||
}
|
||||
itemType := item.Get("type").String()
|
||||
switch {
|
||||
@@ -253,7 +253,7 @@ func AnalyzeToolCallOutputContextCoverageBytes(body []byte) ToolCallOutputContex
|
||||
callID := strings.TrimSpace(item.Get("call_id").String())
|
||||
if callID == "" {
|
||||
missingCallID = true
|
||||
return true
|
||||
return
|
||||
}
|
||||
if outputCallIDs == nil {
|
||||
outputCallIDs = make(map[string]struct{})
|
||||
@@ -262,7 +262,7 @@ func AnalyzeToolCallOutputContextCoverageBytes(body []byte) ToolCallOutputContex
|
||||
case isCodexToolCallContextItemType(itemType):
|
||||
callID := strings.TrimSpace(item.Get("call_id").String())
|
||||
if callID == "" {
|
||||
return true
|
||||
return
|
||||
}
|
||||
if contextIDs == nil {
|
||||
contextIDs = make(map[string]struct{})
|
||||
@@ -271,15 +271,22 @@ func AnalyzeToolCallOutputContextCoverageBytes(body []byte) ToolCallOutputContex
|
||||
case itemType == "item_reference":
|
||||
idValue := strings.TrimSpace(item.Get("id").String())
|
||||
if idValue == "" {
|
||||
return true
|
||||
return
|
||||
}
|
||||
if contextIDs == nil {
|
||||
contextIDs = make(map[string]struct{})
|
||||
}
|
||||
contextIDs[idValue] = struct{}{}
|
||||
}
|
||||
return true
|
||||
})
|
||||
}
|
||||
if input.IsArray() {
|
||||
input.ForEach(func(_, item gjson.Result) bool {
|
||||
analyzeItem(item)
|
||||
return true
|
||||
})
|
||||
} else {
|
||||
analyzeItem(input)
|
||||
}
|
||||
|
||||
if !coverage.HasFunctionCallOutput || missingCallID {
|
||||
return coverage
|
||||
|
||||
@@ -206,6 +206,14 @@ func TestAnalyzeToolCallOutputContextCoverageBytes(t *testing.T) {
|
||||
hasOutput: false,
|
||||
coversAllIDs: false,
|
||||
},
|
||||
{
|
||||
name: "object_tool_output_requires_context_replay",
|
||||
body: map[string]any{"input": map[string]any{
|
||||
"type": "custom_tool_call_output", "call_id": "call_a",
|
||||
}},
|
||||
hasOutput: true,
|
||||
coversAllIDs: false,
|
||||
},
|
||||
{
|
||||
name: "all_outputs_covered_by_context",
|
||||
body: map[string]any{"input": []any{
|
||||
|
||||
@@ -565,7 +565,9 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
}
|
||||
bridgePayloadRaw := currentBridgePayload.payloadRaw
|
||||
bridgePayloadBytes := currentBridgePayload.payloadBytes
|
||||
needsBridgeReplay := currentBridgePayload.previousResponseID != "" || openAIWSRawPayloadHasToolCallOutput(currentBridgePayload.payloadRaw)
|
||||
toolOutputCoverage := AnalyzeToolCallOutputContextCoverageBytes(currentBridgePayload.payloadRaw)
|
||||
needsBridgeReplay := currentBridgePayload.previousResponseID != "" ||
|
||||
(toolOutputCoverage.HasFunctionCallOutput && !toolOutputCoverage.ContextCoversAllCallIDs)
|
||||
turnReplayInput, turnReplayInputExists, replayInputErr := buildOpenAIWSReplayInputSequence(
|
||||
bridgeReplayInput,
|
||||
bridgeReplayInputExists,
|
||||
|
||||
@@ -809,6 +809,35 @@ func TestBuildOpenAIWSReplayInputSequence(t *testing.T) {
|
||||
require.Equal(t, "new", gjson.GetBytes(items[0], "text").String())
|
||||
})
|
||||
|
||||
t.Run("no_previous_response_id_custom_tool_history_does_not_accumulate", func(t *testing.T) {
|
||||
previousFull := []json.RawMessage{
|
||||
json.RawMessage(`{"type":"input_text","text":"stale"}`),
|
||||
json.RawMessage(`{"type":"custom_tool_call","id":"stale_item","call_id":"stale_call","name":"exec","input":"stale"}`),
|
||||
}
|
||||
currentPayload := []byte(`{"input":[
|
||||
{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"},
|
||||
{"type":"custom_tool_call_output","call_id":"call_1","output":"/tmp"},
|
||||
{"type":"input_text","text":"continue"}
|
||||
]}`)
|
||||
|
||||
for range 3 {
|
||||
items, exists, err := buildOpenAIWSReplayInputSequence(
|
||||
previousFull,
|
||||
true,
|
||||
currentPayload,
|
||||
false,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.True(t, exists)
|
||||
require.Len(t, items, 3)
|
||||
require.Equal(t, "custom_tool_call", gjson.GetBytes(items[0], "type").String())
|
||||
require.Equal(t, "call_1", gjson.GetBytes(items[0], "call_id").String())
|
||||
require.Equal(t, "custom_tool_call_output", gjson.GetBytes(items[1], "type").String())
|
||||
require.Equal(t, "call_1", gjson.GetBytes(items[1], "call_id").String())
|
||||
previousFull = append(items, json.RawMessage(`{"type":"custom_tool_call","id":"replayed_item","call_id":"replayed_call","name":"exec","input":"ignored"}`))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("previous_response_id_delta_append", func(t *testing.T) {
|
||||
items, exists, err := buildOpenAIWSReplayInputSequence(
|
||||
lastFull,
|
||||
@@ -823,6 +852,91 @@ func TestBuildOpenAIWSReplayInputSequence(t *testing.T) {
|
||||
require.Equal(t, "world", gjson.GetBytes(items[1], "text").String())
|
||||
})
|
||||
|
||||
t.Run("previous_response_id_filters_orphan_historical_custom_tool_call", func(t *testing.T) {
|
||||
previousFull := []json.RawMessage{
|
||||
json.RawMessage(`{"type":"input_text","text":"hello"}`),
|
||||
json.RawMessage(`{"type":"custom_tool_call","id":"item_orphan","call_id":"call_orphan","name":"exec","input":"pwd"}`),
|
||||
}
|
||||
items, exists, err := buildOpenAIWSReplayInputSequence(
|
||||
previousFull,
|
||||
true,
|
||||
[]byte(`{"previous_response_id":"resp_1","input":[{"role":"user","content":"continue"}]}`),
|
||||
true,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.True(t, exists)
|
||||
require.Len(t, items, 2)
|
||||
require.Equal(t, "hello", gjson.GetBytes(items[0], "text").String())
|
||||
require.Equal(t, "user", gjson.GetBytes(items[1], "role").String())
|
||||
})
|
||||
|
||||
t.Run("previous_response_id_preserves_paired_historical_function_call", func(t *testing.T) {
|
||||
previousFull := []json.RawMessage{
|
||||
json.RawMessage(`{"type":"function_call","id":"item_1","call_id":"call_1","name":"lookup","arguments":"{}"}`),
|
||||
json.RawMessage(`{"type":"function_call_output","call_id":"call_1","output":"ok"}`),
|
||||
}
|
||||
items, exists, err := buildOpenAIWSReplayInputSequence(
|
||||
previousFull,
|
||||
true,
|
||||
[]byte(`{"previous_response_id":"resp_1","input":[{"role":"user","content":"continue"}]}`),
|
||||
true,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.True(t, exists)
|
||||
require.Len(t, items, 3)
|
||||
require.Equal(t, "function_call", gjson.GetBytes(items[0], "type").String())
|
||||
require.Equal(t, "function_call_output", gjson.GetBytes(items[1], "type").String())
|
||||
})
|
||||
|
||||
t.Run("previous_response_id_preserves_paired_historical_custom_tool_call", func(t *testing.T) {
|
||||
previousFull := []json.RawMessage{
|
||||
json.RawMessage(`{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"}`),
|
||||
json.RawMessage(`{"type":"custom_tool_call_output","call_id":"call_1","output":"/tmp"}`),
|
||||
}
|
||||
items, exists, err := buildOpenAIWSReplayInputSequence(
|
||||
previousFull,
|
||||
true,
|
||||
[]byte(`{"previous_response_id":"resp_1","input":[{"role":"user","content":"continue"}]}`),
|
||||
true,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.True(t, exists)
|
||||
require.Len(t, items, 3)
|
||||
require.Equal(t, "custom_tool_call", gjson.GetBytes(items[0], "type").String())
|
||||
require.Equal(t, "custom_tool_call_output", gjson.GetBytes(items[1], "type").String())
|
||||
})
|
||||
|
||||
t.Run("item_reference_does_not_complete_historical_call", func(t *testing.T) {
|
||||
previousFull := []json.RawMessage{
|
||||
json.RawMessage(`{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"}`),
|
||||
}
|
||||
items, exists, err := buildOpenAIWSReplayInputSequence(
|
||||
previousFull,
|
||||
true,
|
||||
[]byte(`{"previous_response_id":"resp_1","input":[{"type":"item_reference","id":"call_1"},{"role":"user","content":"continue"}]}`),
|
||||
true,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.True(t, exists)
|
||||
require.Len(t, items, 2)
|
||||
require.Equal(t, "item_reference", gjson.GetBytes(items[0], "type").String())
|
||||
require.Equal(t, "user", gjson.GetBytes(items[1], "role").String())
|
||||
})
|
||||
|
||||
t.Run("previous_response_id_preserves_current_orphan_custom_tool_call", func(t *testing.T) {
|
||||
items, exists, err := buildOpenAIWSReplayInputSequence(
|
||||
lastFull,
|
||||
true,
|
||||
[]byte(`{"previous_response_id":"resp_1","input":[{"type":"custom_tool_call","id":"item_live","call_id":"call_live","name":"exec","input":"pwd"}]}`),
|
||||
true,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.True(t, exists)
|
||||
require.Len(t, items, 2)
|
||||
require.Equal(t, "custom_tool_call", gjson.GetBytes(items[1], "type").String())
|
||||
require.Equal(t, "call_live", gjson.GetBytes(items[1], "call_id").String())
|
||||
})
|
||||
|
||||
t.Run("previous_response_id_full_input_replace", func(t *testing.T) {
|
||||
items, exists, err := buildOpenAIWSReplayInputSequence(
|
||||
lastFull,
|
||||
|
||||
@@ -610,6 +610,40 @@ func openAIWSRawItemsHaveToolCallContextForOutputs(items []json.RawMessage) bool
|
||||
return true
|
||||
}
|
||||
|
||||
func sanitizeOpenAIWSHistoricalReplayToolCalls(
|
||||
previousItems []json.RawMessage,
|
||||
currentItems []json.RawMessage,
|
||||
) []json.RawMessage {
|
||||
if len(previousItems) == 0 {
|
||||
return cloneOpenAIWSRawMessages(previousItems)
|
||||
}
|
||||
outputCallIDs := make(map[string]struct{})
|
||||
collectOutputCallIDs := func(items []json.RawMessage) {
|
||||
for _, item := range items {
|
||||
if !isCodexToolCallOutputItemType(gjson.GetBytes(item, "type").String()) {
|
||||
continue
|
||||
}
|
||||
if callID := strings.TrimSpace(gjson.GetBytes(item, "call_id").String()); callID != "" {
|
||||
outputCallIDs[callID] = struct{}{}
|
||||
}
|
||||
}
|
||||
}
|
||||
collectOutputCallIDs(previousItems)
|
||||
collectOutputCallIDs(currentItems)
|
||||
|
||||
sanitized := make([]json.RawMessage, 0, len(previousItems))
|
||||
for _, item := range previousItems {
|
||||
if isCodexToolCallContextItemType(gjson.GetBytes(item, "type").String()) {
|
||||
callID := strings.TrimSpace(gjson.GetBytes(item, "call_id").String())
|
||||
if _, paired := outputCallIDs[callID]; !paired {
|
||||
continue
|
||||
}
|
||||
}
|
||||
sanitized = append(sanitized, append(json.RawMessage(nil), item...))
|
||||
}
|
||||
return sanitized
|
||||
}
|
||||
|
||||
func openAIWSRawPayloadHasToolCallOutput(payload []byte) bool {
|
||||
if len(payload) == 0 {
|
||||
return false
|
||||
@@ -648,6 +682,7 @@ func buildOpenAIWSReplayInputSequence(
|
||||
if !previousFullInputExists {
|
||||
return cloneOpenAIWSRawMessages(currentItems), currentExists, nil
|
||||
}
|
||||
previousFullInput = sanitizeOpenAIWSHistoricalReplayToolCalls(previousFullInput, currentItems)
|
||||
if !currentExists || len(currentItems) == 0 {
|
||||
return cloneOpenAIWSRawMessages(previousFullInput), true, nil
|
||||
}
|
||||
|
||||
@@ -416,6 +416,197 @@ func TestOpenAIWSHTTPBridgeAPIKeyReusesClientToolMappingWhenFollowupOmitsTools(t
|
||||
require.False(t, secondInput[2].Get("id").Exists())
|
||||
}
|
||||
|
||||
func TestOpenAIWSHTTPBridgeFullCustomToolHistoryWithoutPreviousResponseIDDoesNotReplay(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
completed := func(responseID string, output string) string {
|
||||
return "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"" + responseID + "\",\"model\":\"gpt-5.1\",\"output\":" + output + ",\"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(completed("resp_1", `[{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"}]`)))},
|
||||
{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(completed("resp_2", `[]`)))},
|
||||
{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(completed("resp_3", `[]`)))},
|
||||
{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(completed("resp_4", `[]`)))},
|
||||
}}
|
||||
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.ResponsesWebsocketsV2 = true
|
||||
cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = true
|
||||
cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes = 1
|
||||
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: 9002, Name: "oauth-full-context", Platform: PlatformOpenAI, Type: AccountTypeOAuth,
|
||||
Credentials: map[string]any{"access_token": "test-token"}, 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, "test-token", 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() }()
|
||||
|
||||
writeAndRead := func(payload string) {
|
||||
writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
require.NoError(t, clientConn.Write(writeCtx, coderws.MessageText, []byte(payload)))
|
||||
cancelWrite()
|
||||
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
_, event, readErr := clientConn.Read(readCtx)
|
||||
cancelRead()
|
||||
require.NoError(t, readErr)
|
||||
require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String())
|
||||
}
|
||||
|
||||
writeAndRead(`{"type":"response.create","model":"gpt-5.1","input":"run pwd"}`)
|
||||
writeAndRead(`{"type":"response.create","model":"gpt-5.1","previous_response_id":"resp_1","input":[{"role":"user","content":"continue without tool output"}]}`)
|
||||
fullContext := `{"type":"response.create","model":"gpt-5.1","input":[{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"},{"type":"custom_tool_call_output","call_id":"call_1","output":"/tmp"},{"role":"user","content":"continue"}]}`
|
||||
writeAndRead(fullContext)
|
||||
writeAndRead(fullContext)
|
||||
|
||||
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, 4)
|
||||
orphanInput := gjson.GetBytes(upstream.bodies[1], "input").Array()
|
||||
require.Len(t, orphanInput, 2)
|
||||
require.Equal(t, "run pwd", orphanInput[0].String())
|
||||
require.Equal(t, "user", orphanInput[1].Get("role").String())
|
||||
for _, body := range upstream.bodies[2:] {
|
||||
input := gjson.GetBytes(body, "input").Array()
|
||||
require.Len(t, input, 3)
|
||||
require.Equal(t, "custom_tool_call", input[0].Get("type").String())
|
||||
require.Equal(t, "call_1", input[0].Get("call_id").String())
|
||||
require.Equal(t, "custom_tool_call_output", input[1].Get("type").String())
|
||||
require.Equal(t, "call_1", input[1].Get("call_id").String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIWSHTTPBridgeObjectToolOutputWithoutPreviousResponseIDReplaysMatchingCall(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
completed := func(responseID string, output string) string {
|
||||
return "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"" + responseID + "\",\"model\":\"gpt-5.1\",\"output\":" + output + ",\"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(completed("resp_1", `[{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"}]`)))},
|
||||
{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(completed("resp_2", `[]`)))},
|
||||
}}
|
||||
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.ResponsesWebsocketsV2 = true
|
||||
cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = true
|
||||
cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes = 1
|
||||
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: 9003, Name: "oauth-output-only", Platform: PlatformOpenAI, Type: AccountTypeOAuth,
|
||||
Credentials: map[string]any{"access_token": "test-token"}, 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, "test-token", 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() }()
|
||||
|
||||
writeAndRead := func(payload string) {
|
||||
writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
require.NoError(t, clientConn.Write(writeCtx, coderws.MessageText, []byte(payload)))
|
||||
cancelWrite()
|
||||
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
_, event, readErr := clientConn.Read(readCtx)
|
||||
cancelRead()
|
||||
require.NoError(t, readErr)
|
||||
require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String())
|
||||
}
|
||||
|
||||
writeAndRead(`{"type":"response.create","model":"gpt-5.1","input":"run pwd"}`)
|
||||
writeAndRead(`{"type":"response.create","model":"gpt-5.1","input":{"type":"custom_tool_call_output","call_id":"call_1","output":"/tmp"}}`)
|
||||
|
||||
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)
|
||||
secondInput := gjson.GetBytes(upstream.bodies[1], "input").Array()
|
||||
require.Len(t, secondInput, 3)
|
||||
require.Equal(t, "custom_tool_call", secondInput[1].Get("type").String())
|
||||
require.Equal(t, "call_1", secondInput[1].Get("call_id").String())
|
||||
require.Equal(t, "custom_tool_call_output", secondInput[2].Get("type").String())
|
||||
require.Equal(t, "call_1", secondInput[2].Get("call_id").String())
|
||||
}
|
||||
|
||||
func TestOpenAIWSHTTPBridgeDecisionKeepsSmallFramesOnWS(t *testing.T) {
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: &config.Config{
|
||||
|
||||
Reference in New Issue
Block a user