Refactor AgentInitializerComponent to support optional memory parameter

This commit is contained in:
Gabriel Luiz Freitas Almeida 2023-12-22 10:39:26 -03:00
commit f943ea26d6

View file

@ -1,7 +1,6 @@
from typing import Callable, List, Union from typing import Callable, List, Optional, Union
from langchain.agents import AgentExecutor, AgentType, initialize_agent, types from langchain.agents import AgentExecutor, AgentType, initialize_agent, types
from langflow import CustomComponent from langflow import CustomComponent
from langflow.field_typing import BaseChatMemory, BaseLanguageModel, Tool from langflow.field_typing import BaseChatMemory, BaseLanguageModel, Tool
@ -20,12 +19,19 @@ class AgentInitializerComponent(CustomComponent):
"memory": {"display_name": "Memory"}, "memory": {"display_name": "Memory"},
"tools": {"display_name": "Tools"}, "tools": {"display_name": "Tools"},
"llm": {"display_name": "Language Model"}, "llm": {"display_name": "Language Model"},
"code": {"advanced": True},
} }
def build( def build(
self, agent: str, llm: BaseLanguageModel, memory: BaseChatMemory, tools: List[Tool], max_iterations: int self,
agent: str,
llm: BaseLanguageModel,
tools: List[Tool],
max_iterations: int,
memory: Optional[BaseChatMemory] = None,
) -> Union[AgentExecutor, Callable]: ) -> Union[AgentExecutor, Callable]:
agent = AgentType(agent) agent = AgentType(agent)
if memory:
return initialize_agent( return initialize_agent(
tools=tools, tools=tools,
llm=llm, llm=llm,
@ -35,3 +41,12 @@ class AgentInitializerComponent(CustomComponent):
handle_parsing_errors=True, handle_parsing_errors=True,
max_iterations=max_iterations, max_iterations=max_iterations,
) )
else:
return initialize_agent(
tools=tools,
llm=llm,
agent=agent,
return_intermediate_steps=True,
handle_parsing_errors=True,
max_iterations=max_iterations,
)