diff --git a/internal/runtime/executor/claude_executor_request.go b/internal/runtime/executor/claude_executor_request.go index 71c2db90..486cf7e3 100644 --- a/internal/runtime/executor/claude_executor_request.go +++ b/internal/runtime/executor/claude_executor_request.go @@ -1392,6 +1392,25 @@ func remapOAuthToolNamesWithBatchedEdits(body []byte, mcpAliases claudeMCPAliasO return true }) } + case "tool_search_tool_result": + toolRefs := part.Get("content.tool_references") + if toolRefs.Exists() && toolRefs.IsArray() { + toolRefs.ForEach(func(_, refPart gjson.Result) bool { + if refPart.Get("type").String() != "tool_reference" { + return true + } + nameResult := refPart.Get("tool_name") + refToolName := nameResult.String() + if newName, renamed := rewriteName(refToolName); renamed { + if !appendStringEdit(nameResult, newName) { + validOffsets = false + return false + } + recordRename(refToolName, newName) + } + return true + }) + } } return validOffsets }) @@ -1620,6 +1639,21 @@ func remapOAuthToolNamesWithOptionsLegacy(body []byte, mcpAliases claudeMCPAlias return true }) } + case "tool_search_tool_result": + toolRefs := part.Get("content.tool_references") + if toolRefs.Exists() && toolRefs.IsArray() { + toolRefs.ForEach(func(refIndex, refPart gjson.Result) bool { + if refPart.Get("type").String() == "tool_reference" { + refToolName := refPart.Get("tool_name").String() + if newName, renamed := rewriteName(refToolName); renamed { + refPath := fmt.Sprintf("messages.%d.content.%d.content.tool_references.%d.tool_name", msgIndex.Int(), contentIndex.Int(), refIndex.Int()) + body, _ = sjson.SetBytes(body, refPath, newName) + recordRename(refToolName, newName) + } + } + return true + }) + } } return true }) @@ -1648,6 +1682,18 @@ type claudeMCPAliasResolver struct { servers map[string]struct{} } +type claudeMCPAliasRestoreError struct { + error +} + +func (e claudeMCPAliasRestoreError) Unwrap() error { + return e.error +} + +func (claudeMCPAliasRestoreError) IsRequestScoped() bool { + return true +} + func newClaudeMCPAliasResolver(reverseMap map[string]string) claudeMCPAliasResolver { resolver := claudeMCPAliasResolver{ exact: reverseMap, @@ -1761,7 +1807,7 @@ func (resolver claudeMCPAliasResolver) resolve(name string) (string, bool, error return matchedOriginal, true, nil } if matchCount > 1 { - return "", false, fmt.Errorf("cannot restore Claude OAuth MCP tool alias %q: matched multiple declared aliases", name) + return "", false, claudeMCPAliasRestoreError{fmt.Errorf("cannot restore Claude OAuth MCP tool alias %q: matched multiple declared aliases", name)} } parts, validAlias := parseClaudeMCPAlias(normalizedName) @@ -1816,10 +1862,10 @@ func (resolver claudeMCPAliasResolver) resolve(name string) (string, bool, error return matchedOriginal, true, nil } if matchCount > 1 { - return "", false, fmt.Errorf("cannot restore Claude OAuth MCP tool alias %q: semantic suffix matches multiple declared tools", name) + return "", false, claudeMCPAliasRestoreError{fmt.Errorf("cannot restore Claude OAuth MCP tool alias %q: semantic suffix matches multiple declared tools", name)} } - return "", false, fmt.Errorf("cannot restore Claude OAuth MCP tool alias %q: no unique request-local match", name) + return "", false, claudeMCPAliasRestoreError{fmt.Errorf("cannot restore Claude OAuth MCP tool alias %q: no unique request-local match", name)} } // reverseRemapOAuthToolNames reverses the tool name mapping for non-stream responses @@ -1880,6 +1926,26 @@ func reverseRemapOAuthToolNames(body []byte, reverseMap map[string]string) ([]by return true }) } + case "tool_search_tool_result": + toolRefs := part.Get("content.tool_references") + if toolRefs.Exists() && toolRefs.IsArray() { + toolRefs.ForEach(func(refIndex, refPart gjson.Result) bool { + if refPart.Get("type").String() != "tool_reference" { + return true + } + toolName := refPart.Get("tool_name").String() + origName, matched, errResolve := resolver.resolve(toolName) + if errResolve != nil { + resolveErr = errResolve + return false + } + if matched { + path := fmt.Sprintf("content.%d.content.tool_references.%d.tool_name", index.Int(), refIndex.Int()) + body, _ = sjson.SetBytes(body, path, origName) + } + return true + }) + } } return resolveErr == nil }) @@ -1928,6 +1994,44 @@ func reverseRemapOAuthToolNamesFromStreamLine(line []byte, reverseMap map[string return line, nil } updated, err = sjson.SetBytes(payload, "content_block.tool_name", origName) + case "tool_search_tool_result": + toolRefs := contentBlock.Get("content.tool_references") + if !toolRefs.Exists() || !toolRefs.IsArray() { + return line, nil + } + updatedPayload := payload + var resolveErr error + hasChange := false + toolRefs.ForEach(func(refIndex, refPart gjson.Result) bool { + if refPart.Get("type").String() != "tool_reference" { + return true + } + toolName := refPart.Get("tool_name").String() + origName, matched, errResolve := resolver.resolve(toolName) + if errResolve != nil { + resolveErr = errResolve + return false + } + if matched { + path := fmt.Sprintf("content_block.content.tool_references.%d.tool_name", refIndex.Int()) + updatedPayload, err = sjson.SetBytes(updatedPayload, path, origName) + if err != nil { + return false + } + hasChange = true + } + return true + }) + if resolveErr != nil { + return line, resolveErr + } + if err != nil { + return line, fmt.Errorf("rewrite Claude OAuth MCP tool alias: %w", err) + } + if !hasChange { + return line, nil + } + updated = updatedPayload default: return line, nil } diff --git a/internal/runtime/executor/claude_executor_request_remap_test.go b/internal/runtime/executor/claude_executor_request_remap_test.go index 8c560263..995ef559 100644 --- a/internal/runtime/executor/claude_executor_request_remap_test.go +++ b/internal/runtime/executor/claude_executor_request_remap_test.go @@ -2,12 +2,14 @@ package executor import ( "bytes" + "errors" "fmt" "maps" "strings" "testing" "github.com/router-for-me/CLIProxyAPI/v7/internal/runtime/executor/helps" + cliproxyexecutor "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/executor" "github.com/tidwall/gjson" ) @@ -343,14 +345,24 @@ func TestReverseRemapOAuthToolNamesRejectsUnsafeMangledAliases(t *testing.T) { for _, test := range tests { t.Run(test.name, func(t *testing.T) { response := []byte(fmt.Sprintf(`{"content":[{"type":"tool_use","id":"toolu_1","name":%q,"input":{}}]}`, test.alias)) - if _, errReverse := reverseRemapOAuthToolNames(response, reverseMap); errReverse == nil || !strings.Contains(errReverse.Error(), test.wantError) { + _, errReverse := reverseRemapOAuthToolNames(response, reverseMap) + if errReverse == nil || !strings.Contains(errReverse.Error(), test.wantError) { t.Fatalf("reverseRemapOAuthToolNames() error = %v, want %q", errReverse, test.wantError) } + var requestErr cliproxyexecutor.RequestScopedError + if !errors.As(errReverse, &requestErr) || !requestErr.IsRequestScoped() { + t.Fatalf("reverseRemapOAuthToolNames() error = %T %v, want request-scoped", errReverse, errReverse) + } line := []byte(fmt.Sprintf(`data: {"type":"content_block_start","index":0,"content_block":{"type":"tool_use","id":"toolu_1","name":%q,"input":{}}}`, test.alias)) - if _, errStream := reverseRemapOAuthToolNamesFromStreamLine(line, reverseMap); errStream == nil || !strings.Contains(errStream.Error(), test.wantError) { + _, errStream := reverseRemapOAuthToolNamesFromStreamLine(line, reverseMap) + if errStream == nil || !strings.Contains(errStream.Error(), test.wantError) { t.Fatalf("reverseRemapOAuthToolNamesFromStreamLine() error = %v, want %q", errStream, test.wantError) } + requestErr = nil + if !errors.As(errStream, &requestErr) || !requestErr.IsRequestScoped() { + t.Fatalf("reverseRemapOAuthToolNamesFromStreamLine() error = %T %v, want request-scoped", errStream, errStream) + } }) } } @@ -542,3 +554,194 @@ func TestRemapKeepsReverseMapEmptyWhenOnlyCallerMCPToolsArePresent(t *testing.T) t.Fatalf("body = %s, want unchanged %s", out, body) } } + +func TestReverseRemapOAuthToolNamesMarksTrailingMarkupFailureRequestScoped(t *testing.T) { + const alias = "mcp__hmzqrngkulqv__xuo7jlxlpzee_clear_thinking" + malformedAlias := alias + "\n