fix: Add get_base_args method and refactor component initialization in Agent (#6026)

* feat: Add get_base_args method to Component class

Introduces a new method to retrieve base initialization arguments for components, including user ID, session ID, and tracing service. This method provides a convenient way to access essential context information during component initialization.

* refactor: Update AgentComponent to use get_base_args method

Modify AgentComponent to pass base initialization arguments when creating CurrentDateComponent and MemoryComponent, ensuring consistent context initialization across components.

* [autofix.ci] apply automated fixes

---------

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Gabriel Luiz Freitas Almeida 2025-01-31 08:58:30 -03:00 • committed by GitHub
commit 58043362b5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 28 additions and 14 deletions

View file

@ -84,8 +84,7 @@ class AgentComponent(ToolCallingAgentComponent):
if not isinstance(self.tools, list): # type: ignore[has-type]
self.tools = []
# Convert CurrentDateComponent to a StructuredTool
current_date_tool = (await CurrentDateComponent().to_toolkit()).pop(0)
# current_date_tool = CurrentDateComponent().to_toolkit()[0]
current_date_tool = (await CurrentDateComponent(**self.get_base_args()).to_toolkit()).pop(0)
if isinstance(current_date_tool, StructuredTool):
self.tools.append(current_date_tool)
else:
@ -122,7 +121,7 @@ class AgentComponent(ToolCallingAgentComponent):
# filter out empty values
memory_kwargs = {k: v for k, v in memory_kwargs.items() if v}
return await MemoryComponent().set(**memory_kwargs).retrieve_messages()
return await MemoryComponent(**self.get_base_args()).set(**memory_kwargs).retrieve_messages()
def get_llm(self):
if isinstance(self.agent_llm, str):

View file

@ -174,6 +174,21 @@ class Component(CustomComponent):
# Return the intersection of the sets
return input_names & output_names
def get_base_args(self):
"""Get the base arguments required for component initialization.
Returns:
dict: A dictionary containing the base arguments:
- _user_id: The ID of the current user
- _session_id: The ID of the current session
- _tracing_service: The tracing service instance for logging/monitoring
"""
return {
"_user_id": self.user_id,
"_session_id": self.session_id,
"_tracing_service": self._tracing_service,
}
@property
def ctx(self):
if not hasattr(self, "graph") or self.graph is None:

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long