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:
parent
2f9cd3e40b
commit
58043362b5
9 changed files with 28 additions and 14 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Add a link
Reference in a new issue