fix: ensure that it's calling as anthropic does

This commit is contained in:
Javier Martinez
2026-07-27 12:57:16 +02:00
parent 872dc93798
commit 77ee6a014c
3 changed files with 95 additions and 37 deletions

View File

@@ -209,6 +209,7 @@ class _ToolDeltaState:
active_tool_block: RawContentBlockStartEvent | None = None
active_tool_raw_id: str | None = None
tool_id_map: dict[str, str] = field(default_factory=dict)
tool_name_map: dict[str, str] = field(default_factory=dict)
finished_tool_raw_ids: set[str] = field(default_factory=set)
last_serialized: dict[str, str] = field(default_factory=dict)
pending_tasks: list[asyncio.Task[_ToolExecutionResult]] = field(
@@ -990,7 +991,10 @@ class AsyncChatEngine:
assistant_message = llm_response.message
self._ensure_update_tool_ids_in_tool_selection(
run, assistant_message, stream_delta_state.tool_state.tool_id_map
run,
assistant_message,
stream_delta_state.tool_state.tool_id_map,
stream_delta_state.tool_state.tool_name_map,
)
self._validate_unique_tool_call_ids(assistant_message)
self._accumulate_usage(run, assistant_message)
@@ -1186,6 +1190,7 @@ class AsyncChatEngine:
tool_state.active_tool_block = use_start
tool_state.active_tool_raw_id = raw_id
tool_state.tool_name_map[raw_id] = tool_call.tool_name or ""
tool_state.last_serialized[raw_id] = ""
if tool_call.tool_kwargs and tool_state.active_tool_block is not None:
@@ -1278,7 +1283,7 @@ class AsyncChatEngine:
final_json = tool_state.last_serialized.get(prev_raw_id, "")
final_obj: Any = json.loads(final_json) if final_json else {}
tool_name = getattr(tool_state.active_tool_block.content_block, "name", None)
tool_name = tool_state.tool_name_map.get(prev_raw_id)
if not isinstance(tool_name, str):
raise TypeError("Active tool block must define a name")
tool_schema = schema_by_name.get(tool_name, {})
@@ -1623,7 +1628,10 @@ class AsyncChatEngine:
@staticmethod
def _ensure_update_tool_ids_in_tool_selection(
run: _Run, message: ChatMessage, tool_id_map: dict[str, str]
run: _Run,
message: ChatMessage,
tool_id_map: dict[str, str],
tool_name_map: dict[str, str] | None = None,
) -> None:
tool_calls = message.additional_kwargs.get("tool_calls", [])
if not isinstance(tool_calls, list):
@@ -1633,6 +1641,8 @@ class AsyncChatEngine:
raw_id = tool_call.tool_id
if raw_id and raw_id in tool_id_map:
tool_call.tool_id = tool_id_map[raw_id]
if tool_name_map and raw_id and raw_id in tool_name_map:
tool_call.tool_name = tool_name_map[raw_id]
@staticmethod
def _validate_unique_tool_call_ids(message: ChatMessage) -> None:

View File

@@ -94,6 +94,27 @@ class ServerToolResultBlock(ToolResultBlock):
description="ID of the server tool use this result answers.",
)
def for_response_mode(
self, response_mode: Literal["anthropic", "zylon"]
) -> Self | None:
if isinstance(self.content, str):
return self
if not isinstance(self.content, Sequence):
block = self.content.for_response_mode(response_mode)
if block is not None:
self.content = block
return self
return None
pruned = [
b
for block in self.content
if (b := block.for_response_mode(response_mode)) is not None
]
if pruned:
self.content = pruned
return self
return None
class WebSearchToolResultBlock(ServerToolResultBlock):
"""Anthropic-shaped result for an internally-executed web_search call.

View File

@@ -35,6 +35,7 @@ from private_gpt.events.models import (
SourceBlock,
TextBlock,
TextDelta,
TextEditorCodeExecutionToolResultBlock,
ToolResultBlock,
ToolUseBlock,
)
@@ -1052,48 +1053,74 @@ async def test_chat_body_validation_tool_result_references_unknown_tool(
@pytest.mark.anyio
async def test_chat_body_validation_mismatched_tool_ids(
async def test_code_execution_expand_and_is_usable(
async_test_client: AsyncClient,
injector: MockInjector,
) -> None:
await mock_llm(
injector,
deltas=[
[
ToolSelection(
tool_id="call_001",
tool_name="create",
tool_kwargs={"path": "potato.md", "file_text": "# Potato"},
),
],
[
ToolSelection(
tool_id="call_002",
tool_name="present_files",
tool_kwargs={"filepaths": ["potato.md"]},
)
],
["Created potato.md."],
],
)
body = {
"messages": [
{"content": "test", "role": "user"},
{"content": "Create a potato.md and present the file", "role": "user"}
],
"tools": [
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "tool_1",
"name": "test_tool",
"input": {},
},
{
"type": "tool_use",
"id": "tool_2",
"name": "test_tool2",
"input": {},
},
],
},
{
"role": "assistant",
"content": [
{
"type": "tool_result",
"tool_use_id": "tool_1",
"content": "result",
},
],
},
]
"name": "code_execution",
"type": "code_execution_v1",
"input_schema": {"type": "object", "properties": {}},
}
],
"system": [{"extensions": ["zylon"]}],
}
response = await async_test_client.post("/v1/messages", json=body)
assert response.status_code == 400
error_detail = response.json()["detail"]
assert any(
"Tool result blocks must match the tool use IDs" in str(err)
for err in error_detail
assert response.status_code == 200
completion: Message = Message.model_validate(response.json())
tool_uses = [
block for block in completion.content if isinstance(block, ToolUseBlock)
]
tool_results = [
block for block in completion.content if isinstance(block, ToolResultBlock)
]
assert len(tool_uses) == 2, f"Expected 2 tool_uses, got {len(tool_uses)}"
assert len(tool_results) == 2, f"Expected 2 tool_result, got {len(tool_results)}"
first_tool_use = tool_uses[0]
first_tool_result = tool_results[0]
assert first_tool_use.name == "text_editor_code_execution", (
f"Expected first tool use to be 'text_editor_code_execution', got {first_tool_use.name}"
)
assert isinstance(first_tool_result, TextEditorCodeExecutionToolResultBlock)
second_tool_use = tool_uses[1]
second_tool_result = tool_results[1]
assert second_tool_use.name == "present_files", (
f"Expected second tool use to be 'present_files', got {second_tool_use.name}"
)
assert isinstance(second_tool_result, ToolResultBlock)
@pytest.mark.anyio