diff --git a/libs/partners/fireworks/langchain_fireworks/chat_models.py b/libs/partners/fireworks/langchain_fireworks/chat_models.py index 78306dbb7a..3b10e157b9 100644 --- a/libs/partners/fireworks/langchain_fireworks/chat_models.py +++ b/libs/partners/fireworks/langchain_fireworks/chat_models.py @@ -343,6 +343,21 @@ def _format_message_content(content: Any) -> Any: return formatted +def _format_tool_call_arguments(arguments: str | dict | None) -> str | dict: + """Preserve invalid historical arguments inside a JSON object for replay.""" + if isinstance(arguments, dict): + return arguments + try: + parsed = json.loads(arguments) if arguments is not None else None + if isinstance(parsed, dict) and arguments is not None: + json.dumps(parsed, allow_nan=False) + return arguments + except ValueError: + logger.debug("Invalid JSON in historical Fireworks tool call arguments") + logger.warning("Wrapping invalid historical Fireworks tool call arguments") + return json.dumps({"__invalid_tool_call_arguments": arguments}, ensure_ascii=False) + + def _convert_message_to_dict(message: BaseMessage) -> dict: """Convert a LangChain message to a dictionary. @@ -392,6 +407,19 @@ def _convert_message_to_dict(message: BaseMessage) -> dict: ] elif "tool_calls" in message.additional_kwargs: message_dict["tool_calls"] = message.additional_kwargs["tool_calls"] + if "tool_calls" in message_dict: + message_dict["tool_calls"] = [ + { + **tool_call, + "function": { + **tool_call["function"], + "arguments": _format_tool_call_arguments( + tool_call["function"]["arguments"] + ), + }, + } + for tool_call in message_dict["tool_calls"] + ] # If tool calls only, content is None not empty string if "tool_calls" in message_dict and message_dict["content"] == "": message_dict["content"] = None diff --git a/libs/partners/fireworks/tests/unit_tests/test_chat_models.py b/libs/partners/fireworks/tests/unit_tests/test_chat_models.py index b665ef0e40..b387c626a1 100644 --- a/libs/partners/fireworks/tests/unit_tests/test_chat_models.py +++ b/libs/partners/fireworks/tests/unit_tests/test_chat_models.py @@ -2,6 +2,7 @@ from __future__ import annotations +import json import logging import os from typing import Any @@ -205,12 +206,71 @@ def test_convert_v1_message_filters_invalid_tool_call_content() -> None: { "type": "function", "id": "call_invalid", - "function": {"name": "get_weather", "arguments": '{"city":'}, + "function": { + "name": "get_weather", + "arguments": json.dumps( + {"__invalid_tool_call_arguments": '{"city":'} + ), + }, } ], } +@pytest.mark.parametrize("raw_history", [False, True]) +@pytest.mark.parametrize( + ("arguments", "wrapped"), + [ + ('{"city":', True), + ("[]", True), + ("null", True), + (None, True), + ('{"x": NaN}', True), + ('{"city": "Paris"}', False), + ], +) +def test_replay_invalid_tool_call_arguments( + arguments: str | None, *, wrapped: bool, raw_history: bool +) -> None: + message = AIMessage( + content="", + tool_calls=[{"name": "get_weather", "args": {"city": "Paris"}, "id": "valid"}], + invalid_tool_calls=[ + {"name": "get_weather", "args": arguments, "id": "invalid", "error": "bad"} + ], + ) + if raw_history: + message.additional_kwargs["tool_calls"] = [ + { + "type": "function", + "id": "invalid", + "function": {"name": "get_weather", "arguments": arguments}, + } + ] + message.tool_calls = [] + message.invalid_tool_calls = [] + original = message.model_dump() + + result = _convert_message_to_dict(message) + + assert message.model_dump() == original + invalid = result["tool_calls"][-1] + assert invalid["id"] == "invalid" + assert invalid["function"]["name"] == "get_weather" + if wrapped: + assert json.loads(invalid["function"]["arguments"]) == { + "__invalid_tool_call_arguments": arguments + } + else: + assert invalid["function"]["arguments"] == arguments + if not raw_history: + assert json.loads(result["tool_calls"][0]["function"]["arguments"]) == { + "city": "Paris" + } + tool_result = ToolMessage(content="Invalid JSON", tool_call_id="invalid") + assert _convert_message_to_dict(tool_result)["tool_call_id"] == invalid["id"] + + def test_sanitize_chat_completions_content_passthrough_non_text_block() -> None: blocks = [{"type": "image_url", "image_url": {"url": "https://x/y.png"}}] assert _sanitize_chat_completions_content(blocks) == blocks