diff --git a/config.example.yaml b/config.example.yaml index 9b7dd443..5020eaef 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -222,7 +222,7 @@ codex: identity-confuse: false # Disable forcing the official Codex User-Agent and Originator headers on HTTP/SSE and WebSocket requests. disable-codex-cloaking: false - # When true, optimize Codex Desktop and codex-tui requests for multi-agent v2. + # When true, optimize Codex Desktop, codex-tui, and codex_cli_rs requests for multi-agent v2. # This refreshes Codex spawn_agent model details, removes message parameter encryption, # normalizes encrypted agent_message content for Codex, and converts agent_message input # into standard user messages for non-Codex upstream protocols. diff --git a/internal/client/codex/optimize-multi-agent-v2/optimize_multi_agent_v2.go b/internal/client/codex/optimize-multi-agent-v2/optimize_multi_agent_v2.go index 3b20ecd6..fbfde007 100644 --- a/internal/client/codex/optimize-multi-agent-v2/optimize_multi_agent_v2.go +++ b/internal/client/codex/optimize-multi-agent-v2/optimize_multi_agent_v2.go @@ -27,6 +27,10 @@ const ( codexOptimizedCollaborationNamePrefix = codexOptimizedCollaborationNamespace + "__" ) +// CodexMultiAgentV2ToolsPreparedContextKey marks a request whose collaboration +// tool definitions were prepared at the Responses API boundary. +const CodexMultiAgentV2ToolsPreparedContextKey = "codex_multi_agent_v2_tools_prepared" + // codexCollaborationMessageTools are the collaboration tool names whose // parameters.properties.message.encrypted field must be stripped so that // message content remains readable by the proxy. @@ -75,6 +79,24 @@ func TranslateRequestWithCodexMultiAgentV2(ctx context.Context, headers http.Hea return sdktranslator.TranslateRequest(from, to, model, payload, stream) } +// PrepareCodexMultiAgentV2Tools prepares collaboration tool definitions at the +// Responses API boundary without changing the collaboration namespace. +func PrepareCodexMultiAgentV2Tools(ctx context.Context, headers http.Header, payload []byte, enabled, homeEnabled bool) ([]byte, bool) { + if !codexMultiAgentV2ClientEnabled(ctx, headers, enabled) { + return payload, false + } + + updated := removeCodexCollaborationMessageEncryption(payload, codexCollaborationMessageToolPaths(payload)) + toolPaths := codexSpawnAgentToolPaths(updated) + if len(toolPaths) == 0 || hasCodexOptimizedCollaborationConflict(updated) { + return updated, true + } + + models := codexSpawnAgentModelsForRequest(ctx, headers, homeEnabled) + updated = rewriteCodexSpawnAgentTools(updated, toolPaths, models) + return updated, true +} + // OptimizeCodexMultiAgentV2Request rewrites an eligible spawn_agent request and // reports whether the collaboration namespace was renamed for upstream use. func OptimizeCodexMultiAgentV2Request(ctx context.Context, headers http.Header, payload []byte, cfg *config.Config) ([]byte, bool) { @@ -82,18 +104,37 @@ func OptimizeCodexMultiAgentV2Request(ctx context.Context, headers http.Header, return payload, false } updated := rewriteCodexAgentMessageContent(payload) - updated = removeCodexCollaborationMessageEncryption(updated, codexCollaborationMessageToolPaths(updated)) + if codexMultiAgentV2ToolsPrepared(ctx) { + updated = removeCodexCollaborationMessageEncryption(updated, codexCollaborationMessageToolPaths(updated)) + } else { + updated, _ = PrepareCodexMultiAgentV2Tools(ctx, headers, updated, cfg.Codex.OptimizeMultiAgentV2, cfg.Home.Enabled) + } toolPaths := codexSpawnAgentToolPaths(updated) if len(toolPaths) == 0 || hasCodexOptimizedCollaborationConflict(updated) { return updated, false } - models := codexSpawnAgentModelsForRequest(ctx, headers, cfg.Home.Enabled) - updated = rewriteCodexSpawnAgentTools(updated, toolPaths, models) return optimizeCodexCollaborationNamespace(updated, toolPaths) } func codexMultiAgentV2Enabled(ctx context.Context, headers http.Header, cfg *config.Config) bool { - return cfg != nil && cfg.Codex.OptimizeMultiAgentV2 && isCodexMultiAgentClient(codexClientUserAgent(ctx, headers)) + return cfg != nil && codexMultiAgentV2ClientEnabled(ctx, headers, cfg.Codex.OptimizeMultiAgentV2) +} + +func codexMultiAgentV2ClientEnabled(ctx context.Context, headers http.Header, enabled bool) bool { + return enabled && isCodexMultiAgentClient(codexClientUserAgent(ctx, headers)) +} + +func codexMultiAgentV2ToolsPrepared(ctx context.Context) bool { + if ctx == nil { + return false + } + ginCtx, ok := ctx.Value("gin").(*gin.Context) + if !ok || ginCtx == nil { + return false + } + prepared, ok := ginCtx.Get(CodexMultiAgentV2ToolsPreparedContextKey) + isPrepared, _ := prepared.(bool) + return ok && isPrepared } func codexClientUserAgent(ctx context.Context, headers http.Header) string { @@ -127,7 +168,10 @@ func headerValueCaseInsensitive(headers http.Header, name string) string { func isCodexMultiAgentClient(userAgent string) bool { userAgent = strings.TrimSpace(userAgent) - return strings.HasPrefix(userAgent, "Codex Desktop/") || strings.HasPrefix(userAgent, "codex-tui/") + return strings.HasPrefix(userAgent, "Codex Desktop/") || + strings.HasPrefix(userAgent, "codex-tui/") || + userAgent == "codex_cli_rs" || + strings.HasPrefix(userAgent, "codex_cli_rs/") } func codexSpawnAgentModelsForRequest(ctx context.Context, headers http.Header, homeEnabled bool) []codexSpawnAgentModel { diff --git a/internal/client/codex/optimize-multi-agent-v2/optimize_multi_agent_v2_test.go b/internal/client/codex/optimize-multi-agent-v2/optimize_multi_agent_v2_test.go index 112c3265..01a64484 100644 --- a/internal/client/codex/optimize-multi-agent-v2/optimize_multi_agent_v2_test.go +++ b/internal/client/codex/optimize-multi-agent-v2/optimize_multi_agent_v2_test.go @@ -33,6 +33,16 @@ func TestIsCodexMultiAgentClient(t *testing.T) { userAgent: "codex-tui/0.145.0 (Mac OS 26.5.2; arm64) iTerm.app/3.6.11 (codex-tui; 0.145.0)", want: true, }, + { + name: "codex cli rs", + userAgent: "codex_cli_rs/0.144.1 (Mac OS 26.3.1; arm64) iTerm.app/3.6.9", + want: true, + }, + { + name: "bare codex cli rs", + userAgent: "codex_cli_rs", + want: true, + }, { name: "other client", userAgent: "curl/8.7.1", @@ -309,6 +319,64 @@ func TestRewriteCodexSpawnAgentDescriptionEnabledOptimizesTool(t *testing.T) { } } +func TestPrepareCodexMultiAgentV2ToolsOnlyPreparesToolDefinitions(t *testing.T) { + t.Parallel() + + payload := []byte(`{ + "input":[ + {"type":"agent_message","content":[{"type":"encrypted_content","encrypted_content":"task"}]}, + {"type":"additional_tools","role":"developer","tools":[ + {"type":"namespace","name":"collaboration","tools":[ + {"type":"function","name":"spawn_agent","description":"Spawns an agent.","parameters":{"properties":{"message":{"encrypted":true}}}}, + {"type":"function","name":"send_message","parameters":{"properties":{"message":{"encrypted":true}}}} + ]} + ]} + ] + }`) + headers := http.Header{"User-Agent": []string{"codex_cli_rs/0.144.1"}} + got, prepared := PrepareCodexMultiAgentV2Tools(context.Background(), headers, payload, true, false) + if !prepared { + t.Fatal("Codex CLI request was not marked prepared") + } + if messageType := gjson.GetBytes(got, "input.0.content.0.type").String(); messageType != "encrypted_content" { + t.Fatalf("agent_message content type = %q, want encrypted_content", messageType) + } + if namespace := gjson.GetBytes(got, "input.1.tools.0.name").String(); namespace != codexCollaborationNamespace { + t.Fatalf("namespace = %q, want %q", namespace, codexCollaborationNamespace) + } + for _, path := range []string{"input.1.tools.0.tools.0", "input.1.tools.0.tools.1"} { + if encrypted := gjson.GetBytes(got, path+".parameters.properties.message.encrypted"); encrypted.Exists() { + t.Fatalf("%s message.encrypted was not removed: %s", path, encrypted.Raw) + } + } +} + +func TestOptimizeCodexMultiAgentV2RequestSkipsPreparedToolRefresh(t *testing.T) { + t.Parallel() + + request := httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + request.Header.Set("User-Agent", "codex_cli_rs/0.144.1") + ginContext, _ := gin.CreateTestContext(httptest.NewRecorder()) + ginContext.Request = request + ginContext.Set(CodexMultiAgentV2ToolsPreparedContextKey, true) + ctx := context.WithValue(context.Background(), "gin", ginContext) + + payload := []byte(`{"tools":[{"type":"namespace","name":"collaboration","tools":[{"type":"function","name":"spawn_agent","description":"Available model overrides (optional; inherited parent model is preferred): +- old-model: Old model. +Spawns an agent.","parameters":{"properties":{"message":{"encrypted":true}}}}]}]}`) + cfg := &config.Config{Codex: config.CodexConfig{OptimizeMultiAgentV2: true}} + got, optimized := OptimizeCodexMultiAgentV2Request(ctx, nil, payload, cfg) + if !optimized { + t.Fatal("collaboration namespace was not optimized") + } + if description := gjson.GetBytes(got, "tools.0.tools.0.description").String(); !strings.Contains(description, "old-model") { + t.Fatalf("prepared spawn_agent description was refreshed: %q", description) + } + if encrypted := gjson.GetBytes(got, "tools.0.tools.0.parameters.properties.message.encrypted"); encrypted.Exists() { + t.Fatalf("message.encrypted was not removed: %s", encrypted.Raw) + } +} + func TestOptimizeCodexMultiAgentV2RequestNormalizesAgentMessageContentOnly(t *testing.T) { t.Parallel() diff --git a/sdk/api/handlers/openai/openai_responses_handlers.go b/sdk/api/handlers/openai/openai_responses_handlers.go index e9063b86..cb45b959 100644 --- a/sdk/api/handlers/openai/openai_responses_handlers.go +++ b/sdk/api/handlers/openai/openai_responses_handlers.go @@ -16,6 +16,7 @@ import ( "sort" "github.com/gin-gonic/gin" + "github.com/router-for-me/CLIProxyAPI/v7/internal/client/codex/optimize-multi-agent-v2" . "github.com/router-for-me/CLIProxyAPI/v7/internal/constant" "github.com/router-for-me/CLIProxyAPI/v7/internal/interfaces" "github.com/router-for-me/CLIProxyAPI/v7/internal/registry" @@ -363,6 +364,35 @@ func (h *OpenAIResponsesAPIHandler) OpenAIResponsesModels(c *gin.Context) { }) } +func (h *OpenAIResponsesAPIHandler) prepareCodexMultiAgentV2Tools(c *gin.Context, payload []byte) []byte { + if h == nil || h.Cfg == nil { + return payload + } + + requestCtx := context.Background() + if c != nil && c.Request != nil { + requestCtx = c.Request.Context() + } + requestCtx = context.WithValue(requestCtx, "gin", c) + + var requestHeaders http.Header + if c != nil && c.Request != nil { + requestHeaders = c.Request.Header + } + homeEnabled := h.AuthManager != nil && h.AuthManager.HomeEnabled() + updated, prepared := multiagentv2.PrepareCodexMultiAgentV2Tools( + requestCtx, + requestHeaders, + payload, + h.Cfg.CodexOptimizeMultiAgentV2, + homeEnabled, + ) + if prepared && c != nil { + c.Set(multiagentv2.CodexMultiAgentV2ToolsPreparedContextKey, true) + } + return updated +} + // Responses handles the /v1/responses endpoint. // It determines whether the request is for a streaming or non-streaming response // and calls the appropriate handler based on the model provider. @@ -382,6 +412,8 @@ func (h *OpenAIResponsesAPIHandler) Responses(c *gin.Context) { return } + rawJSON = h.prepareCodexMultiAgentV2Tools(c, rawJSON) + // Check if the client requested a streaming response. streamResult := gjson.GetBytes(rawJSON, "stream") if streamResult.Type == gjson.True { diff --git a/sdk/api/handlers/openai/openai_responses_multi_agent_test.go b/sdk/api/handlers/openai/openai_responses_multi_agent_test.go new file mode 100644 index 00000000..0eae0953 --- /dev/null +++ b/sdk/api/handlers/openai/openai_responses_multi_agent_test.go @@ -0,0 +1,201 @@ +package openai + +import ( + "bytes" + "context" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/gin-gonic/gin" + "github.com/gorilla/websocket" + multiagentv2 "github.com/router-for-me/CLIProxyAPI/v7/internal/client/codex/optimize-multi-agent-v2" + "github.com/router-for-me/CLIProxyAPI/v7/internal/registry" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/api/handlers" + coreauth "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" + coreexecutor "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/executor" + sdkconfig "github.com/router-for-me/CLIProxyAPI/v7/sdk/config" + "github.com/tidwall/gjson" +) + +func TestPrepareCodexMultiAgentV2ToolsAtResponsesBoundary(t *testing.T) { + t.Parallel() + + gin.SetMode(gin.TestMode) + base := handlers.NewBaseAPIHandlers(&sdkconfig.SDKConfig{CodexOptimizeMultiAgentV2: true}, nil) + handler := NewOpenAIResponsesAPIHandler(base) + request := httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + request.Header.Set("User-Agent", "codex_cli_rs/0.144.1") + ginContext, _ := gin.CreateTestContext(httptest.NewRecorder()) + ginContext.Request = request + + payload := []byte(`{ + "tools":[{"type":"namespace","name":"collaboration","tools":[ + {"type":"function","name":"spawn_agent","description":"Spawns an agent.","parameters":{"properties":{"message":{"encrypted":true}}}}, + {"type":"function","name":"send_message","parameters":{"properties":{"message":{"encrypted":true}}}} + ]}] + }`) + got := handler.prepareCodexMultiAgentV2Tools(ginContext, payload) + + if namespace := gjson.GetBytes(got, "tools.0.name").String(); namespace != "collaboration" { + t.Fatalf("namespace = %q, want collaboration", namespace) + } + for _, path := range []string{"tools.0.tools.0", "tools.0.tools.1"} { + if encrypted := gjson.GetBytes(got, path+".parameters.properties.message.encrypted"); encrypted.Exists() { + t.Fatalf("%s message.encrypted was not removed: %s", path, encrypted.Raw) + } + } + prepared, exists := ginContext.Get(multiagentv2.CodexMultiAgentV2ToolsPreparedContextKey) + if !exists || prepared != true { + t.Fatalf("prepared marker = %#v, want true", prepared) + } +} + +func TestResponsesPreparesCodexMultiAgentV2ToolsForHTTPAndSSE(t *testing.T) { + t.Parallel() + + for _, stream := range []bool{false, true} { + t.Run(fmt.Sprintf("stream=%t", stream), func(t *testing.T) { + executor := &responsesMultiAgentCaptureExecutor{} + handler, modelID := newResponsesMultiAgentTestHandler(t, executor) + router := gin.New() + router.POST("/v1/responses", handler.Responses) + + payload := fmt.Sprintf(`{"model":%q,"stream":%t,"tools":[{"type":"namespace","name":"collaboration","tools":[{"type":"function","name":"spawn_agent","description":"Spawns an agent.","parameters":{"properties":{"message":{"encrypted":true}}}}]}]}`, modelID, stream) + request := httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewBufferString(payload)) + request.Header.Set("User-Agent", "codex_cli_rs/0.144.1") + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, request) + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, want 200; body=%s", recorder.Code, recorder.Body.String()) + } + + payloads := executor.Payloads() + if len(payloads) != 1 { + t.Fatalf("captured payload count = %d, want 1", len(payloads)) + } + captured := payloads[0] + if encrypted := gjson.GetBytes(captured, "tools.0.tools.0.parameters.properties.message.encrypted"); encrypted.Exists() { + t.Fatalf("message.encrypted was not removed: %s", captured) + } + if namespace := gjson.GetBytes(captured, "tools.0.name").String(); namespace != "collaboration" { + t.Fatalf("namespace = %q, want collaboration", namespace) + } + }) + } +} + +type responsesMultiAgentCaptureExecutor struct { + websocketDirectCaptureExecutor +} + +func (e *responsesMultiAgentCaptureExecutor) Execute(_ context.Context, _ *coreauth.Auth, req coreexecutor.Request, _ coreexecutor.Options) (coreexecutor.Response, error) { + e.mu.Lock() + e.payloads = append(e.payloads, bytes.Clone(req.Payload)) + e.mu.Unlock() + return coreexecutor.Response{Payload: []byte(`{"id":"resp-1","output":[]}`)}, nil +} + +func (e *responsesMultiAgentCaptureExecutor) ExecuteStream(_ context.Context, _ *coreauth.Auth, req coreexecutor.Request, _ coreexecutor.Options) (*coreexecutor.StreamResult, error) { + e.mu.Lock() + e.payloads = append(e.payloads, bytes.Clone(req.Payload)) + e.mu.Unlock() + chunks := make(chan coreexecutor.StreamChunk, 1) + chunks <- coreexecutor.StreamChunk{Payload: []byte("data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp-1\",\"output\":[]}}\n\n")} + close(chunks) + return &coreexecutor.StreamResult{Chunks: chunks}, nil +} + +func newResponsesMultiAgentTestHandler(t *testing.T, executor *responsesMultiAgentCaptureExecutor) (*OpenAIResponsesAPIHandler, string) { + t.Helper() + + modelID := "responses-multi-agent-test-model" + authID := "responses-multi-agent-test-auth" + manager := coreauth.NewManager(nil, nil, nil) + manager.RegisterExecutor(executor) + auth := &coreauth.Auth{ID: authID, Provider: "codex", Status: coreauth.StatusActive, ProxyURL: "direct"} + if _, errRegister := manager.Register(context.Background(), auth); errRegister != nil { + t.Fatalf("Register auth: %v", errRegister) + } + registry.GetGlobalRegistry().RegisterClient(authID, auth.Provider, []*registry.ModelInfo{{ID: modelID}}) + t.Cleanup(func() { + registry.GetGlobalRegistry().UnregisterClient(authID) + }) + + base := handlers.NewBaseAPIHandlers(&sdkconfig.SDKConfig{CodexOptimizeMultiAgentV2: true}, manager) + return NewOpenAIResponsesAPIHandler(base), modelID +} + +func TestResponsesWebsocketPreparesCodexMultiAgentV2Tools(t *testing.T) { + gin.SetMode(gin.TestMode) + executor := &websocketDirectCaptureExecutor{provider: "codex"} + manager := coreauth.NewManager(nil, nil, nil) + manager.RegisterExecutor(executor) + auth := &coreauth.Auth{ID: "responses-multi-agent-ws-auth", Provider: "codex", Status: coreauth.StatusActive, ProxyURL: "direct"} + if _, errRegister := manager.Register(context.Background(), auth); errRegister != nil { + t.Fatalf("Register auth: %v", errRegister) + } + modelID := "responses-multi-agent-ws-model" + registry.GetGlobalRegistry().RegisterClient(auth.ID, auth.Provider, []*registry.ModelInfo{{ID: modelID}}) + t.Cleanup(func() { + registry.GetGlobalRegistry().UnregisterClient(auth.ID) + }) + + base := handlers.NewBaseAPIHandlers(&sdkconfig.SDKConfig{CodexOptimizeMultiAgentV2: true}, manager) + handler := NewOpenAIResponsesAPIHandler(base) + router := gin.New() + router.GET("/v1/responses", handler.ResponsesWebsocket) + server := httptest.NewServer(router) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + "/v1/responses" + conn, _, errDial := websocket.DefaultDialer.Dial(wsURL, http.Header{"User-Agent": []string{"codex_cli_rs/0.144.1"}}) + if errDial != nil { + t.Fatalf("dial websocket: %v", errDial) + } + defer func() { _ = conn.Close() }() + + request := fmt.Sprintf(`{"type":"response.create","model":%q,"input":[],"tools":[{"type":"namespace","name":"collaboration","tools":[{"type":"function","name":"spawn_agent","description":"Spawns an agent.","parameters":{"properties":{"message":{"encrypted":true}}}}]}]}`, modelID) + if errWrite := conn.WriteMessage(websocket.TextMessage, []byte(request)); errWrite != nil { + t.Fatalf("write websocket request: %v", errWrite) + } + if _, _, errRead := conn.ReadMessage(); errRead != nil { + t.Fatalf("read websocket response: %v", errRead) + } + + payloads := executor.Payloads() + if len(payloads) != 1 { + t.Fatalf("captured payload count = %d, want 1", len(payloads)) + } + captured := payloads[0] + if encrypted := gjson.GetBytes(captured, "tools.0.tools.0.parameters.properties.message.encrypted"); encrypted.Exists() { + t.Fatalf("message.encrypted was not removed: %s", captured) + } + if namespace := gjson.GetBytes(captured, "tools.0.name").String(); namespace != "collaboration" { + t.Fatalf("namespace = %q, want collaboration", namespace) + } +} + +func TestPrepareCodexMultiAgentV2ToolsAtResponsesBoundarySkipsOtherClients(t *testing.T) { + t.Parallel() + + gin.SetMode(gin.TestMode) + base := handlers.NewBaseAPIHandlers(&sdkconfig.SDKConfig{CodexOptimizeMultiAgentV2: true}, nil) + handler := NewOpenAIResponsesAPIHandler(base) + request := httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + request.Header.Set("User-Agent", "curl/8.7.1") + ginContext, _ := gin.CreateTestContext(httptest.NewRecorder()) + ginContext.Request = request + + payload := []byte(`{"tools":[{"type":"function","name":"send_message","parameters":{"properties":{"message":{"encrypted":true}}}}]}`) + got := handler.prepareCodexMultiAgentV2Tools(ginContext, payload) + + if string(got) != string(payload) { + t.Fatalf("other client payload changed: %s", got) + } + if _, exists := ginContext.Get(multiagentv2.CodexMultiAgentV2ToolsPreparedContextKey); exists { + t.Fatal("other client unexpectedly received prepared marker") + } +} diff --git a/sdk/api/handlers/openai/openai_responses_websocket.go b/sdk/api/handlers/openai/openai_responses_websocket.go index afcd7b8e..fce86e82 100644 --- a/sdk/api/handlers/openai/openai_responses_websocket.go +++ b/sdk/api/handlers/openai/openai_responses_websocket.go @@ -455,6 +455,9 @@ func (h *OpenAIResponsesAPIHandler) ResponsesWebsocket(c *gin.Context) { } continue } + + requestJSON = h.prepareCodexMultiAgentV2Tools(c, requestJSON) + if !useUpstreamWebsocketPassthrough && shouldHandleResponsesWebsocketPrewarmLocally(payload, lastRequest, false) { if updated, errDelete := sjson.DeleteBytes(requestJSON, "generate"); errDelete == nil { requestJSON = updated