Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
103 changes: 103 additions & 0 deletions apps/application/flow/anthropic_tool_content.py
Original file line number Diff line number Diff line change
@@ -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
69 changes: 59 additions & 10 deletions apps/application/flow/tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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()
Expand All @@ -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
Expand All @@ -561,15 +570,21 @@ 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 片段放这里)
# ----------------------------------------------------------------
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'] 方式(旧格式/非流式情况)
Expand All @@ -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
Expand Down Expand Up @@ -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):
Expand All @@ -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
Expand Down
72 changes: 72 additions & 0 deletions tests/test_anthropic_tool_content.py
Original file line number Diff line number Diff line change
@@ -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()