Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
tool_validator.py222 linesDownload Raw Back to prebuilt
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 
codekingpro/portable-devtools · Team Ai