tools.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380
  1. import inspect
  2. from collections.abc import Awaitable, Callable, Iterable, Sequence
  3. from copy import deepcopy
  4. from dataclasses import dataclass
  5. from typing import Any
  6. from agent_lab.application.events import (
  7. ConfirmationPolicy,
  8. CalendarSchedulePort,
  9. DeviceVolumePort,
  10. EventDefinition,
  11. EventExecutionContext,
  12. EventKernel,
  13. EventRegistry,
  14. EventRequest,
  15. EventResult,
  16. EventSource,
  17. EventStatus,
  18. ResultPolicy,
  19. RiskLevel,
  20. SessionTerminationPort,
  21. WebSearchPort,
  22. build_builtin_event_definitions,
  23. )
  24. from agent_lab.domain.events import EVENT_BLOCK_END, EVENT_BLOCK_START, ToolCallEvent
  25. from agent_lab.domain.messages import ChatMessage
  26. ToolExecutionContext = EventExecutionContext
  27. ToolHandler = Callable[
  28. [ToolCallEvent],
  29. dict[str, Any] | Awaitable[dict[str, Any]],
  30. ]
  31. ToolArgumentResolver = Callable[
  32. [ToolCallEvent, ToolExecutionContext],
  33. dict[str, Any],
  34. ]
  35. ToolArgumentNormalizer = Callable[[dict[str, Any]], dict[str, Any]]
  36. @dataclass(frozen=True)
  37. class ToolDefinition:
  38. name: str
  39. description: str
  40. parameters: dict[str, Any]
  41. handler: ToolHandler
  42. argument_resolver: ToolArgumentResolver | None = None
  43. schema_version: str = "1"
  44. fallback_allowed: bool = True
  45. result_policy: ResultPolicy = ResultPolicy.LLM_FOLLOW_UP
  46. confirmation_policy: ConfirmationPolicy = ConfirmationPolicy.NONE
  47. risk_level: RiskLevel = RiskLevel.LOW
  48. idempotency_key_fields: tuple[str, ...] = ()
  49. concurrency_class: str | None = None
  50. conflict_keys: tuple[str, ...] = ()
  51. timeout_seconds: float | None = None
  52. terminal: bool = False
  53. normalizer: ToolArgumentNormalizer | None = None
  54. class ToolRegistry:
  55. def __init__(
  56. self,
  57. definitions: Iterable[ToolDefinition | EventDefinition],
  58. ) -> None:
  59. self._definitions: dict[str, EventDefinition] = {}
  60. event_definitions: list[EventDefinition] = []
  61. for definition in definitions:
  62. if definition.name in self._definitions:
  63. raise ValueError(f"duplicate event definition: {definition.name}")
  64. event_definition = (
  65. definition
  66. if isinstance(definition, EventDefinition)
  67. else self._to_event_definition(definition)
  68. )
  69. self._definitions[definition.name] = event_definition
  70. event_definitions.append(event_definition)
  71. self.event_registry = EventRegistry(event_definitions)
  72. self.kernel = EventKernel(self.event_registry)
  73. def available_tools(self) -> list[dict[str, Any]]:
  74. return self.event_registry.catalog()
  75. def chat_event_system_message(self, enabled_names: Iterable[str]) -> str:
  76. enabled = set(enabled_names)
  77. event_lines = [
  78. f"- {definition.name}: {definition.description}"
  79. for definition in self._definitions.values()
  80. if definition.name in enabled
  81. ]
  82. if not event_lines:
  83. return ""
  84. return "\n".join(
  85. [
  86. "You may request EventAgent work using this text protocol.",
  87. "Available events:",
  88. *event_lines,
  89. (
  90. "First write the user-facing reply normally. If events are needed, "
  91. "append this block after the visible reply:"
  92. ),
  93. EVENT_BLOCK_START,
  94. "event_name",
  95. EVENT_BLOCK_END,
  96. "Use exact event names only, one per line. Do not include parameters.",
  97. ]
  98. )
  99. def tool_schema(self, name: str) -> dict[str, Any] | None:
  100. return self.event_registry.tool_schema(name)
  101. def provider_tool_schemas(
  102. self,
  103. enabled_names: Iterable[str],
  104. ) -> list[dict[str, Any]]:
  105. return [
  106. schema
  107. for schema in (
  108. self.tool_schema(name)
  109. for name in enabled_names
  110. )
  111. if schema is not None
  112. ]
  113. def handle(
  114. self,
  115. event: ToolCallEvent,
  116. context: ToolExecutionContext | None = None,
  117. ) -> dict[str, Any]:
  118. result = self.kernel.execute_sync(
  119. self.event_request(event),
  120. context=context or ToolExecutionContext(history=()),
  121. )
  122. return self.tool_payload(result)
  123. def execute(self, event: ToolCallEvent) -> dict[str, Any]:
  124. result = self.kernel.execute_sync(
  125. self.event_request(event, source=EventSource.PROVIDER_RESOLVED)
  126. )
  127. return self.tool_payload(result)
  128. async def handle_async(
  129. self,
  130. event: ToolCallEvent,
  131. context: ToolExecutionContext | None = None,
  132. ) -> dict[str, Any]:
  133. result = await self.kernel.execute(
  134. self.event_request(event),
  135. context=context or ToolExecutionContext(history=()),
  136. )
  137. return self.tool_payload(result)
  138. async def execute_async(
  139. self,
  140. event: ToolCallEvent,
  141. *,
  142. enabled_names: Iterable[str] | None = None,
  143. ) -> dict[str, Any]:
  144. result = await self.kernel.execute(
  145. self.event_request(event, source=EventSource.PROVIDER_RESOLVED),
  146. enabled_names=enabled_names,
  147. )
  148. return self.tool_payload(result)
  149. def event_request(
  150. self,
  151. event: ToolCallEvent,
  152. source: EventSource = EventSource.TEXT_EVENT,
  153. ) -> EventRequest:
  154. return EventRequest(
  155. id=event.id,
  156. name=event.name,
  157. arguments=deepcopy(event.arguments),
  158. raw_arguments=event.raw_arguments,
  159. source=source,
  160. )
  161. def tool_payload(self, result: EventResult) -> dict[str, Any]:
  162. if result.status is EventStatus.SUCCESS:
  163. return result.payload
  164. error = result.error or result.status.value
  165. if result.status is EventStatus.UNKNOWN:
  166. error = "unknown tool"
  167. elif result.status is EventStatus.DISABLED:
  168. error = "tool disabled"
  169. elif result.status is EventStatus.DEFINITION_ERROR:
  170. error = error.replace(
  171. "event argument validation failed:",
  172. "tool definition validation failed:",
  173. 1,
  174. )
  175. elif result.status is EventStatus.HANDLER_ERROR:
  176. error = error.replace("event handler failed:", "tool handler failed:", 1)
  177. error = error.replace("EventKernel.execute", "execute_async")
  178. return {"tool": result.event_name, "error": error}
  179. def _to_event_definition(self, definition: ToolDefinition) -> EventDefinition:
  180. def resolver(
  181. request: EventRequest,
  182. context: EventExecutionContext,
  183. ) -> dict[str, Any]:
  184. event = self._tool_event(request)
  185. if definition.argument_resolver is not None:
  186. return definition.argument_resolver(event, context)
  187. if context.history:
  188. return {}
  189. return deepcopy(request.arguments)
  190. if inspect.iscoroutinefunction(definition.handler):
  191. async def handler(request: EventRequest) -> dict[str, Any]:
  192. return await definition.handler(self._tool_event(request))
  193. else:
  194. def handler(
  195. request: EventRequest,
  196. ) -> dict[str, Any] | Awaitable[dict[str, Any]]:
  197. return definition.handler(self._tool_event(request))
  198. return EventDefinition(
  199. name=definition.name,
  200. description=definition.description,
  201. parameters=deepcopy(definition.parameters),
  202. handler=handler,
  203. resolver=resolver,
  204. normalizer=definition.normalizer,
  205. schema_version=definition.schema_version,
  206. fallback_allowed=definition.fallback_allowed,
  207. result_policy=definition.result_policy,
  208. confirmation_policy=definition.confirmation_policy,
  209. risk_level=definition.risk_level,
  210. idempotency_key_fields=definition.idempotency_key_fields,
  211. concurrency_class=definition.concurrency_class,
  212. conflict_keys=definition.conflict_keys,
  213. timeout_seconds=definition.timeout_seconds,
  214. terminal=definition.terminal,
  215. )
  216. def _tool_event(self, request: EventRequest) -> ToolCallEvent:
  217. return ToolCallEvent(
  218. id=request.id,
  219. name=request.name,
  220. arguments=deepcopy(request.arguments),
  221. raw_arguments=request.raw_arguments,
  222. )
  223. def build_default_tool_registry(
  224. *,
  225. session_termination_port: SessionTerminationPort | None = None,
  226. device_volume_port: DeviceVolumePort | None = None,
  227. calendar_schedule_port: CalendarSchedulePort | None = None,
  228. web_search_port: WebSearchPort | None = None,
  229. clock: Callable[[], Any] | None = None,
  230. ) -> ToolRegistry:
  231. return ToolRegistry(
  232. [
  233. ToolDefinition(
  234. name="handoff_note",
  235. description="Send a note to the event agent.",
  236. parameters={
  237. "type": "object",
  238. "properties": {
  239. "message": {"type": "string"},
  240. },
  241. "required": ["message"],
  242. },
  243. handler=_handle_handoff_note,
  244. argument_resolver=_resolve_handoff_note_arguments,
  245. ),
  246. ToolDefinition(
  247. name="mock_search",
  248. description="Search mock external knowledge for the current turn.",
  249. parameters={
  250. "type": "object",
  251. "properties": {
  252. "query": {"type": "string"},
  253. },
  254. "required": ["query"],
  255. },
  256. handler=_handle_mock_search,
  257. argument_resolver=_resolve_mock_search_arguments,
  258. ),
  259. ToolDefinition(
  260. name="mock_ticket",
  261. description="Create a mock support ticket for follow-up work.",
  262. parameters={
  263. "type": "object",
  264. "properties": {
  265. "title": {"type": "string"},
  266. },
  267. "required": ["title"],
  268. },
  269. handler=_handle_mock_ticket,
  270. argument_resolver=_resolve_mock_ticket_arguments,
  271. ),
  272. *build_builtin_event_definitions(
  273. session_termination_port=session_termination_port,
  274. device_volume_port=device_volume_port,
  275. calendar_schedule_port=calendar_schedule_port,
  276. web_search_port=web_search_port,
  277. clock=clock,
  278. ),
  279. ]
  280. )
  281. def _resolve_handoff_note_arguments(
  282. event: ToolCallEvent,
  283. context: ToolExecutionContext,
  284. ) -> dict[str, Any]:
  285. content = _latest_content(context.history, preferred_roles=("assistant", "user"))
  286. if not content and not context.history:
  287. content = str(event.arguments.get("message", ""))
  288. return {"message": content}
  289. def _resolve_mock_search_arguments(
  290. event: ToolCallEvent,
  291. context: ToolExecutionContext,
  292. ) -> dict[str, Any]:
  293. query = _latest_content(context.history, preferred_roles=("assistant", "user"))
  294. if not query and not context.history:
  295. query = str(event.arguments.get("query", ""))
  296. return {"query": query}
  297. def _resolve_mock_ticket_arguments(
  298. event: ToolCallEvent,
  299. context: ToolExecutionContext,
  300. ) -> dict[str, Any]:
  301. title = _latest_content(context.history, preferred_roles=("assistant", "user"))
  302. if not title and not context.history:
  303. title = str(event.arguments.get("title", ""))
  304. return {"title": title}
  305. def _latest_content(
  306. history: Sequence[ChatMessage],
  307. preferred_roles: tuple[str, ...],
  308. ) -> str:
  309. for role in preferred_roles:
  310. for message in reversed(history):
  311. if message.role == role and message.content.strip():
  312. return message.content.strip()
  313. return ""
  314. def _handle_handoff_note(event: ToolCallEvent) -> dict[str, Any]:
  315. return {
  316. "tool": "handoff_note",
  317. "message": str(event.arguments.get("message", "")),
  318. }
  319. def _handle_mock_search(event: ToolCallEvent) -> dict[str, Any]:
  320. query = str(event.arguments.get("query", ""))
  321. return {
  322. "tool": "mock_search",
  323. "query": query,
  324. "results": [
  325. {
  326. "title": "Mock knowledge base result",
  327. "snippet": f"Simulated external context for: {query}",
  328. }
  329. ],
  330. }
  331. def _handle_mock_ticket(event: ToolCallEvent) -> dict[str, Any]:
  332. title = str(event.arguments.get("title", ""))
  333. return {
  334. "tool": "mock_ticket",
  335. "ticket_id": "MOCK-001",
  336. "title": title,
  337. "status": "created",
  338. }