| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380 |
- import inspect
- from collections.abc import Awaitable, Callable, Iterable, Sequence
- from copy import deepcopy
- from dataclasses import dataclass
- from typing import Any
- from agent_lab.application.events import (
- ConfirmationPolicy,
- CalendarSchedulePort,
- DeviceVolumePort,
- EventDefinition,
- EventExecutionContext,
- EventKernel,
- EventRegistry,
- EventRequest,
- EventResult,
- EventSource,
- EventStatus,
- ResultPolicy,
- RiskLevel,
- SessionTerminationPort,
- WebSearchPort,
- build_builtin_event_definitions,
- )
- from agent_lab.domain.events import EVENT_BLOCK_END, EVENT_BLOCK_START, ToolCallEvent
- from agent_lab.domain.messages import ChatMessage
- ToolExecutionContext = EventExecutionContext
- ToolHandler = Callable[
- [ToolCallEvent],
- dict[str, Any] | Awaitable[dict[str, Any]],
- ]
- ToolArgumentResolver = Callable[
- [ToolCallEvent, ToolExecutionContext],
- dict[str, Any],
- ]
- ToolArgumentNormalizer = Callable[[dict[str, Any]], dict[str, Any]]
- @dataclass(frozen=True)
- class ToolDefinition:
- name: str
- description: str
- parameters: dict[str, Any]
- handler: ToolHandler
- argument_resolver: ToolArgumentResolver | None = None
- schema_version: str = "1"
- fallback_allowed: bool = True
- result_policy: ResultPolicy = ResultPolicy.LLM_FOLLOW_UP
- confirmation_policy: ConfirmationPolicy = ConfirmationPolicy.NONE
- risk_level: RiskLevel = RiskLevel.LOW
- idempotency_key_fields: tuple[str, ...] = ()
- concurrency_class: str | None = None
- conflict_keys: tuple[str, ...] = ()
- timeout_seconds: float | None = None
- terminal: bool = False
- normalizer: ToolArgumentNormalizer | None = None
- class ToolRegistry:
- def __init__(
- self,
- definitions: Iterable[ToolDefinition | EventDefinition],
- ) -> None:
- self._definitions: dict[str, EventDefinition] = {}
- event_definitions: list[EventDefinition] = []
- for definition in definitions:
- if definition.name in self._definitions:
- raise ValueError(f"duplicate event definition: {definition.name}")
- event_definition = (
- definition
- if isinstance(definition, EventDefinition)
- else self._to_event_definition(definition)
- )
- self._definitions[definition.name] = event_definition
- event_definitions.append(event_definition)
- self.event_registry = EventRegistry(event_definitions)
- self.kernel = EventKernel(self.event_registry)
- def available_tools(self) -> list[dict[str, Any]]:
- return self.event_registry.catalog()
- def chat_event_system_message(self, enabled_names: Iterable[str]) -> str:
- enabled = set(enabled_names)
- event_lines = [
- f"- {definition.name}: {definition.description}"
- for definition in self._definitions.values()
- if definition.name in enabled
- ]
- if not event_lines:
- return ""
- return "\n".join(
- [
- "You may request EventAgent work using this text protocol.",
- "Available events:",
- *event_lines,
- (
- "First write the user-facing reply normally. If events are needed, "
- "append this block after the visible reply:"
- ),
- EVENT_BLOCK_START,
- "event_name",
- EVENT_BLOCK_END,
- "Use exact event names only, one per line. Do not include parameters.",
- ]
- )
- def tool_schema(self, name: str) -> dict[str, Any] | None:
- return self.event_registry.tool_schema(name)
- def provider_tool_schemas(
- self,
- enabled_names: Iterable[str],
- ) -> list[dict[str, Any]]:
- return [
- schema
- for schema in (
- self.tool_schema(name)
- for name in enabled_names
- )
- if schema is not None
- ]
- def handle(
- self,
- event: ToolCallEvent,
- context: ToolExecutionContext | None = None,
- ) -> dict[str, Any]:
- result = self.kernel.execute_sync(
- self.event_request(event),
- context=context or ToolExecutionContext(history=()),
- )
- return self.tool_payload(result)
- def execute(self, event: ToolCallEvent) -> dict[str, Any]:
- result = self.kernel.execute_sync(
- self.event_request(event, source=EventSource.PROVIDER_RESOLVED)
- )
- return self.tool_payload(result)
- async def handle_async(
- self,
- event: ToolCallEvent,
- context: ToolExecutionContext | None = None,
- ) -> dict[str, Any]:
- result = await self.kernel.execute(
- self.event_request(event),
- context=context or ToolExecutionContext(history=()),
- )
- return self.tool_payload(result)
- async def execute_async(
- self,
- event: ToolCallEvent,
- *,
- enabled_names: Iterable[str] | None = None,
- ) -> dict[str, Any]:
- result = await self.kernel.execute(
- self.event_request(event, source=EventSource.PROVIDER_RESOLVED),
- enabled_names=enabled_names,
- )
- return self.tool_payload(result)
- def event_request(
- self,
- event: ToolCallEvent,
- source: EventSource = EventSource.TEXT_EVENT,
- ) -> EventRequest:
- return EventRequest(
- id=event.id,
- name=event.name,
- arguments=deepcopy(event.arguments),
- raw_arguments=event.raw_arguments,
- source=source,
- )
- def tool_payload(self, result: EventResult) -> dict[str, Any]:
- if result.status is EventStatus.SUCCESS:
- return result.payload
- error = result.error or result.status.value
- if result.status is EventStatus.UNKNOWN:
- error = "unknown tool"
- elif result.status is EventStatus.DISABLED:
- error = "tool disabled"
- elif result.status is EventStatus.DEFINITION_ERROR:
- error = error.replace(
- "event argument validation failed:",
- "tool definition validation failed:",
- 1,
- )
- elif result.status is EventStatus.HANDLER_ERROR:
- error = error.replace("event handler failed:", "tool handler failed:", 1)
- error = error.replace("EventKernel.execute", "execute_async")
- return {"tool": result.event_name, "error": error}
- def _to_event_definition(self, definition: ToolDefinition) -> EventDefinition:
- def resolver(
- request: EventRequest,
- context: EventExecutionContext,
- ) -> dict[str, Any]:
- event = self._tool_event(request)
- if definition.argument_resolver is not None:
- return definition.argument_resolver(event, context)
- if context.history:
- return {}
- return deepcopy(request.arguments)
- if inspect.iscoroutinefunction(definition.handler):
- async def handler(request: EventRequest) -> dict[str, Any]:
- return await definition.handler(self._tool_event(request))
- else:
- def handler(
- request: EventRequest,
- ) -> dict[str, Any] | Awaitable[dict[str, Any]]:
- return definition.handler(self._tool_event(request))
- return EventDefinition(
- name=definition.name,
- description=definition.description,
- parameters=deepcopy(definition.parameters),
- handler=handler,
- resolver=resolver,
- normalizer=definition.normalizer,
- schema_version=definition.schema_version,
- fallback_allowed=definition.fallback_allowed,
- result_policy=definition.result_policy,
- confirmation_policy=definition.confirmation_policy,
- risk_level=definition.risk_level,
- idempotency_key_fields=definition.idempotency_key_fields,
- concurrency_class=definition.concurrency_class,
- conflict_keys=definition.conflict_keys,
- timeout_seconds=definition.timeout_seconds,
- terminal=definition.terminal,
- )
- def _tool_event(self, request: EventRequest) -> ToolCallEvent:
- return ToolCallEvent(
- id=request.id,
- name=request.name,
- arguments=deepcopy(request.arguments),
- raw_arguments=request.raw_arguments,
- )
- def build_default_tool_registry(
- *,
- session_termination_port: SessionTerminationPort | None = None,
- device_volume_port: DeviceVolumePort | None = None,
- calendar_schedule_port: CalendarSchedulePort | None = None,
- web_search_port: WebSearchPort | None = None,
- clock: Callable[[], Any] | None = None,
- ) -> ToolRegistry:
- return ToolRegistry(
- [
- ToolDefinition(
- name="handoff_note",
- description="Send a note to the event agent.",
- parameters={
- "type": "object",
- "properties": {
- "message": {"type": "string"},
- },
- "required": ["message"],
- },
- handler=_handle_handoff_note,
- argument_resolver=_resolve_handoff_note_arguments,
- ),
- ToolDefinition(
- name="mock_search",
- description="Search mock external knowledge for the current turn.",
- parameters={
- "type": "object",
- "properties": {
- "query": {"type": "string"},
- },
- "required": ["query"],
- },
- handler=_handle_mock_search,
- argument_resolver=_resolve_mock_search_arguments,
- ),
- ToolDefinition(
- name="mock_ticket",
- description="Create a mock support ticket for follow-up work.",
- parameters={
- "type": "object",
- "properties": {
- "title": {"type": "string"},
- },
- "required": ["title"],
- },
- handler=_handle_mock_ticket,
- argument_resolver=_resolve_mock_ticket_arguments,
- ),
- *build_builtin_event_definitions(
- session_termination_port=session_termination_port,
- device_volume_port=device_volume_port,
- calendar_schedule_port=calendar_schedule_port,
- web_search_port=web_search_port,
- clock=clock,
- ),
- ]
- )
- def _resolve_handoff_note_arguments(
- event: ToolCallEvent,
- context: ToolExecutionContext,
- ) -> dict[str, Any]:
- 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,
- context: ToolExecutionContext,
- ) -> dict[str, Any]:
- 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,
- context: ToolExecutionContext,
- ) -> dict[str, Any]:
- 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}
- def _latest_content(
- history: Sequence[ChatMessage],
- preferred_roles: tuple[str, ...],
- ) -> str:
- for role in preferred_roles:
- for message in reversed(history):
- if message.role == role and message.content.strip():
- return message.content.strip()
- return ""
- def _handle_handoff_note(event: ToolCallEvent) -> dict[str, Any]:
- return {
- "tool": "handoff_note",
- "message": str(event.arguments.get("message", "")),
- }
- def _handle_mock_search(event: ToolCallEvent) -> dict[str, Any]:
- query = str(event.arguments.get("query", ""))
- return {
- "tool": "mock_search",
- "query": query,
- "results": [
- {
- "title": "Mock knowledge base result",
- "snippet": f"Simulated external context for: {query}",
- }
- ],
- }
- def _handle_mock_ticket(event: ToolCallEvent) -> dict[str, Any]:
- title = str(event.arguments.get("title", ""))
- return {
- "tool": "mock_ticket",
- "ticket_id": "MOCK-001",
- "title": title,
- "status": "created",
- }
|