codekingpro/portable-devtools
114k
1"""Messages for tools."""2 3import json4from typing import Any, Literal, cast, overload5from uuid import UUID6 7from pydantic import Field, model_validator8from typing_extensions import NotRequired, TypedDict, override9 10from langchain_core.messages import content as types11from langchain_core.messages.base import BaseMessage, BaseMessageChunk, merge_content12from langchain_core.messages.content import InvalidToolCall13from langchain_core.utils._merge import merge_dicts, merge_obj14 15 16class ToolOutputMixin:17 """Mixin for objects that tools can return directly.18 19 If a custom BaseTool is invoked with a `ToolCall` and the output of custom code is20 not an instance of `ToolOutputMixin`, the output will automatically be coerced to21 a string and wrapped in a `ToolMessage`.22 23 """24 25 26class ToolMessage(BaseMessage, ToolOutputMixin):27 """Message for passing the result of executing a tool back to a model.28 29 `ToolMessage` objects contain the result of a tool invocation. Typically, the result30 is encoded inside the `content` field.31 32 `tool_call_id` is used to associate the tool call request with the tool call33 response. Useful in situations where a chat model is able to request multiple tool34 calls in parallel.35 36 Example:37 A `ToolMessage` representing a result of `42` from a tool call with id38 39 ```python40 from langchain_core.messages import ToolMessage41 42 ToolMessage(content="42", tool_call_id="call_Jja7J89XsjrOLA5r!MEOW!SL")43 ```44 45 Example:46 A `ToolMessage` where only part of the tool output is sent to the model47 and the full output is passed in to artifact.48 49 ```python50 from langchain_core.messages import ToolMessage51 52 tool_output = {53 "stdout": "From the graph we can see that the correlation between "54 "x and y is ...",55 "stderr": None,56 "artifacts": {"type": "image", "base64_data": "/9j/4gIcSU..."},57 }58 59 ToolMessage(60 content=tool_output["stdout"],61 artifact=tool_output,62 tool_call_id="call_Jja7J89XsjrOLA5r!MEOW!SL",63 )64 ```65 """66 67 tool_call_id: str68 """Tool call that this message is responding to."""69 70 type: Literal["tool"] = "tool"71 """The type of the message (used for serialization)."""72 73 artifact: Any = None74 """Artifact of the Tool execution which is not meant to be sent to the model.75 76 Should only be specified if it is different from the message content, e.g. if only77 a subset of the full tool output is being passed as message content but the full78 output is needed in other parts of the code.79 80 """81 82 status: Literal["success", "error"] = "success"83 """Status of the tool invocation."""84 85 additional_kwargs: dict = Field(default_factory=dict, repr=False)86 """Currently inherited from `BaseMessage`, but not used."""87 response_metadata: dict = Field(default_factory=dict, repr=False)88 """Currently inherited from `BaseMessage`, but not used."""89 90 @model_validator(mode="before")91 @classmethod92 def coerce_args(cls, values: dict) -> dict:93 """Coerce the model arguments to the correct types.94 95 Args:96 values: The model arguments.97 98 """99 content = values["content"]100 if isinstance(content, tuple):101 content = list(content)102 103 if not isinstance(content, (str, list)):104 try:105 values["content"] = str(content)106 except ValueError as e:107 msg = (108 "ToolMessage content should be a string or a list of string/dicts. "109 f"Received:\n\n{content=}\n\n which could not be coerced into a "110 "string."111 )112 raise ValueError(msg) from e113 elif isinstance(content, list):114 values["content"] = []115 for i, x in enumerate(content):116 if not isinstance(x, (str, dict)):117 try:118 values["content"].append(str(x))119 except ValueError as e:120 msg = (121 "ToolMessage content should be a string or a list of "122 "string/dicts. Received a list but "123 f"element ToolMessage.content[{i}] is not a dict and could "124 f"not be coerced to a string.:\n\n{x}"125 )126 raise ValueError(msg) from e127 else:128 values["content"].append(x)129 130 tool_call_id = values["tool_call_id"]131 if isinstance(tool_call_id, (UUID, int, float)):132 values["tool_call_id"] = str(tool_call_id)133 return values134 135 @overload136 def __init__(137 self,138 content: str | list[str | dict],139 **kwargs: Any,140 ) -> None: ...141 142 @overload143 def __init__(144 self,145 content: str | list[str | dict] | None = None,146 content_blocks: list[types.ContentBlock] | None = None,147 **kwargs: Any,148 ) -> None: ...149 150 def __init__(151 self,152 content: str | list[str | dict] | None = None,153 content_blocks: list[types.ContentBlock] | None = None,154 **kwargs: Any,155 ) -> None:156 """Initialize a `ToolMessage`.157 158 Specify `content` as positional arg or `content_blocks` for typing.159 160 Args:161 content: The contents of the message.162 content_blocks: Typed standard content.163 **kwargs: Additional fields.164 """165 if content_blocks is not None:166 super().__init__(167 content=cast("str | list[str | dict]", content_blocks),168 **kwargs,169 )170 else:171 super().__init__(content=content, **kwargs)172 173 174class ToolMessageChunk(ToolMessage, BaseMessageChunk):175 """Tool Message chunk."""176 177 # Ignoring mypy re-assignment here since we're overriding the value178 # to make sure that the chunk variant can be discriminated from the179 # non-chunk variant.180 type: Literal["ToolMessageChunk"] = "ToolMessageChunk" # type: ignore[assignment]181 182 @override183 def __add__(self, other: Any) -> BaseMessageChunk: # type: ignore[override]184 if isinstance(other, ToolMessageChunk):185 if self.tool_call_id != other.tool_call_id:186 msg = "Cannot concatenate ToolMessageChunks with different names."187 raise ValueError(msg)188 189 return self.__class__(190 tool_call_id=self.tool_call_id,191 content=merge_content(self.content, other.content),192 artifact=merge_obj(self.artifact, other.artifact),193 additional_kwargs=merge_dicts(194 self.additional_kwargs, other.additional_kwargs195 ),196 response_metadata=merge_dicts(197 self.response_metadata, other.response_metadata198 ),199 id=self.id,200 status=_merge_status(self.status, other.status),201 )202 203 return super().__add__(other)204 205 206class ToolCall(TypedDict):207 """Represents an AI's request to call a tool.208 209 Example:210 ```python211 {"name": "foo", "args": {"a": 1}, "id": "123"}212 ```213 214 This represents a request to call the tool named `'foo'` with arguments215 `{"a": 1}` and an identifier of `'123'`.216 217 !!! note "Factory function"218 219 `tool_call` may also be used as a factory to create a `ToolCall`. Benefits220 include:221 222 * Required arguments strictly validated at creation time223 """224 225 name: str226 """The name of the tool to be called."""227 228 args: dict[str, Any]229 """The arguments to the tool call as a dictionary."""230 231 id: str | None232 """An identifier associated with the tool call.233 234 An identifier is needed to associate a tool call request with a tool235 call result in events when multiple concurrent tool calls are made.236 """237 238 type: NotRequired[Literal["tool_call"]]239 """Used for discrimination."""240 241 242def tool_call(243 *,244 name: str,245 args: dict[str, Any],246 id: str | None,247) -> ToolCall:248 """Create a tool call.249 250 Args:251 name: The name of the tool to be called.252 args: The arguments to the tool call as a dictionary.253 id: An identifier associated with the tool call.254 255 Returns:256 The created tool call.257 """258 return ToolCall(name=name, args=args, id=id, type="tool_call")259 260 261class ToolCallChunk(TypedDict):262 """A chunk of a tool call (yielded when streaming).263 264 When merging `ToolCallChunk` objects (e.g., via `AIMessageChunk.__add__`), all265 string attributes are concatenated. Chunks are only merged if their values of266 `index` are equal and not `None`.267 268 Example:269 ```python270 left_chunks = [ToolCallChunk(name="foo", args='{"a":', index=0)]271 right_chunks = [ToolCallChunk(name=None, args="1}", index=0)]272 273 (274 AIMessageChunk(content="", tool_call_chunks=left_chunks)275 + AIMessageChunk(content="", tool_call_chunks=right_chunks)276 ).tool_call_chunks == [ToolCallChunk(name="foo", args='{"a":1}', index=0)]277 ```278 """279 280 name: str | None281 """The name of the tool to be called."""282 283 args: str | None284 """The arguments to the tool call as a JSON-parseable string."""285 286 id: str | None287 """An identifier associated with the tool call.288 289 An identifier is needed to associate a tool call request with a tool290 call result in events when multiple concurrent tool calls are made.291 """292 293 index: int | None294 """The index of the tool call in a sequence.295 296 Used for merging chunks.297 """298 299 type: NotRequired[Literal["tool_call_chunk"]]300 """Used for discrimination."""301 302 303def tool_call_chunk(304 *,305 name: str | None = None,306 args: str | None = None,307 id: str | None = None,308 index: int | None = None,309) -> ToolCallChunk:310 """Create a tool call chunk.311 312 Args:313 name: The name of the tool to be called.314 args: The arguments to the tool call as a JSON string.315 id: An identifier associated with the tool call.316 index: The index of the tool call in a sequence.317 318 Returns:319 The created tool call chunk.320 """321 return ToolCallChunk(322 name=name, args=args, id=id, index=index, type="tool_call_chunk"323 )324 325 326def invalid_tool_call(327 *,328 name: str | None = None,329 args: str | None = None,330 id: str | None = None,331 error: str | None = None,332) -> InvalidToolCall:333 """Create an invalid tool call.334 335 Args:336 name: The name of the tool to be called.337 args: The arguments to the tool call as a JSON string.338 id: An identifier associated with the tool call.339 error: An error message associated with the tool call.340 341 Returns:342 The created invalid tool call.343 """344 return InvalidToolCall(345 name=name, args=args, id=id, error=error, type="invalid_tool_call"346 )347 348 349def default_tool_parser(350 raw_tool_calls: list[dict],351) -> tuple[list[ToolCall], list[InvalidToolCall]]:352 """Best-effort parsing of tools.353 354 Args:355 raw_tool_calls: List of raw tool call dicts to parse.356 357 Returns:358 A list of tool calls and invalid tool calls.359 """360 tool_calls = []361 invalid_tool_calls = []362 for raw_tool_call in raw_tool_calls:363 if "function" not in raw_tool_call:364 continue365 function_name = raw_tool_call["function"]["name"]366 try:367 function_args = json.loads(raw_tool_call["function"]["arguments"])368 parsed = tool_call(369 name=function_name or "",370 args=function_args or {},371 id=raw_tool_call.get("id"),372 )373 tool_calls.append(parsed)374 except json.JSONDecodeError:375 invalid_tool_calls.append(376 invalid_tool_call(377 name=function_name,378 args=raw_tool_call["function"]["arguments"],379 id=raw_tool_call.get("id"),380 error=None,381 )382 )383 return tool_calls, invalid_tool_calls384 385 386def default_tool_chunk_parser(raw_tool_calls: list[dict]) -> list[ToolCallChunk]:387 """Best-effort parsing of tool chunks.388 389 Args:390 raw_tool_calls: List of raw tool call dicts to parse.391 392 Returns:393 List of parsed ToolCallChunk objects.394 """395 tool_call_chunks = []396 for tool_call in raw_tool_calls:397 if "function" not in tool_call:398 function_args = None399 function_name = None400 else:401 function_args = tool_call["function"]["arguments"]402 function_name = tool_call["function"]["name"]403 parsed = tool_call_chunk(404 name=function_name,405 args=function_args,406 id=tool_call.get("id"),407 index=tool_call.get("index"),408 )409 tool_call_chunks.append(parsed)410 return tool_call_chunks411 412 413def _merge_status(414 left: Literal["success", "error"], right: Literal["success", "error"]415) -> Literal["success", "error"]:416 return "error" if "error" in {left, right} else "success"417 