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
71 changes: 52 additions & 19 deletions src/langchain_claude_code/claude_chat_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,54 @@
from langchain_claude_code.claude_code_tools import ClaudeTool, normalize_tools


def _generation_info_from_result(msg: "ResultMessage") -> dict[str, Any]:
"""Build a ``generation_info`` dict from an SDK ``ResultMessage``.

Mirrors the SDK's per-request metadata onto LangChain's
``generation_info`` so downstream consumers can read it off the
final AIMessage's ``response_metadata``. Used by both ``_generate``
(non-streaming) and ``_astream`` (streaming) so the two paths
surface the same field set — previously they diverged: the
non-streaming path preserved ``num_turns`` / ``is_error`` but
missed ``finish_reason``, while the streaming path emitted
``finish_reason`` but dropped ``num_turns`` / ``is_error``
entirely. Neither path preserved ``stop_reason``.

Field shape (all keys present unless explicitly noted):

- ``total_cost_usd``, ``duration_ms``, ``duration_api_ms``,
``session_id`` — SDK fields
- ``num_turns``, ``is_error`` — SDK fields
- ``finish_reason`` — LangChain
convention; ``"error"`` when ``msg.is_error`` else ``"stop"``
- ``stop_reason`` — granular SDK
reason (``"end_turn"``, ``"max_turns"``, ``"max_tokens"``,
etc.); only included when present (newer SDK only) AND
non-``None``. Access is ``getattr``-guarded so the helper
works on the SDK ``>= 0.1.10`` floor this package declares.
- ``usage`` — only when
``msg.usage`` is non-empty
"""
info: dict[str, Any] = {
"total_cost_usd": msg.total_cost_usd,
"duration_ms": msg.duration_ms,
"duration_api_ms": msg.duration_api_ms,
"session_id": msg.session_id,
"num_turns": msg.num_turns,
"is_error": msg.is_error,
"finish_reason": "error" if msg.is_error else "stop",
}
# ``stop_reason`` was added to ResultMessage in a later SDK
# release; use getattr so the helper stays compatible with the
# >= 0.1.10 floor pinned in pyproject.toml.
stop_reason = getattr(msg, "stop_reason", None)
if stop_reason is not None:
info["stop_reason"] = stop_reason
if msg.usage:
info["usage"] = msg.usage
return info


class ClaudeCodeChatModel(BaseChatModel):
"""LangChain chat model wrapping Claude Code Agent SDK.

Expand Down Expand Up @@ -341,16 +389,7 @@ async def _aquery(

elif isinstance(msg, ResultMessage):
self._last_result = msg
generation_info = {
"total_cost_usd": msg.total_cost_usd,
"duration_ms": msg.duration_ms,
"duration_api_ms": msg.duration_api_ms,
"num_turns": msg.num_turns,
"session_id": msg.session_id,
"is_error": msg.is_error,
}
if msg.usage:
generation_info["usage"] = msg.usage
generation_info = _generation_info_from_result(msg)

captured = self._tool_results_var.get()
if captured:
Expand Down Expand Up @@ -503,15 +542,9 @@ async def _astream(
elif isinstance(msg, ResultMessage):
self._last_result = msg

generation_info: dict[str, Any] = {
"total_cost_usd": msg.total_cost_usd,
"duration_ms": msg.duration_ms,
"duration_api_ms": msg.duration_api_ms,
"session_id": msg.session_id,
"finish_reason": "stop" if not msg.is_error else "error",
}
if msg.usage:
generation_info["usage"] = msg.usage
generation_info: dict[str, Any] = (
_generation_info_from_result(msg)
)
if tool_calls_buffer:
generation_info["internal_tool_calls"] = tool_calls_buffer
if tool_results_buffer:
Expand Down
173 changes: 173 additions & 0 deletions tests/test_claude_chat_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -319,6 +319,179 @@ async def test_tool_result_blocks_append_to_content(self):
self.assertIn("tool_results", res.response_metadata)
self.assertEqual(res.response_metadata["tool_results"][0]["tool_use_id"], "t1")

async def test_astream_preserves_full_result_message_fields(self):
"""``_astream``'s result chunk's ``generation_info`` must
carry the full SDK ``ResultMessage`` field set so callers
reading ``finished_message.response_metadata`` see the same
information they'd get from the non-streaming ``_generate``
path. Previously the streaming path dropped ``stop_reason``,
``num_turns``, and ``is_error`` entirely — collapsing
``is_error`` into a binary ``finish_reason`` — which made
the two code paths produce asymmetric AIMessage shapes.

``stop_reason`` is set via ``setattr`` because the SDK
``>= 0.1.10`` floor this package declares predates that
field's addition to ``ResultMessage``; on newer SDKs the
helper picks it up via ``getattr`` without an SDK bump.
"""
rm = ResultMessage(
subtype="result",
duration_ms=10,
duration_api_ms=5,
is_error=False,
num_turns=3,
session_id="sess-stream-full",
total_cost_usd=0.001,
usage={"input_tokens": 10, "output_tokens": 5},
result=None,
structured_output=None,
)
# Simulate newer-SDK stop_reason field via setattr.
try:
rm.stop_reason = "end_turn" # type: ignore[attr-defined]
except (AttributeError, TypeError):
pass
StubClaudeSDKClient.preset_responses = [
AssistantMessage(
content=[TextBlock(text="ok")],
model="test", parent_tool_use_id=None, error=None,
),
rm,
]

model = ClaudeCodeChatModel()
with patch(
"langchain_claude_code.claude_chat_model.ClaudeSDKClient",
StubClaudeSDKClient,
):
chunks = [
chunk async for chunk in model._astream([HumanMessage(content="hi")])
]

final = chunks[-1].generation_info
# Granular SDK fields all present:
self.assertEqual(final["num_turns"], 3)
self.assertEqual(final["is_error"], False)
# LangChain convention finish_reason also present (binary):
self.assertEqual(final["finish_reason"], "stop")
# stop_reason is preserved when present on the SDK message.
# On SDK 0.1.10 (no stop_reason field) the key is omitted —
# that's a forward-compatibility-only assertion; gate it.
if hasattr(StubClaudeSDKClient.preset_responses[1], "stop_reason"):
self.assertEqual(final.get("stop_reason"), "end_turn")
# Existing fields untouched:
self.assertEqual(final["total_cost_usd"], 0.001)
self.assertEqual(final["session_id"], "sess-stream-full")
self.assertEqual(final["usage"], {"input_tokens": 10, "output_tokens": 5})

async def test_astream_finish_reason_error_on_is_error_true(self):
"""When ``ResultMessage.is_error`` is True, ``finish_reason``
renders as ``"error"`` (LangChain convention) and ``is_error``
is preserved as its own field for callers who need the boolean
directly."""
rm = ResultMessage(
subtype="result",
duration_ms=1, duration_api_ms=1,
is_error=True,
num_turns=50,
session_id="sess-err",
total_cost_usd=None,
usage=None,
result=None,
structured_output=None,
)
try:
rm.stop_reason = "max_turns" # type: ignore[attr-defined]
except (AttributeError, TypeError):
pass
StubClaudeSDKClient.preset_responses = [rm]
model = ClaudeCodeChatModel()
with patch(
"langchain_claude_code.claude_chat_model.ClaudeSDKClient",
StubClaudeSDKClient,
):
chunks = [
chunk async for chunk in model._astream([HumanMessage(content="x")])
]
final = chunks[-1].generation_info
self.assertEqual(final["is_error"], True)
self.assertEqual(final["finish_reason"], "error")
self.assertEqual(final["num_turns"], 50)
if hasattr(StubClaudeSDKClient.preset_responses[0], "stop_reason"):
self.assertEqual(final.get("stop_reason"), "max_turns")

async def test_generate_preserves_full_result_message_fields(self):
"""Non-streaming ``_generate`` path mirrors the streaming
path's generation_info shape — the same helper builds both,
so the field set is symmetric. Pre-fix, ``_generate``
preserved ``num_turns`` / ``is_error`` but emitted no
``finish_reason``, while ``_astream`` did the opposite."""
rm = ResultMessage(
subtype="result",
duration_ms=8, duration_api_ms=4,
is_error=False,
num_turns=2,
session_id="sess-gen-full",
total_cost_usd=0.005,
usage={"output_tokens": 7},
result=None,
structured_output=None,
)
try:
rm.stop_reason = "end_turn" # type: ignore[attr-defined]
except (AttributeError, TypeError):
pass
StubClaudeSDKClient.preset_responses = [
AssistantMessage(
content=[TextBlock(text="done")],
model="test", parent_tool_use_id=None, error=None,
),
rm,
]
model = ClaudeCodeChatModel()
with patch(
"langchain_claude_code.claude_chat_model.ClaudeSDKClient",
StubClaudeSDKClient,
):
res = await model.ainvoke([HumanMessage(content="hi")])
md = res.response_metadata
# All SDK fields preserved + finish_reason added per
# LangChain convention.
self.assertEqual(md["num_turns"], 2)
self.assertEqual(md["is_error"], False)
self.assertEqual(md["finish_reason"], "stop")
self.assertEqual(md["session_id"], "sess-gen-full")
if hasattr(rm, "stop_reason"):
self.assertEqual(md.get("stop_reason"), "end_turn")

def test_generation_info_from_result_helper(self):
"""Direct unit test of the helper that both paths share:
produces a consistent generation_info dict from any
ResultMessage. Omits ``stop_reason`` when the SDK doesn't
carry it (older SDK) or when present-but-``None``. Omits
``usage`` only when empty."""
from langchain_claude_code.claude_chat_model import (
_generation_info_from_result,
)

msg = ResultMessage(
subtype="result",
duration_ms=1, duration_api_ms=1,
is_error=False,
num_turns=1,
session_id="s1",
total_cost_usd=0.0,
usage=None, # falsy → key omitted
result=None,
structured_output=None,
)
info = _generation_info_from_result(msg)
self.assertEqual(info["num_turns"], 1)
self.assertEqual(info["is_error"], False)
self.assertEqual(info["finish_reason"], "stop")
self.assertNotIn("stop_reason", info) # absent on SDK 0.1.10
self.assertNotIn("usage", info) # None → omitted


if __name__ == "__main__":
unittest.main()