|
|
@@ -7,12 +7,6 @@ from agent_lab.domain.events import ToolCallEvent
|
|
|
from agent_lab.domain.messages import ChatMessage
|
|
|
|
|
|
|
|
|
-ToolHandler = Callable[[ToolCallEvent], dict[str, Any]]
|
|
|
-ToolArgumentResolver = Callable[
|
|
|
- [ToolCallEvent, Sequence[ChatMessage]],
|
|
|
- dict[str, Any],
|
|
|
-]
|
|
|
-
|
|
|
EVENT_NAME_ONLY_PARAMETERS: dict[str, Any] = {
|
|
|
"type": "object",
|
|
|
"properties": {},
|
|
|
@@ -20,6 +14,20 @@ EVENT_NAME_ONLY_PARAMETERS: dict[str, Any] = {
|
|
|
}
|
|
|
|
|
|
|
|
|
+@dataclass(frozen=True)
|
|
|
+class ToolExecutionContext:
|
|
|
+ history: Sequence[ChatMessage]
|
|
|
+ system_prompt: str = ""
|
|
|
+ extra_body: dict[str, Any] | None = None
|
|
|
+
|
|
|
+
|
|
|
+ToolHandler = Callable[[ToolCallEvent], dict[str, Any]]
|
|
|
+ToolArgumentResolver = Callable[
|
|
|
+ [ToolCallEvent, ToolExecutionContext],
|
|
|
+ dict[str, Any],
|
|
|
+]
|
|
|
+
|
|
|
+
|
|
|
@dataclass(frozen=True)
|
|
|
class ToolDefinition:
|
|
|
name: str
|
|
|
@@ -61,7 +69,7 @@ class ToolRegistry:
|
|
|
def handle(
|
|
|
self,
|
|
|
event: ToolCallEvent,
|
|
|
- history: Sequence[ChatMessage] = (),
|
|
|
+ context: ToolExecutionContext | None = None,
|
|
|
) -> dict[str, Any]:
|
|
|
definition = self._definitions.get(event.name)
|
|
|
if definition is None:
|
|
|
@@ -69,9 +77,14 @@ class ToolRegistry:
|
|
|
"tool": event.name,
|
|
|
"error": "unknown tool",
|
|
|
}
|
|
|
+ resolved_context = context or ToolExecutionContext(history=())
|
|
|
resolved_event = event.model_copy(
|
|
|
update={
|
|
|
- "arguments": self._resolve_arguments(definition, event, history),
|
|
|
+ "arguments": self._resolve_arguments(
|
|
|
+ definition,
|
|
|
+ event,
|
|
|
+ resolved_context,
|
|
|
+ ),
|
|
|
}
|
|
|
)
|
|
|
return definition.handler(resolved_event)
|
|
|
@@ -80,11 +93,11 @@ class ToolRegistry:
|
|
|
self,
|
|
|
definition: ToolDefinition,
|
|
|
event: ToolCallEvent,
|
|
|
- history: Sequence[ChatMessage],
|
|
|
+ context: ToolExecutionContext,
|
|
|
) -> dict[str, Any]:
|
|
|
if definition.argument_resolver is not None:
|
|
|
- return definition.argument_resolver(event, history)
|
|
|
- if history:
|
|
|
+ return definition.argument_resolver(event, context)
|
|
|
+ if context.history:
|
|
|
return {}
|
|
|
return deepcopy(event.arguments)
|
|
|
|
|
|
@@ -137,30 +150,30 @@ def build_default_tool_registry() -> ToolRegistry:
|
|
|
|
|
|
def _resolve_handoff_note_arguments(
|
|
|
event: ToolCallEvent,
|
|
|
- history: Sequence[ChatMessage],
|
|
|
+ context: ToolExecutionContext,
|
|
|
) -> dict[str, Any]:
|
|
|
- content = _latest_content(history, preferred_roles=("assistant", "user"))
|
|
|
- if not content and not history:
|
|
|
+ content = _latest_content(context.history, preferred_roles=("assistant", "user"))
|
|
|
+ if not content and not context.history:
|
|
|
content = str(event.arguments.get("message", ""))
|
|
|
return {"message": content}
|
|
|
|
|
|
|
|
|
def _resolve_mock_search_arguments(
|
|
|
event: ToolCallEvent,
|
|
|
- history: Sequence[ChatMessage],
|
|
|
+ context: ToolExecutionContext,
|
|
|
) -> dict[str, Any]:
|
|
|
- query = _latest_content(history, preferred_roles=("assistant", "user"))
|
|
|
- if not query and not history:
|
|
|
+ query = _latest_content(context.history, preferred_roles=("assistant", "user"))
|
|
|
+ if not query and not context.history:
|
|
|
query = str(event.arguments.get("query", ""))
|
|
|
return {"query": query}
|
|
|
|
|
|
|
|
|
def _resolve_mock_ticket_arguments(
|
|
|
event: ToolCallEvent,
|
|
|
- history: Sequence[ChatMessage],
|
|
|
+ context: ToolExecutionContext,
|
|
|
) -> dict[str, Any]:
|
|
|
- title = _latest_content(history, preferred_roles=("assistant", "user"))
|
|
|
- if not title and not history:
|
|
|
+ title = _latest_content(context.history, preferred_roles=("assistant", "user"))
|
|
|
+ if not title and not context.history:
|
|
|
title = str(event.arguments.get("title", ""))
|
|
|
return {"title": title}
|
|
|
|