refactor: adjust ToolCallingAgent to use new IO and ChatHistory, added helper functions

This commit is contained in:
Cezar Vasconcelos 2024-06-21 01:54:04 +00:00
commit 33129c7466

View file

@ -2,63 +2,99 @@ from typing import List, Optional
from langchain.agents.tool_calling_agent.base import create_tool_calling_agent from langchain.agents.tool_calling_agent.base import create_tool_calling_agent
from langchain_core.prompts import ChatPromptTemplate from langchain_core.prompts import ChatPromptTemplate
from langchain.agents import AgentExecutor
from langflow.base.agents.agent import LCAgentComponent from langchain_core.messages import BaseMessage, HumanMessage, AIMessage
from langflow.schema.message import Message
from langflow.custom import Component
from langflow.io import HandleInput, TextInput, BoolInput, Output
from langflow.field_typing import LanguageModel, Text, Tool from langflow.field_typing import LanguageModel, Text, Tool
from langflow.schema import Data from langflow.schema import Data
class ToolCallingAgentComponent(LCAgentComponent): class ToolCallingAgentComponent(Component):
display_name: str = "Tool Calling Agent" display_name: str = "Tool Calling Agent"
description: str = "Agent that uses tools. Only models that are compatible with function calling are supported." description: str = "Agent that uses tools. Only models that are compatible with function calling are supported."
icon = "Agent"
def build_config(self): inputs = [
return { HandleInput(
"llm": {"display_name": "LLM"}, name="llm",
"tools": {"display_name": "Tools"}, display_name="LLM",
"user_prompt": { input_types=["LanguageModel"],
"display_name": "Prompt", ),
"multiline": True, HandleInput(
"info": "This prompt must contain 'input' key.", name="tools",
}, display_name="Tools",
"handle_parsing_errors": { input_types=["Tool"],
"display_name": "Handle Parsing Errors", is_list=True,
"info": "If True, the agent will handle parsing errors. If False, the agent will raise an error.", ),
"advanced": True, TextInput(
}, name="user_prompt",
"memory": { display_name="Prompt",
"display_name": "Memory", info="This prompt must contain 'input' key.",
"info": "Memory to use for the agent.", value="{input}",
}, ),
"input_value": { BoolInput(
"display_name": "Inputs", name="handle_parsing_errors",
"info": "Input text to pass to the agent.", display_name="Handle Parsing Errors",
}, info="If True, the agent will handle parsing errors. If False, the agent will raise an error.",
} advanced=True,
value=True,
),
HandleInput(
name="memory",
display_name="Memory",
input_types=["Data"],
info="Memory to use for the agent.",
),
TextInput(
name="input_value",
display_name="Inputs",
info="Input text to pass to the agent.",
),
]
async def build( outputs = [
self, Output(display_name="Text", name="text_output", method="run_agent"),
input_value: str, ]
llm: LanguageModel,
tools: List[Tool], async def run_agent(self) -> Message:
user_prompt: str = "{input}", if "input" not in self.user_prompt:
message_history: Optional[List[Data]] = None,
system_message: str = "You are a helpful assistant",
handle_parsing_errors: bool = True,
) -> Text:
if "input" not in user_prompt:
raise ValueError("Prompt must contain 'input' key.") raise ValueError("Prompt must contain 'input' key.")
messages = [ messages = [
("system", system_message), ("system", "You are a helpful assistant"),
( (
"placeholder", "placeholder",
"{chat_history}", "{chat_history}",
), ),
("human", user_prompt), ("human", self.user_prompt),
("placeholder", "{agent_scratchpad}"), ("placeholder", "{agent_scratchpad}"),
] ]
prompt = ChatPromptTemplate.from_messages(messages) prompt = ChatPromptTemplate.from_messages(messages)
agent = create_tool_calling_agent(llm, tools, prompt) agent = create_tool_calling_agent(self.llm, self.tools, prompt)
result = await self.run_agent(agent, input_value, tools, message_history, handle_parsing_errors)
runnable = AgentExecutor.from_agent_and_tools(
agent=agent,
tools=self.tools,
verbose=True,
handle_parsing_errors=self.handle_parsing_errors,
)
input_dict: dict[str, str | list[BaseMessage]] = {"input": self.input_value}
if hasattr(self, "memory") and self.memory:
input_dict["chat_history"] = self.convert_chat_history(self.memory)
result = await runnable.ainvoke(input_dict)
self.status = result self.status = result
return result
if "output" not in result:
raise ValueError("Output key not found in result. Tried 'output'.")
result_string = result["output"]
return Message(text=result_string)
def convert_chat_history(self, chat_history: List[Data]) -> List[Dict[str, str]]:
messages = []
for item in chat_history:
role = "user" if item.sender == "User" else "assistant"
messages.append({"role": role, "content": item.text})
return messages