kernel.py 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516
  1. from __future__ import annotations
  2. import inspect
  3. import json
  4. from collections.abc import Awaitable, Callable, Iterable
  5. from dataclasses import dataclass, replace
  6. from typing import Any
  7. from jsonschema.exceptions import ValidationError
  8. from agent_lab.application.events.models import (
  9. EventDefinition,
  10. EventArgumentResolution,
  11. EventExecutionContext,
  12. EventRequest,
  13. EventResult,
  14. EventSource,
  15. EventStatus,
  16. ResolvedEventArguments,
  17. )
  18. from agent_lab.application.events.registry import EventRegistry
  19. ArgumentFallbackResult = ResolvedEventArguments | None
  20. ArgumentFallback = Callable[
  21. [EventDefinition, EventRequest, EventExecutionContext],
  22. ArgumentFallbackResult | Awaitable[ArgumentFallbackResult],
  23. ]
  24. @dataclass(frozen=True)
  25. class _ValidationResult:
  26. status: EventStatus | None = None
  27. error: str | None = None
  28. missing_required: bool = False
  29. class EventKernel:
  30. def __init__(
  31. self,
  32. registry: EventRegistry,
  33. argument_fallback: ArgumentFallback | None = None,
  34. ) -> None:
  35. self.registry = registry
  36. self.argument_fallback = argument_fallback
  37. async def execute(
  38. self,
  39. request: EventRequest,
  40. *,
  41. enabled_names: Iterable[str] | None = None,
  42. context: EventExecutionContext | None = None,
  43. ) -> EventResult:
  44. definition, early_result = self._lookup(request, enabled_names)
  45. if early_result is not None:
  46. return early_result
  47. assert definition is not None
  48. resolved_context = context or EventExecutionContext()
  49. resolved, resolution_complete, resolution_error = (
  50. self._resolve_deterministic(
  51. definition,
  52. request,
  53. resolved_context,
  54. )
  55. )
  56. if resolution_error is not None:
  57. return resolution_error
  58. assert resolved is not None
  59. validation = self._validate(definition, resolved.arguments)
  60. if not resolution_complete and validation.status is None:
  61. validation = _ValidationResult(
  62. status=EventStatus.INVALID_ARGUMENTS,
  63. error="event arguments incomplete",
  64. )
  65. used_fallback = False
  66. if (
  67. validation.status is not EventStatus.DEFINITION_ERROR
  68. and (validation.missing_required or not resolution_complete)
  69. and request.source is not EventSource.PROVIDER_RESOLVED
  70. and definition.fallback_allowed
  71. and self.argument_fallback is not None
  72. ):
  73. used_fallback = True
  74. fallback_request = replace(
  75. request,
  76. arguments=resolved.arguments,
  77. raw_arguments=resolved.raw_arguments,
  78. )
  79. try:
  80. fallback_value = self.argument_fallback(
  81. definition,
  82. fallback_request,
  83. resolved_context,
  84. )
  85. if inspect.isawaitable(fallback_value):
  86. fallback_value = await fallback_value
  87. except Exception as exc:
  88. return self._resolution_error(
  89. definition,
  90. request,
  91. f"event argument fallback failed: {exc}",
  92. resolved=resolved,
  93. used_fallback=True,
  94. )
  95. if fallback_value is not None:
  96. resolved, fallback_error = self._normalize_fallback(
  97. definition,
  98. request,
  99. fallback_value,
  100. )
  101. if fallback_error is not None:
  102. return fallback_error
  103. assert resolved is not None
  104. validation = self._validate(
  105. definition,
  106. resolved.arguments,
  107. )
  108. if validation.status is not None:
  109. return self._result(
  110. definition,
  111. request,
  112. validation.status,
  113. resolved=resolved,
  114. error=validation.error,
  115. used_fallback=used_fallback,
  116. )
  117. resolved_request = replace(
  118. request,
  119. arguments=resolved.arguments,
  120. raw_arguments=resolved.raw_arguments,
  121. )
  122. try:
  123. payload = definition.handler(resolved_request)
  124. if inspect.isawaitable(payload):
  125. payload = await payload
  126. except Exception as exc:
  127. return self._result(
  128. definition,
  129. request,
  130. EventStatus.HANDLER_ERROR,
  131. resolved=resolved,
  132. error=f"event handler failed: {exc}",
  133. used_fallback=used_fallback,
  134. )
  135. if not isinstance(payload, dict):
  136. return self._result(
  137. definition,
  138. request,
  139. EventStatus.HANDLER_ERROR,
  140. resolved=resolved,
  141. error="event handler returned non-object payload",
  142. used_fallback=used_fallback,
  143. )
  144. return self._result(
  145. definition,
  146. request,
  147. EventStatus.SUCCESS,
  148. resolved=resolved,
  149. payload=dict(payload),
  150. used_fallback=used_fallback,
  151. )
  152. def execute_sync(
  153. self,
  154. request: EventRequest,
  155. *,
  156. enabled_names: Iterable[str] | None = None,
  157. context: EventExecutionContext | None = None,
  158. ) -> EventResult:
  159. definition, early_result = self._lookup(request, enabled_names)
  160. if early_result is not None:
  161. return early_result
  162. assert definition is not None
  163. resolved, resolution_complete, resolution_error = self._resolve_deterministic(
  164. definition,
  165. request,
  166. context or EventExecutionContext(),
  167. )
  168. if resolution_error is not None:
  169. return resolution_error
  170. assert resolved is not None
  171. validation = self._validate(definition, resolved.arguments)
  172. if not resolution_complete and validation.status is None:
  173. validation = _ValidationResult(
  174. status=EventStatus.INVALID_ARGUMENTS,
  175. error="event arguments incomplete",
  176. )
  177. if validation.status is not None:
  178. return self._result(
  179. definition,
  180. request,
  181. validation.status,
  182. resolved=resolved,
  183. error=validation.error,
  184. )
  185. if inspect.iscoroutinefunction(definition.handler):
  186. return self._result(
  187. definition,
  188. request,
  189. EventStatus.HANDLER_ERROR,
  190. resolved=resolved,
  191. error="event handler failed: async event handlers require EventKernel.execute",
  192. )
  193. resolved_request = replace(
  194. request,
  195. arguments=resolved.arguments,
  196. raw_arguments=resolved.raw_arguments,
  197. )
  198. try:
  199. payload = definition.handler(resolved_request)
  200. except Exception as exc:
  201. return self._result(
  202. definition,
  203. request,
  204. EventStatus.HANDLER_ERROR,
  205. resolved=resolved,
  206. error=f"event handler failed: {exc}",
  207. )
  208. if inspect.isawaitable(payload):
  209. if inspect.iscoroutine(payload):
  210. payload.close()
  211. return self._result(
  212. definition,
  213. request,
  214. EventStatus.HANDLER_ERROR,
  215. resolved=resolved,
  216. error="event handler failed: async event handlers require EventKernel.execute",
  217. )
  218. if not isinstance(payload, dict):
  219. return self._result(
  220. definition,
  221. request,
  222. EventStatus.HANDLER_ERROR,
  223. resolved=resolved,
  224. error="event handler returned non-object payload",
  225. )
  226. return self._result(
  227. definition,
  228. request,
  229. EventStatus.SUCCESS,
  230. resolved=resolved,
  231. payload=dict(payload),
  232. )
  233. def _lookup(
  234. self,
  235. request: EventRequest,
  236. enabled_names: Iterable[str] | None,
  237. ) -> tuple[EventDefinition | None, EventResult | None]:
  238. definition = self.registry.definition(request.name)
  239. if definition is None:
  240. return None, EventResult(
  241. event_id=request.id,
  242. event_name=request.name,
  243. status=EventStatus.UNKNOWN,
  244. source=request.source,
  245. arguments=dict(request.arguments),
  246. raw_arguments=request.raw_arguments,
  247. error="unknown event",
  248. )
  249. if enabled_names is not None and request.name not in set(enabled_names):
  250. return definition, self._result(
  251. definition,
  252. request,
  253. EventStatus.DISABLED,
  254. resolved=ResolvedEventArguments(
  255. event_name=request.name,
  256. arguments=dict(request.arguments),
  257. raw_arguments=request.raw_arguments,
  258. ),
  259. error="event disabled",
  260. )
  261. return definition, None
  262. def _resolve_deterministic(
  263. self,
  264. definition: EventDefinition,
  265. request: EventRequest,
  266. context: EventExecutionContext,
  267. ) -> tuple[ResolvedEventArguments | None, bool, EventResult | None]:
  268. if request.source is EventSource.PROVIDER_RESOLVED:
  269. return (
  270. ResolvedEventArguments(
  271. event_name=request.name,
  272. arguments=dict(request.arguments),
  273. raw_arguments=request.raw_arguments,
  274. ),
  275. True,
  276. None,
  277. )
  278. try:
  279. value = (
  280. definition.resolver(request, context)
  281. if definition.resolver is not None
  282. else dict(request.arguments)
  283. )
  284. except Exception as exc:
  285. return None, False, self._resolution_error(
  286. definition,
  287. request,
  288. f"event argument resolver failed: {exc}",
  289. )
  290. if isinstance(value, EventArgumentResolution):
  291. arguments = value.arguments
  292. complete = value.complete
  293. else:
  294. arguments = value
  295. complete = True
  296. if not isinstance(arguments, dict) or not isinstance(complete, bool):
  297. return None, False, self._resolution_error(
  298. definition,
  299. request,
  300. "event argument resolver returned invalid payload",
  301. )
  302. try:
  303. raw_arguments = self._json_arguments(arguments)
  304. except (TypeError, ValueError) as exc:
  305. return None, False, self._resolution_error(
  306. definition,
  307. request,
  308. f"event argument resolver failed to serialize: {exc}",
  309. )
  310. return (
  311. ResolvedEventArguments(
  312. event_name=request.name,
  313. arguments=dict(arguments),
  314. raw_arguments=raw_arguments,
  315. ),
  316. complete,
  317. None,
  318. )
  319. def _normalize_fallback(
  320. self,
  321. definition: EventDefinition,
  322. request: EventRequest,
  323. value: Any,
  324. ) -> tuple[ResolvedEventArguments | None, EventResult | None]:
  325. if not isinstance(value, ResolvedEventArguments):
  326. return None, self._resolution_error(
  327. definition,
  328. request,
  329. "event argument fallback returned invalid payload",
  330. used_fallback=True,
  331. )
  332. if not isinstance(value.arguments, dict):
  333. return None, self._resolution_error(
  334. definition,
  335. request,
  336. "event argument fallback returned invalid payload",
  337. used_fallback=True,
  338. )
  339. if value.event_name != request.name:
  340. return None, self._resolution_error(
  341. definition,
  342. request,
  343. f"fallback returned tool {value.event_name} for {request.name}",
  344. resolved=value,
  345. used_fallback=True,
  346. )
  347. try:
  348. raw_arguments = json.loads(
  349. value.raw_arguments,
  350. parse_constant=_reject_json_constant,
  351. )
  352. except (TypeError, ValueError, json.JSONDecodeError):
  353. return None, self._resolution_error(
  354. definition,
  355. request,
  356. "fallback raw arguments are not valid JSON",
  357. resolved=value,
  358. used_fallback=True,
  359. )
  360. try:
  361. self._json_arguments(value.arguments)
  362. except (TypeError, ValueError):
  363. return None, self._resolution_error(
  364. definition,
  365. request,
  366. "fallback parsed arguments are not valid JSON",
  367. resolved=value,
  368. used_fallback=True,
  369. )
  370. if not _json_values_equal(raw_arguments, value.arguments):
  371. return None, self._resolution_error(
  372. definition,
  373. request,
  374. "fallback raw arguments do not match parsed arguments",
  375. resolved=value,
  376. used_fallback=True,
  377. )
  378. return value, None
  379. def _validate(
  380. self,
  381. definition: EventDefinition,
  382. arguments: dict[str, Any],
  383. ) -> _ValidationResult:
  384. validator = self.registry.validator(definition.name)
  385. assert validator is not None
  386. try:
  387. errors = list(validator.iter_errors(arguments))
  388. except Exception as exc:
  389. return _ValidationResult(
  390. status=EventStatus.DEFINITION_ERROR,
  391. error=f"event argument validation failed: {exc}",
  392. )
  393. if not errors:
  394. return _ValidationResult()
  395. missing_errors = [error for error in errors if error.validator == "required"]
  396. error = missing_errors[0] if missing_errors else errors[0]
  397. return _ValidationResult(
  398. status=EventStatus.INVALID_ARGUMENTS,
  399. error=self._format_validation_error(error),
  400. missing_required=bool(missing_errors),
  401. )
  402. def _format_validation_error(self, error: ValidationError) -> str:
  403. if error.validator == "required":
  404. missing = [
  405. name
  406. for name in error.validator_value
  407. if name not in error.instance
  408. ]
  409. return f"missing required arguments: {', '.join(missing)}"
  410. if error.validator == "type" and error.path:
  411. return (
  412. f"invalid argument type for {error.path[-1]}: "
  413. f"expected {error.validator_value}"
  414. )
  415. return f"invalid event arguments: {error.message}"
  416. def _resolution_error(
  417. self,
  418. definition: EventDefinition,
  419. request: EventRequest,
  420. error: str,
  421. *,
  422. resolved: ResolvedEventArguments | None = None,
  423. used_fallback: bool = False,
  424. ) -> EventResult:
  425. return self._result(
  426. definition,
  427. request,
  428. EventStatus.RESOLUTION_ERROR,
  429. resolved=resolved
  430. or ResolvedEventArguments(
  431. event_name=request.name,
  432. arguments=dict(request.arguments),
  433. raw_arguments=request.raw_arguments,
  434. ),
  435. error=error,
  436. used_fallback=used_fallback,
  437. )
  438. def _result(
  439. self,
  440. definition: EventDefinition,
  441. request: EventRequest,
  442. status: EventStatus,
  443. *,
  444. resolved: ResolvedEventArguments,
  445. payload: dict[str, Any] | None = None,
  446. error: str | None = None,
  447. used_fallback: bool = False,
  448. ) -> EventResult:
  449. return EventResult(
  450. event_id=request.id,
  451. event_name=request.name,
  452. status=status,
  453. source=request.source,
  454. arguments=dict(resolved.arguments),
  455. raw_arguments=resolved.raw_arguments,
  456. payload=payload or {},
  457. error=error,
  458. used_fallback=used_fallback,
  459. result_policy=definition.result_policy,
  460. confirmation_policy=definition.confirmation_policy,
  461. risk_level=definition.risk_level,
  462. idempotency_key_fields=definition.idempotency_key_fields,
  463. concurrency_class=definition.concurrency_class,
  464. conflict_keys=definition.conflict_keys,
  465. timeout_seconds=definition.timeout_seconds,
  466. terminal=definition.terminal,
  467. )
  468. def _json_arguments(self, arguments: dict[str, Any]) -> str:
  469. return json.dumps(
  470. arguments,
  471. ensure_ascii=False,
  472. separators=(",", ":"),
  473. allow_nan=False,
  474. )
  475. def _reject_json_constant(value: str) -> None:
  476. raise ValueError(f"non-standard JSON constant: {value}")
  477. def _json_values_equal(left: Any, right: Any) -> bool:
  478. if type(left) is not type(right):
  479. return False
  480. if isinstance(left, dict):
  481. return left.keys() == right.keys() and all(
  482. _json_values_equal(left[key], right[key]) for key in left
  483. )
  484. if isinstance(left, list):
  485. return len(left) == len(right) and all(
  486. _json_values_equal(left_item, right_item)
  487. for left_item, right_item in zip(left, right, strict=True)
  488. )
  489. return left == right