refactor: adjust ToolCallingAgent to use new IO and ChatHistory, added helper functions
This commit is contained in:
parent
806cbf3842
commit
33129c7466
1 changed files with 78 additions and 42 deletions
|
|
@ -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
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue