diff --git a/libs/core/langchain_core/messages/utils.py b/libs/core/langchain_core/messages/utils.py index 2b1a1ba7b6..86619c699a 100644 --- a/libs/core/langchain_core/messages/utils.py +++ b/libs/core/langchain_core/messages/utils.py @@ -47,7 +47,7 @@ from langchain_core.messages.human import HumanMessage, HumanMessageChunk from langchain_core.messages.modifier import RemoveMessage from langchain_core.messages.system import SystemMessage, SystemMessageChunk from langchain_core.messages.tool import ToolCall, ToolMessage, ToolMessageChunk -from langchain_core.utils.function_calling import convert_to_openai_tool +from langchain_core.utils.pydantic import model_json_schema as get_model_json_schema if TYPE_CHECKING: from langchain_core.language_models import BaseLanguageModel @@ -2274,9 +2274,10 @@ def count_tokens_approximately( using the **most recent** AI message that has `usage_metadata['total_tokens']`. The scaling factor is: `AI_total_tokens / approx_tokens_up_to_that_AI_message` - tools: List of tools to include in the token count. Each tool can be either - a `BaseTool` instance or a dict representing a tool schema. `BaseTool` - instances are converted to OpenAI tool format before counting. + tools: List of tools to include in the token count. Each tool can be a + `BaseTool` instance, a dict representing a tool schema, or a plain + callable. `BaseTool` instances and callables are converted to + OpenAI tool format before counting. Returns: Approximate number of tokens in the messages (and tools, if provided). @@ -2304,7 +2305,24 @@ def count_tokens_approximately( if tools: tools_chars = 0 for tool in tools: - tool_dict = tool if isinstance(tool, dict) else convert_to_openai_tool(tool) + if isinstance(tool, dict): + tool_dict = tool + else: + # tool_call_schema is memoized per instance + schema = tool.tool_call_schema + if isinstance(schema, dict): + parameters = dict(schema) + else: + parameters = dict(get_model_json_schema(schema)) + # Drop the schema's own `title`/`description` to avoid double-counting: + # they're duplicated at the top level. + parameters.pop("title", None) + parameters.pop("description", None) + tool_dict = { + "name": tool.name, + "description": tool.description, + "parameters": parameters, + } tools_chars += len(json.dumps(tool_dict)) token_count += math.ceil(tools_chars / chars_per_token) diff --git a/libs/core/tests/unit_tests/messages/test_utils.py b/libs/core/tests/unit_tests/messages/test_utils.py index 6a93029095..552f727a5a 100644 --- a/libs/core/tests/unit_tests/messages/test_utils.py +++ b/libs/core/tests/unit_tests/messages/test_utils.py @@ -29,7 +29,7 @@ from langchain_core.messages.utils import ( merge_message_runs, trim_messages, ) -from langchain_core.tools import BaseTool, tool +from langchain_core.tools import BaseTool, StructuredTool, tool @pytest.mark.parametrize("msg_cls", [HumanMessage, AIMessage, SystemMessage]) @@ -2961,6 +2961,33 @@ def test_count_tokens_approximately_with_tools() -> None: assert count_empty_tools == base_count +def test_count_tokens_approximately_basetool_dict_args_schema() -> None: + """`BaseTool.tool_call_schema` can itself already be a plain dict. + + This happens when `args_schema` is given as a raw JSON schema instead of + a Pydantic model class -- there's no model class to call + `model_json_schema()` on in that case. + """ + schema_dict = { + "title": "GetWeatherInput", + "type": "object", + "properties": {"location": {"type": "string"}}, + "required": ["location"], + } + weather_tool = StructuredTool.from_function( + func=lambda location: f"Weather in {location}", + name="get_weather", + description="Get the weather for a location.", + args_schema=schema_dict, + ) + assert isinstance(weather_tool.tool_call_schema, dict) + + messages = [HumanMessage(content="Hello")] + base_count = count_tokens_approximately(messages) + count_with_tool = count_tokens_approximately(messages, tools=[weather_tool]) + assert count_with_tool > base_count + + # --------------------------------------------------------------------------- # `Serializable` constructor-envelope wire-shape acceptance in `_convert_to_message` # ---------------------------------------------------------------------------