From afd5a34c28ccf2d6195e4852f0d0fcc2ef4988ed Mon Sep 17 00:00:00 2001 From: lorenzozanee Date: Wed, 16 Sep 2026 02:53:23 +0800 Subject: [PATCH] fix: assemble Anthropic streamed tool_use blocks before replay --- .../flow/anthropic_tool_content.py | 103 ++++++++++++++++++ apps/application/flow/tools.py | 69 ++++++++++-- tests/test_anthropic_tool_content.py | 72 ++++++++++++ 3 files changed, 234 insertions(+), 10 deletions(-) create mode 100644 apps/application/flow/anthropic_tool_content.py create mode 100644 tests/test_anthropic_tool_content.py diff --git a/apps/application/flow/anthropic_tool_content.py b/apps/application/flow/anthropic_tool_content.py new file mode 100644 index 00000000000..b6f04cfd76b --- /dev/null +++ b/apps/application/flow/anthropic_tool_content.py @@ -0,0 +1,103 @@ +# coding=utf-8 +"""Normalize Anthropic streamed tool content before message replay.""" + +import json + +ANTHROPIC_TOOL_STOP_REASONS = {"tool_use"} + + +def is_anthropic_tool_finish(response_metadata, chunk_position=None): + metadata = response_metadata or {} + return ( + metadata.get("finish_reason") == "tool_calls" + or metadata.get("stop_reason") in ANTHROPIC_TOOL_STOP_REASONS + or chunk_position == "last" + ) + + +def collect_input_json_deltas(content): + collected = {} + if not isinstance(content, list): + return collected + for block in content: + if not isinstance(block, dict) or block.get("type") != "input_json_delta": + continue + index = block.get("index") + if index is None: + raise ValueError("Anthropic input_json_delta is missing its content-block index") + fragment = block.get("partial_json", block.get("input", "")) + if not isinstance(fragment, str): + fragment = json.dumps(fragment, ensure_ascii=False) + collected[index] = collected.get(index, "") + fragment + return collected + + +def _parse_tool_input(value): + if isinstance(value, (dict, list)): + return value + if not value: + return None + try: + parsed = json.loads(value) + except (TypeError, ValueError, json.JSONDecodeError) as error: + raise ValueError("Invalid Anthropic tool input JSON") from error + if not isinstance(parsed, (dict, list)): + raise ValueError("Anthropic tool input must be an object or array") + return parsed + + +def _input_from_fragments(block, fragments): + if not fragments: + return None + for fragment in fragments.values(): + if not isinstance(fragment, dict): + continue + if block.get("id") and fragment.get("id") == block["id"]: + return _parse_tool_input(fragment.get("arguments")) + if block.get("index") is not None and fragment.get("index") == block["index"]: + return _parse_tool_input(fragment.get("arguments")) + return None + + +def _input_from_tool_calls(block, tool_calls): + for tool_call in tool_calls or []: + if not isinstance(tool_call, dict): + continue + if block.get("id") and tool_call.get("id") == block["id"]: + return _parse_tool_input(tool_call.get("args") or tool_call.get("arguments")) + if block.get("index") is not None and tool_call.get("index") == block["index"]: + return _parse_tool_input(tool_call.get("args") or tool_call.get("arguments")) + return None + + +def finalize_anthropic_assistant_content(content, tool_calls=None, fragments=None): + if not isinstance(content, list): + return content + + json_by_index = collect_input_json_deltas(content) + finalized = [] + for block in content: + if not isinstance(block, dict): + finalized.append(block) + continue + block_type = block.get("type") + if block_type == "input_json_delta": + continue + if block_type in ("text", "text_delta") and not str(block.get("text") or "").strip(): + continue + if block_type == "tool_use": + normalized = dict(block) + normalized.pop("partial_json", None) + input_value = _input_from_fragments(normalized, fragments) + if input_value is None: + input_value = _input_from_tool_calls(normalized, tool_calls) + if input_value is None: + input_value = _parse_tool_input(json_by_index.get(normalized.get("index"))) + if input_value is not None and not normalized.get("input"): + normalized["input"] = input_value + elif normalized.get("input") == "": + normalized["input"] = {} + finalized.append(normalized) + continue + finalized.append(block) + return finalized diff --git a/apps/application/flow/tools.py b/apps/application/flow/tools.py index 504c5f45494..46ed1f9dd25 100644 --- a/apps/application/flow/tools.py +++ b/apps/application/flow/tools.py @@ -45,6 +45,12 @@ from langchain_core.messages import AIMessageChunk, BaseMessage, BaseMessageChunk, ToolMessage from langchain_core.tools import StructuredTool from langchain_core.utils._merge import merge_lists as _original_merge_lists +from langchain_mcp_adapters.client import MultiServerMCPClient +from .anthropic_tool_content import ( + collect_input_json_deltas, + finalize_anthropic_assistant_content, + is_anthropic_tool_finish, +) from langgraph.checkpoint.memory import MemorySaver from maxkb.const import CONFIG from pydantic import Field, create_model @@ -480,6 +486,7 @@ async def _yield_mcp_response( tool_calls_info = {} # tool_id -> {'name': ..., 'input': ...} # key(index/id) -> {'id': ..., 'name': ..., 'arguments': ...} _tool_fragments = {} + _anthropic_chunks = [] def _merge_arguments(entry, part_args): if not isinstance(part_args, str): @@ -517,10 +524,12 @@ def _get_fragment_key(idx, raw_id): return f"id:{_extract_tool_id(str(raw_id).strip())}" return None - def _upsert_fragment(key, raw_id, func_name, part_args): + def _upsert_fragment(key, raw_id, func_name, part_args, index=None): if key is None: return - entry = _tool_fragments.setdefault(key, {"id": "", "name": "", "arguments": ""}) + entry = _tool_fragments.setdefault(key, {"id": "", "name": "", "arguments": "", "index": index}) + if index is not None: + entry["index"] = index if raw_id and str(raw_id).strip(): new_id = str(raw_id).strip() @@ -546,7 +555,7 @@ def _upsert_fragment(key, raw_id, func_name, part_args): for tc_chunk in chunk[0].tool_call_chunks or []: raw_id = tc_chunk.get("id") key = _get_fragment_key(tc_chunk.get("index"), raw_id) - _upsert_fragment(key, raw_id, tc_chunk.get("name"), tc_chunk.get("args", "")) + _upsert_fragment(key, raw_id, tc_chunk.get("name"), tc_chunk.get("args", ""), tc_chunk.get("index")) # ---------------------------------------------------------------- # 1.1 兼容部分模型将工具调用放在 chunk.tool_calls,且 tool_call_chunks @@ -561,7 +570,7 @@ def _upsert_fragment(key, raw_id, func_name, part_args): if has_tool_call_chunks and (part_args == "" or part_args == {} or part_args == []): part_args = "" key = _get_fragment_key(tool_call.get("index"), raw_id) - _upsert_fragment(key, raw_id, tool_call.get("name"), part_args) + _upsert_fragment(key, raw_id, tool_call.get("name"), part_args, tool_call.get("index")) # ---------------------------------------------------------------- # 1.2 兼容 invalid_tool_calls 分片(部分模型会把中间 JSON 片段放这里) @@ -569,7 +578,13 @@ def _upsert_fragment(key, raw_id, func_name, part_args): for invalid_tool_call in chunk[0].invalid_tool_calls or []: raw_id = invalid_tool_call.get("id") key = _get_fragment_key(invalid_tool_call.get("index"), raw_id) - _upsert_fragment(key, raw_id, invalid_tool_call.get("name"), invalid_tool_call.get("args", "")) + _upsert_fragment( + key, + raw_id, + invalid_tool_call.get("name"), + invalid_tool_call.get("args", ""), + invalid_tool_call.get("index"), + ) # ---------------------------------------------------------------- # 2. 兼容 additional_kwargs['tool_calls'] 方式(旧格式/非流式情况) @@ -585,14 +600,28 @@ def _upsert_fragment(key, raw_id, func_name, part_args): func_name = tool_call.get("name") part_args = tool_call.get("arguments", "") key = _get_fragment_key(tool_call.get("index"), raw_id) - _upsert_fragment(key, raw_id, func_name, part_args) + _upsert_fragment(key, raw_id, func_name, part_args, tool_call.get("index")) + + if isinstance(chunk[0].content, list): + for block in chunk[0].content: + if not isinstance(block, dict) or block.get("type") != "tool_use": + continue + key = _get_fragment_key(block.get("index"), block.get("id")) + _upsert_fragment( + key, + block.get("id"), + block.get("name"), + block.get("input", ""), + block.get("index"), + ) + for index, partial_json in collect_input_json_deltas(chunk[0].content).items(): + key = _get_fragment_key(index, None) + _upsert_fragment(key, None, None, partial_json, index) # ---------------------------------------------------------------- # 3. 检测工具调用结束,更新 tool_calls_info # ---------------------------------------------------------------- - is_finish_chunk = ( - chunk[0].response_metadata.get("finish_reason") == "tool_calls" or chunk[0].chunk_position == "last" - ) + is_finish_chunk = is_anthropic_tool_finish(chunk[0].response_metadata, chunk[0].chunk_position) if is_finish_chunk: # 在 finish chunk 时,将所有未完成的 fragment 标记完成并更新 tool_calls_info @@ -673,6 +702,26 @@ def _upsert_fragment(key, raw_id, func_name, part_args): fixed_tool_calls.append(tc) chunk[0].additional_kwargs["tool_calls"] = fixed_tool_calls + has_anthropic_content = isinstance(chunk[0].content, list) and any( + isinstance(block, dict) and block.get("type") in ("tool_use", "input_json_delta") + for block in chunk[0].content + ) + if _anthropic_chunks or has_anthropic_content: + _anthropic_chunks.append(chunk[0]) + if not is_finish_chunk: + continue + combined_chunk = _anthropic_chunks[0] + for buffered_chunk in _anthropic_chunks[1:]: + combined_chunk += buffered_chunk + _anthropic_chunks.clear() + if isinstance(combined_chunk.content, list): + combined_chunk.content = finalize_anthropic_assistant_content( + combined_chunk.content, + tool_calls=combined_chunk.tool_calls, + fragments=_tool_fragments, + ) + chunk[0] = combined_chunk + yield chunk[0] if mcp_output_enable and isinstance(chunk[0], ToolMessage): @@ -699,7 +748,7 @@ def _upsert_fragment(key, raw_id, func_name, part_args): if tool_lib_id: await save_tool_record(tool_lib_id, tool_info, tool_result, source_id, source_type) tool_result = json.dumps(text_result, ensure_ascii=False) - except Exception as e: + except Exception: tool_result = chunk[0].content content = generate_tool_message_complete( tool_info.get("icon", ""), tool_info["name"], tool_info["input"], tool_result diff --git a/tests/test_anthropic_tool_content.py b/tests/test_anthropic_tool_content.py new file mode 100644 index 00000000000..a2e9a593736 --- /dev/null +++ b/tests/test_anthropic_tool_content.py @@ -0,0 +1,72 @@ +import sys +import unittest +from pathlib import Path + +from langchain_core.messages import AIMessageChunk + +sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "apps")) + +from application.flow.anthropic_tool_content import ( + collect_input_json_deltas, + finalize_anthropic_assistant_content, + is_anthropic_tool_finish, +) + + +class AnthropicToolContentTests(unittest.TestCase): + def test_finish_detects_anthropic_stop_reason(self): + self.assertTrue(is_anthropic_tool_finish({"stop_reason": "tool_use"})) + self.assertTrue(is_anthropic_tool_finish({"finish_reason": "tool_calls"})) + self.assertTrue(is_anthropic_tool_finish({}, chunk_position="last")) + self.assertFalse(is_anthropic_tool_finish({"stop_reason": "end_turn"})) + + def test_strips_deltas_and_fills_tool_input(self): + content = [ + {"type": "text", "text": ""}, + {"type": "tool_use", "id": "toolu_1", "name": "lookup", "input": {}, "index": 1}, + {"type": "input_json_delta", "index": 1, "partial_json": '{"query":'}, + {"type": "input_json_delta", "index": 1, "partial_json": ' "MaxKB"}'}, + ] + + finalized = finalize_anthropic_assistant_content(content) + + self.assertEqual( + finalized, + [{"type": "tool_use", "id": "toolu_1", "name": "lookup", "input": {"query": "MaxKB"}, "index": 1}], + ) + + def test_fills_tool_use_input_from_tool_calls(self): + content = [{"type": "tool_use", "id": "toolu_2", "name": "lookup", "input": ""}] + tool_calls = [{"id": "toolu_2", "name": "lookup", "args": {"query": "value"}}] + + finalized = finalize_anthropic_assistant_content(content, tool_calls=tool_calls) + + self.assertEqual(finalized[0]["input"], {"query": "value"}) + + def test_collect_input_json_deltas_concatenates_fragments(self): + content = [ + {"type": "input_json_delta", "index": 0, "partial_json": '{"a":'}, + {"type": "input_json_delta", "index": 0, "partial_json": " 1}"}, + ] + + self.assertEqual(collect_input_json_deltas(content), {0: '{"a": 1}'}) + + def test_finalizes_langchain_merged_stream(self): + chunks = [ + AIMessageChunk(content=[{"type": "tool_use", "id": "toolu_3", "name": "lookup", "input": {}, "index": 0}]), + AIMessageChunk(content=[{"type": "input_json_delta", "index": 0, "partial_json": '{"q":'}]), + AIMessageChunk(content=[{"type": "input_json_delta", "index": 0, "partial_json": '"v"}'}]), + ] + merged = chunks[0] + chunks[1] + chunks[2] + + finalized = finalize_anthropic_assistant_content( + merged.content, + fragments={"idx:0": {"id": "toolu_3", "index": 0, "arguments": '{"q":"v"}'}}, + ) + + self.assertEqual(finalized[0]["input"], {"q": "v"}) + self.assertNotIn("partial_json", finalized[0]) + + +if __name__ == "__main__": + unittest.main()