Files
langchain/docs/modules/agents/tools/tool_input_validation.ipynb
2023-04-17 21:35:29 -07:00

15 KiB

Tool Input Schema

By default, tools infer the argument schema by inspecting the function signature. For more strict requirements, custom input schema can be specified, along with custom validation logic.

In [1]:
from langchain.agents import initialize_agent, Tool
from langchain.agents import AgentType
from langchain.llms import OpenAI
from langchain.tools.python.tool import PythonREPLTool
from pydantic import BaseModel
from pydantic import Field, root_validator
from typing import Dict, Any
In [2]:
llm = OpenAI(temperature=0)
In [3]:
class ToolInputModel(BaseModel):
    query: str = Field(...)
    
    @root_validator
    def validate_query(cls, values: Dict[str, Any]) -> Dict:
        # Note: this is NOT a safe REPL! This is used for instructive purposes only
        if "import os" in values["query"]:
            raise ValueError("'import os' not permitted in this python REPL.")
        return values
    
tool = PythonREPLTool(args_schema=ToolInputModel)
In [4]:
agent = initialize_agent([tool], llm, agent=AgentType.ZERO_SHOT_REACT_DESCRIPTION, verbose=True)
In [5]:
# This will succeed, since there aren't any arguments that will be triggered during validation
agent.run("Run a python function that adds 2 and 2")
Out [5]:

> Entering new AgentExecutor chain...
 I need to define a function that adds two numbers
Action: Python REPL
Action Input: def add_two_numbers(a, b):
    return a + b I need to call the function
Action: Python REPL
Action Input: print(add_two_numbers(2, 2))4
 I now know the final answer
Final Answer: 4

> Finished chain.
'4'
In [6]:
# This will fail, because the attempt to import os will trigger a validation error
agent.run("Run a python function that imports os and lists the dir")

> Entering new AgentExecutor chain...
 I need to import os and then list the dir
Action: Python REPL
Action Input: import os; print(os.listdir())
---------------------------------------------------------------------------
ValidationError                           Traceback (most recent call last)
Cell In[6], line 1
----> 1 agent.run("Run a python function that imports os and lists the dir")

File ~/code/lc/lckg/langchain/chains/base.py:213, in Chain.run(self, *args, **kwargs)
    211     if len(args) != 1:
    212         raise ValueError("`run` supports only one positional argument.")
--> 213     return self(args[0])[self.output_keys[0]]
    215 if kwargs and not args:
    216     return self(kwargs)[self.output_keys[0]]

File ~/code/lc/lckg/langchain/chains/base.py:116, in Chain.__call__(self, inputs, return_only_outputs)
    114 except (KeyboardInterrupt, Exception) as e:
    115     self.callback_manager.on_chain_error(e, verbose=self.verbose)
--> 116     raise e
    117 self.callback_manager.on_chain_end(outputs, verbose=self.verbose)
    118 return self.prep_outputs(inputs, outputs, return_only_outputs)

File ~/code/lc/lckg/langchain/chains/base.py:113, in Chain.__call__(self, inputs, return_only_outputs)
    107 self.callback_manager.on_chain_start(
    108     {"name": self.__class__.__name__},
    109     inputs,
    110     verbose=self.verbose,
    111 )
    112 try:
--> 113     outputs = self._call(inputs)
    114 except (KeyboardInterrupt, Exception) as e:
    115     self.callback_manager.on_chain_error(e, verbose=self.verbose)

File ~/code/lc/lckg/langchain/agents/agent.py:792, in AgentExecutor._call(self, inputs)
    790 # We now enter the agent loop (until it returns something).
    791 while self._should_continue(iterations, time_elapsed):
--> 792     next_step_output = self._take_next_step(
    793         name_to_tool_map, color_mapping, inputs, intermediate_steps
    794     )
    795     if isinstance(next_step_output, AgentFinish):
    796         return self._return(next_step_output, intermediate_steps)

File ~/code/lc/lckg/langchain/agents/agent.py:695, in AgentExecutor._take_next_step(self, name_to_tool_map, color_mapping, inputs, intermediate_steps)
    693         tool_run_kwargs["llm_prefix"] = ""
    694     # We then call the tool on the tool input to get an observation
--> 695     observation = tool.run(
    696         agent_action.tool_input,
    697         verbose=self.verbose,
    698         color=color,
    699         **tool_run_kwargs,
    700     )
    701 else:
    702     tool_run_kwargs = self.agent.tool_run_logging_kwargs()

File ~/code/lc/lckg/langchain/tools/base.py:146, in BaseTool.run(self, tool_input, verbose, start_color, color, **kwargs)
    137 def run(
    138     self,
    139     tool_input: Union[str, Dict],
   (...)
    143     **kwargs: Any,
    144 ) -> str:
    145     """Run the tool."""
--> 146     run_input = self._parse_input(tool_input)
    147     if not self.verbose and verbose is not None:
    148         verbose_ = verbose

File ~/code/lc/lckg/langchain/tools/base.py:112, in BaseTool._parse_input(self, tool_input)
    110     tool_input = {field_name: tool_input}
    111 if pydantic_input_type is not None:
--> 112     return pydantic_input_type.parse_obj(tool_input)
    113 else:
    114     raise ValueError(
    115         f"args_schema required for tool {self.name} in order to"
    116         f" accept input of type {type(tool_input)}"
    117     )

File ~/code/lc/lckg/.venv/lib/python3.11/site-packages/pydantic/main.py:526, in pydantic.main.BaseModel.parse_obj()

File ~/code/lc/lckg/.venv/lib/python3.11/site-packages/pydantic/main.py:341, in pydantic.main.BaseModel.__init__()

ValidationError: 1 validation error for ToolInputModel
__root__
  'import os' not  (type=value_error)
In [ ]: