codekingpro/portable-devtools
114k
1"""This module provides a ValidationNode class that can be used to validate tool calls2in a langchain graph. It applies a pydantic schema to tool_calls in the models' outputs,3and returns a ToolMessage with the validated content. If the schema is not valid, it4returns a ToolMessage with the error message. The ValidationNode can be used in a5StateGraph with a "messages" key. If multiple tool calls are requested, they will be run in parallel.6"""7 8from collections.abc import Callable, Sequence9from typing import (10 Any,11 cast,12)13 14from langchain_core.messages import (15 AIMessage,16 AnyMessage,17 ToolCall,18 ToolMessage,19)20from langchain_core.runnables import (21 RunnableConfig,22)23from langchain_core.runnables.config import get_executor_for_config24from langchain_core.tools import BaseTool, create_schema_from_function25from langchain_core.utils.pydantic import is_basemodel_subclass26from langgraph._internal._runnable import RunnableCallable27from langgraph.warnings import LangGraphDeprecatedSinceV1028from pydantic import BaseModel, ValidationError29from pydantic.v1 import BaseModel as BaseModelV130from pydantic.v1 import ValidationError as ValidationErrorV131from typing_extensions import deprecated32 33 34def _default_format_error(35 error: BaseException,36 call: ToolCall,37 schema: type[BaseModel] | type[BaseModelV1],38) -> str:39 """Default error formatting function."""40 return f"{repr(error)}\n\nRespond after fixing all validation errors."41 42 43@deprecated(44 "ValidationNode is deprecated. Please use `create_agent` from `langchain.agents` with custom tool error handling.",45 category=LangGraphDeprecatedSinceV10,46)47class ValidationNode(RunnableCallable):48 """A node that validates all tools requests from the last `AIMessage`.49 50 It can be used either in `StateGraph` with a `'messages'` key.51 52 !!! note53 54 This node does not actually **run** the tools, it only validates the tool calls,55 which is useful for extraction and other use cases where you need to generate56 structured output that conforms to a complex schema without losing the original57 messages and tool IDs (for use in multi-turn conversations).58 59 Returns:60 (Union[Dict[str, List[ToolMessage]], Sequence[ToolMessage]]): A list of61 `ToolMessage` objects with the validated content or error messages.62 63 Example:64 ```python title="Example usage for re-prompting the model to generate a valid response:"65 from typing import Literal, Annotated66 from typing_extensions import TypedDict67 68 from langchain_anthropic import ChatAnthropic69 from pydantic import BaseModel, field_validator70 71 from langgraph.graph import END, START, StateGraph72 from langgraph.prebuilt import ValidationNode73 from langgraph.graph.message import add_messages74 75 class SelectNumber(BaseModel):76 a: int77 78 @field_validator("a")79 def a_must_be_meaningful(cls, v):80 if v != 37:81 raise ValueError("Only 37 is allowed")82 return v83 84 builder = StateGraph(Annotated[list, add_messages])85 llm = ChatAnthropic(model="claude-3-5-haiku-latest").bind_tools([SelectNumber])86 builder.add_node("model", llm)87 builder.add_node("validation", ValidationNode([SelectNumber]))88 builder.add_edge(START, "model")89 90 def should_validate(state: list) -> Literal["validation", "__end__"]:91 if state[-1].tool_calls:92 return "validation"93 return END94 95 builder.add_conditional_edges("model", should_validate)96 97 def should_reprompt(state: list) -> Literal["model", "__end__"]:98 for msg in state[::-1]:99 # None of the tool calls were errors100 if msg.type == "ai":101 return END102 if msg.additional_kwargs.get("is_error"):103 return "model"104 return END105 106 builder.add_conditional_edges("validation", should_reprompt)107 108 graph = builder.compile()109 res = graph.invoke(("user", "Select a number, any number"))110 # Show the retry logic111 for msg in res:112 msg.pretty_print()113 ```114 """115 116 def __init__(117 self,118 schemas: Sequence[BaseTool | type[BaseModel] | Callable],119 *,120 format_error: Callable[[BaseException, ToolCall, type[BaseModel]], str]121 | None = None,122 name: str = "validation",123 tags: list[str] | None = None,124 ) -> None:125 """Initialize the ValidationNode.126 127 Args:128 schemas: A list of schemas to validate the tool calls with. These can be129 any of the following:130 - A pydantic BaseModel class131 - A BaseTool instance (the args_schema will be used)132 - A function (a schema will be created from the function signature)133 format_error: A function that takes an exception, a ToolCall, and a schema134 and returns a formatted error string. By default, it returns the135 exception repr and a message to respond after fixing validation errors.136 name: The name of the node.137 tags: A list of tags to add to the node.138 """139 super().__init__(self._func, None, name=name, tags=tags, trace=False)140 self._format_error = format_error or _default_format_error141 self.schemas_by_name: dict[str, type[BaseModel]] = {}142 for schema in schemas:143 if isinstance(schema, BaseTool):144 if schema.args_schema is None:145 raise ValueError(146 f"Tool {schema.name} does not have an args_schema defined."147 )148 elif not isinstance(149 schema.args_schema, type150 ) or not is_basemodel_subclass(schema.args_schema):151 raise ValueError(152 "Validation node only works with tools that have a pydantic BaseModel args_schema. "153 f"Got {schema.name} with args_schema: {schema.args_schema}."154 )155 self.schemas_by_name[schema.name] = schema.args_schema156 elif isinstance(schema, type) and issubclass(157 schema, (BaseModel, BaseModelV1)158 ):159 self.schemas_by_name[schema.__name__] = cast(type[BaseModel], schema)160 elif callable(schema):161 base_model = create_schema_from_function("Validation", schema)162 self.schemas_by_name[schema.__name__] = base_model163 else:164 raise ValueError(165 f"Unsupported input to ValidationNode. Expected BaseModel, tool or function. Got: {type(schema)}."166 )167 168 def _get_message(169 self, input: list[AnyMessage] | dict[str, Any]170 ) -> tuple[str, AIMessage]:171 """Extract the last AIMessage from the input."""172 if isinstance(input, list):173 output_type = "list"174 messages: list = input175 elif messages := input.get("messages", []):176 output_type = "dict"177 else:178 raise ValueError("No message found in input")179 message: AnyMessage = messages[-1]180 if not isinstance(message, AIMessage):181 raise ValueError("Last message is not an AIMessage")182 return output_type, message183 184 def _func(185 self, input: list[AnyMessage] | dict[str, Any], config: RunnableConfig186 ) -> Any:187 """Validate and run tool calls synchronously."""188 output_type, message = self._get_message(input)189 190 def run_one(call: ToolCall) -> ToolMessage:191 schema = self.schemas_by_name[call["name"]]192 try:193 if issubclass(schema, BaseModel):194 output = schema.model_validate(call["args"])195 content = output.model_dump_json()196 elif issubclass(schema, BaseModelV1):197 output = schema.validate(call["args"])198 content = output.json()199 else:200 raise ValueError(201 f"Unsupported schema type: {type(schema)}. Expected BaseModel or BaseModelV1."202 )203 return ToolMessage(204 content=content,205 name=call["name"],206 tool_call_id=cast(str, call["id"]),207 )208 except (ValidationError, ValidationErrorV1) as e:209 return ToolMessage(210 content=self._format_error(e, call, schema),211 name=call["name"],212 tool_call_id=cast(str, call["id"]),213 additional_kwargs={"is_error": True},214 )215 216 with get_executor_for_config(config) as executor:217 outputs = [*executor.map(run_one, message.tool_calls)]218 if output_type == "list":219 return outputs220 else:221 return {"messages": outputs}222 