test_event_kernel.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401
  1. from __future__ import annotations
  2. from typing import Any
  3. import pytest
  4. from agent_lab.application.events import (
  5. ConfirmationPolicy,
  6. EventDefinition,
  7. EventExecutionContext,
  8. EventKernel,
  9. EventRegistry,
  10. EventRequest,
  11. EventSource,
  12. EventStatus,
  13. ResultPolicy,
  14. RiskLevel,
  15. )
  16. from agent_lab.application.tools import (
  17. ToolDefinition,
  18. ToolExecutionContext,
  19. ToolRegistry,
  20. )
  21. from agent_lab.domain.events import ToolCallEvent
  22. from agent_lab.domain.messages import ChatMessage
  23. def _definition(
  24. name: str = "example.lookup",
  25. **overrides: Any,
  26. ) -> EventDefinition:
  27. values: dict[str, Any] = {
  28. "name": name,
  29. "description": "Look up an example value.",
  30. "parameters": {
  31. "type": "object",
  32. "properties": {"query": {"type": "string"}},
  33. "required": ["query"],
  34. },
  35. "handler": lambda request: {
  36. "event": request.name,
  37. "query": request.arguments["query"],
  38. },
  39. }
  40. values.update(overrides)
  41. return EventDefinition(**values)
  42. def test_registry_registers_flat_definitions_and_filters_enabled_catalog():
  43. registry = EventRegistry(
  44. [_definition("example.lookup"), _definition("device.inspect")]
  45. )
  46. assert [item["name"] for item in registry.catalog()] == [
  47. "example.lookup",
  48. "device.inspect",
  49. ]
  50. assert registry.catalog(["device.inspect"]) == [
  51. {
  52. "name": "device.inspect",
  53. "description": "Look up an example value.",
  54. "parameters": {
  55. "type": "object",
  56. "properties": {"query": {"type": "string"}},
  57. "required": ["query"],
  58. },
  59. }
  60. ]
  61. assert registry.tool_schema("example.lookup")["function"]["name"] == (
  62. "example.lookup"
  63. )
  64. assert registry.tool_schema("missing") is None
  65. def test_registry_rejects_duplicate_definition_names():
  66. with pytest.raises(ValueError, match="duplicate event definition: example.lookup"):
  67. EventRegistry([_definition(), _definition()])
  68. def test_tool_registry_public_api_remains_compatible():
  69. registry = ToolRegistry(
  70. [
  71. ToolDefinition(
  72. name="compat.lookup",
  73. description="Look up compatibility data.",
  74. parameters={
  75. "type": "object",
  76. "properties": {"query": {"type": "string"}},
  77. "required": ["query"],
  78. },
  79. handler=lambda event: {
  80. "tool": event.name,
  81. "query": event.arguments["query"],
  82. },
  83. argument_resolver=lambda event, context: {
  84. "query": context.history[-1].content
  85. },
  86. )
  87. ]
  88. )
  89. event = ToolCallEvent(
  90. id="call-1",
  91. name="compat.lookup",
  92. arguments={"query": "provider"},
  93. raw_arguments='{"query":"provider"}',
  94. )
  95. assert registry.available_tools() == [
  96. {
  97. "name": "compat.lookup",
  98. "description": "Look up compatibility data.",
  99. "parameters": {
  100. "type": "object",
  101. "properties": {"query": {"type": "string"}},
  102. "required": ["query"],
  103. },
  104. }
  105. ]
  106. assert "- compat.lookup: Look up compatibility data." in (
  107. registry.chat_event_system_message(["compat.lookup"])
  108. )
  109. assert registry.tool_schema("compat.lookup")["function"]["name"] == (
  110. "compat.lookup"
  111. )
  112. assert registry.handle(
  113. event,
  114. ToolExecutionContext(history=[ChatMessage(role="user", content="history")]),
  115. ) == {"tool": "compat.lookup", "query": "history"}
  116. assert registry.execute(event) == {"tool": "compat.lookup", "query": "provider"}
  117. @pytest.mark.asyncio
  118. async def test_kernel_executes_complete_deterministic_arguments_without_fallback():
  119. fallback_calls: list[str] = []
  120. async def fallback(*args: Any) -> dict[str, Any]:
  121. fallback_calls.append("called")
  122. return {"query": "fallback"}
  123. registry = EventRegistry(
  124. [_definition(resolver=lambda request, context: {"query": "deterministic"})]
  125. )
  126. result = await EventKernel(registry, argument_fallback=fallback).execute(
  127. EventRequest(id="event-1", name="example.lookup"),
  128. enabled_names=["example.lookup"],
  129. )
  130. assert result.status is EventStatus.SUCCESS
  131. assert result.arguments == {"query": "deterministic"}
  132. assert result.payload == {"event": "example.lookup", "query": "deterministic"}
  133. assert result.used_fallback is False
  134. assert fallback_calls == []
  135. @pytest.mark.asyncio
  136. async def test_kernel_calls_fallback_once_when_required_arguments_are_incomplete():
  137. fallback_calls: list[dict[str, Any]] = []
  138. async def fallback(
  139. definition: EventDefinition,
  140. request: EventRequest,
  141. context: EventExecutionContext,
  142. ) -> dict[str, Any]:
  143. fallback_calls.append(dict(request.arguments))
  144. return {"query": "resolved once"}
  145. registry = EventRegistry([_definition(resolver=lambda request, context: {})])
  146. result = await EventKernel(registry, argument_fallback=fallback).execute(
  147. EventRequest(id="event-1", name="example.lookup"),
  148. enabled_names=["example.lookup"],
  149. )
  150. assert result.status is EventStatus.SUCCESS
  151. assert result.arguments == {"query": "resolved once"}
  152. assert result.used_fallback is True
  153. assert fallback_calls == [{}]
  154. @pytest.mark.asyncio
  155. async def test_kernel_does_not_fallback_when_definition_disallows_it():
  156. fallback_calls = 0
  157. async def fallback(*args: Any) -> dict[str, Any]:
  158. nonlocal fallback_calls
  159. fallback_calls += 1
  160. return {"query": "not allowed"}
  161. registry = EventRegistry(
  162. [
  163. _definition(
  164. resolver=lambda request, context: {},
  165. fallback_allowed=False,
  166. )
  167. ]
  168. )
  169. result = await EventKernel(registry, argument_fallback=fallback).execute(
  170. EventRequest(id="event-1", name="example.lookup"),
  171. enabled_names=["example.lookup"],
  172. )
  173. assert result.status is EventStatus.INVALID_ARGUMENTS
  174. assert result.error == "missing required arguments: query"
  175. assert result.used_fallback is False
  176. assert fallback_calls == 0
  177. @pytest.mark.asyncio
  178. async def test_provider_resolved_arguments_are_not_rewritten_or_fallen_back():
  179. resolver_calls = 0
  180. fallback_calls = 0
  181. def resolver(*args: Any) -> dict[str, Any]:
  182. nonlocal resolver_calls
  183. resolver_calls += 1
  184. return {"query": "rewritten"}
  185. async def fallback(*args: Any) -> dict[str, Any]:
  186. nonlocal fallback_calls
  187. fallback_calls += 1
  188. return {"query": "fallback"}
  189. registry = EventRegistry([_definition(resolver=resolver)])
  190. result = await EventKernel(registry, argument_fallback=fallback).execute(
  191. EventRequest(
  192. id="event-1",
  193. name="example.lookup",
  194. arguments={"query": "provider value"},
  195. source=EventSource.PROVIDER_RESOLVED,
  196. ),
  197. enabled_names=["example.lookup"],
  198. )
  199. assert result.status is EventStatus.SUCCESS
  200. assert result.arguments == {"query": "provider value"}
  201. assert resolver_calls == 0
  202. assert fallback_calls == 0
  203. @pytest.mark.asyncio
  204. async def test_provider_resolved_missing_arguments_return_invalid_without_fallback():
  205. fallback_calls = 0
  206. async def fallback(*args: Any) -> dict[str, Any]:
  207. nonlocal fallback_calls
  208. fallback_calls += 1
  209. return {"query": "fallback"}
  210. result = await EventKernel(
  211. EventRegistry([_definition()]), argument_fallback=fallback
  212. ).execute(
  213. EventRequest(
  214. id="event-1",
  215. name="example.lookup",
  216. arguments={},
  217. source=EventSource.PROVIDER_RESOLVED,
  218. ),
  219. enabled_names=["example.lookup"],
  220. )
  221. assert result.status is EventStatus.INVALID_ARGUMENTS
  222. assert result.error == "missing required arguments: query"
  223. assert fallback_calls == 0
  224. @pytest.mark.asyncio
  225. @pytest.mark.parametrize(
  226. ("event_request", "enabled_names", "expected_status", "expected_error"),
  227. [
  228. (
  229. EventRequest(id="event-1", name="missing"),
  230. ["missing"],
  231. EventStatus.UNKNOWN,
  232. "unknown event",
  233. ),
  234. (
  235. EventRequest(id="event-1", name="example.lookup"),
  236. [],
  237. EventStatus.DISABLED,
  238. "event disabled",
  239. ),
  240. (
  241. EventRequest(
  242. id="event-1",
  243. name="example.lookup",
  244. arguments={"query": 42},
  245. source=EventSource.PROVIDER_RESOLVED,
  246. ),
  247. ["example.lookup"],
  248. EventStatus.INVALID_ARGUMENTS,
  249. "invalid argument type for query: expected string",
  250. ),
  251. ],
  252. )
  253. async def test_kernel_normalizes_lookup_and_validation_failures(
  254. event_request: EventRequest,
  255. enabled_names: list[str],
  256. expected_status: EventStatus,
  257. expected_error: str,
  258. ):
  259. result = await EventKernel(EventRegistry([_definition()])).execute(
  260. event_request,
  261. enabled_names=enabled_names,
  262. )
  263. assert result.status is expected_status
  264. assert result.error == expected_error
  265. @pytest.mark.asyncio
  266. async def test_kernel_normalizes_handler_exceptions():
  267. def fail(request: EventRequest) -> dict[str, Any]:
  268. raise RuntimeError("boom")
  269. result = await EventKernel(
  270. EventRegistry([_definition(handler=fail)])
  271. ).execute(
  272. EventRequest(
  273. id="event-1",
  274. name="example.lookup",
  275. arguments={"query": "value"},
  276. source=EventSource.PROVIDER_RESOLVED,
  277. ),
  278. enabled_names=["example.lookup"],
  279. )
  280. assert result.status is EventStatus.HANDLER_ERROR
  281. assert result.error == "event handler failed: boom"
  282. @pytest.mark.asyncio
  283. async def test_kernel_rejects_boolean_for_json_number_arguments():
  284. definition = _definition(
  285. parameters={
  286. "type": "object",
  287. "properties": {"query": {"type": "number"}},
  288. "required": ["query"],
  289. }
  290. )
  291. result = await EventKernel(EventRegistry([definition])).execute(
  292. EventRequest(
  293. id="event-1",
  294. name="example.lookup",
  295. arguments={"query": True},
  296. source=EventSource.PROVIDER_RESOLVED,
  297. ),
  298. enabled_names=["example.lookup"],
  299. )
  300. assert result.status is EventStatus.INVALID_ARGUMENTS
  301. assert result.error == "invalid argument type for query: expected number"
  302. @pytest.mark.asyncio
  303. async def test_definition_metadata_survives_registration_and_result_creation():
  304. definition = _definition(
  305. result_policy=ResultPolicy.TEMPLATE_FOLLOW_UP,
  306. confirmation_policy=ConfirmationPolicy.REQUIRED,
  307. risk_level=RiskLevel.HIGH,
  308. idempotency_key_fields=("session_id", "event_id"),
  309. concurrency_class="device-write",
  310. conflict_keys=("device",),
  311. timeout_seconds=1.5,
  312. terminal=True,
  313. )
  314. registry = EventRegistry([definition])
  315. result = await EventKernel(registry).execute(
  316. EventRequest(
  317. id="event-1",
  318. name="example.lookup",
  319. arguments={"query": "value"},
  320. source=EventSource.PROVIDER_RESOLVED,
  321. ),
  322. enabled_names=["example.lookup"],
  323. )
  324. registered = registry.definition("example.lookup")
  325. assert registered is definition
  326. assert result.result_policy is ResultPolicy.TEMPLATE_FOLLOW_UP
  327. assert result.confirmation_policy is ConfirmationPolicy.REQUIRED
  328. assert result.risk_level is RiskLevel.HIGH
  329. assert result.idempotency_key_fields == ("session_id", "event_id")
  330. assert result.concurrency_class == "device-write"
  331. assert result.conflict_keys == ("device",)
  332. assert result.timeout_seconds == 1.5
  333. assert result.terminal is True
  334. def test_kernel_source_has_no_builtin_event_name_branches():
  335. from pathlib import Path
  336. source = Path("src/agent_lab/application/events/kernel.py").read_text()
  337. assert "handoff_note" not in source
  338. assert "mock_search" not in source
  339. assert "mock_ticket" not in source