|
|
|
@@ -413,6 +413,191 @@ 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", `[]`)))},
|
|
|
|
|
}}
|
|
|
|
|
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"}`)
|
|
|
|
|
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, 3)
|
|
|
|
|
for _, body := range upstream.bodies[1:] {
|
|
|
|
|
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{
|
|
|
|
|