kernel.py 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511
  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_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 validation.status is not None:
  173. return self._result(
  174. definition,
  175. request,
  176. validation.status,
  177. resolved=resolved,
  178. error=validation.error,
  179. )
  180. if inspect.iscoroutinefunction(definition.handler):
  181. return self._result(
  182. definition,
  183. request,
  184. EventStatus.HANDLER_ERROR,
  185. resolved=resolved,
  186. error="event handler failed: async event handlers require EventKernel.execute",
  187. )
  188. resolved_request = replace(
  189. request,
  190. arguments=resolved.arguments,
  191. raw_arguments=resolved.raw_arguments,
  192. )
  193. try:
  194. payload = definition.handler(resolved_request)
  195. except Exception as exc:
  196. return self._result(
  197. definition,
  198. request,
  199. EventStatus.HANDLER_ERROR,
  200. resolved=resolved,
  201. error=f"event handler failed: {exc}",
  202. )
  203. if inspect.isawaitable(payload):
  204. if inspect.iscoroutine(payload):
  205. payload.close()
  206. return self._result(
  207. definition,
  208. request,
  209. EventStatus.HANDLER_ERROR,
  210. resolved=resolved,
  211. error="event handler failed: async event handlers require EventKernel.execute",
  212. )
  213. if not isinstance(payload, dict):
  214. return self._result(
  215. definition,
  216. request,
  217. EventStatus.HANDLER_ERROR,
  218. resolved=resolved,
  219. error="event handler returned non-object payload",
  220. )
  221. return self._result(
  222. definition,
  223. request,
  224. EventStatus.SUCCESS,
  225. resolved=resolved,
  226. payload=dict(payload),
  227. )
  228. def _lookup(
  229. self,
  230. request: EventRequest,
  231. enabled_names: Iterable[str] | None,
  232. ) -> tuple[EventDefinition | None, EventResult | None]:
  233. definition = self.registry.definition(request.name)
  234. if definition is None:
  235. return None, EventResult(
  236. event_id=request.id,
  237. event_name=request.name,
  238. status=EventStatus.UNKNOWN,
  239. source=request.source,
  240. arguments=dict(request.arguments),
  241. raw_arguments=request.raw_arguments,
  242. error="unknown event",
  243. )
  244. if enabled_names is not None and request.name not in set(enabled_names):
  245. return definition, self._result(
  246. definition,
  247. request,
  248. EventStatus.DISABLED,
  249. resolved=ResolvedEventArguments(
  250. event_name=request.name,
  251. arguments=dict(request.arguments),
  252. raw_arguments=request.raw_arguments,
  253. ),
  254. error="event disabled",
  255. )
  256. return definition, None
  257. def _resolve_deterministic(
  258. self,
  259. definition: EventDefinition,
  260. request: EventRequest,
  261. context: EventExecutionContext,
  262. ) -> tuple[ResolvedEventArguments | None, bool, EventResult | None]:
  263. if request.source is EventSource.PROVIDER_RESOLVED:
  264. return (
  265. ResolvedEventArguments(
  266. event_name=request.name,
  267. arguments=dict(request.arguments),
  268. raw_arguments=request.raw_arguments,
  269. ),
  270. True,
  271. None,
  272. )
  273. try:
  274. value = (
  275. definition.resolver(request, context)
  276. if definition.resolver is not None
  277. else dict(request.arguments)
  278. )
  279. except Exception as exc:
  280. return None, False, self._resolution_error(
  281. definition,
  282. request,
  283. f"event argument resolver failed: {exc}",
  284. )
  285. if isinstance(value, EventArgumentResolution):
  286. arguments = value.arguments
  287. complete = value.complete
  288. else:
  289. arguments = value
  290. complete = True
  291. if not isinstance(arguments, dict) or not isinstance(complete, bool):
  292. return None, False, self._resolution_error(
  293. definition,
  294. request,
  295. "event argument resolver returned invalid payload",
  296. )
  297. try:
  298. raw_arguments = self._json_arguments(arguments)
  299. except (TypeError, ValueError) as exc:
  300. return None, False, self._resolution_error(
  301. definition,
  302. request,
  303. f"event argument resolver failed to serialize: {exc}",
  304. )
  305. return (
  306. ResolvedEventArguments(
  307. event_name=request.name,
  308. arguments=dict(arguments),
  309. raw_arguments=raw_arguments,
  310. ),
  311. complete,
  312. None,
  313. )
  314. def _normalize_fallback(
  315. self,
  316. definition: EventDefinition,
  317. request: EventRequest,
  318. value: Any,
  319. ) -> tuple[ResolvedEventArguments | None, EventResult | None]:
  320. if not isinstance(value, ResolvedEventArguments):
  321. return None, self._resolution_error(
  322. definition,
  323. request,
  324. "event argument fallback returned invalid payload",
  325. used_fallback=True,
  326. )
  327. if not isinstance(value.arguments, dict):
  328. return None, self._resolution_error(
  329. definition,
  330. request,
  331. "event argument fallback returned invalid payload",
  332. used_fallback=True,
  333. )
  334. if value.event_name != request.name:
  335. return None, self._resolution_error(
  336. definition,
  337. request,
  338. f"fallback returned tool {value.event_name} for {request.name}",
  339. resolved=value,
  340. used_fallback=True,
  341. )
  342. try:
  343. raw_arguments = json.loads(
  344. value.raw_arguments,
  345. parse_constant=_reject_json_constant,
  346. )
  347. except (TypeError, ValueError, json.JSONDecodeError):
  348. return None, self._resolution_error(
  349. definition,
  350. request,
  351. "fallback raw arguments are not valid JSON",
  352. resolved=value,
  353. used_fallback=True,
  354. )
  355. try:
  356. self._json_arguments(value.arguments)
  357. except (TypeError, ValueError):
  358. return None, self._resolution_error(
  359. definition,
  360. request,
  361. "fallback parsed arguments are not valid JSON",
  362. resolved=value,
  363. used_fallback=True,
  364. )
  365. if not _json_values_equal(raw_arguments, value.arguments):
  366. return None, self._resolution_error(
  367. definition,
  368. request,
  369. "fallback raw arguments do not match parsed arguments",
  370. resolved=value,
  371. used_fallback=True,
  372. )
  373. return value, None
  374. def _validate(
  375. self,
  376. definition: EventDefinition,
  377. arguments: dict[str, Any],
  378. ) -> _ValidationResult:
  379. validator = self.registry.validator(definition.name)
  380. assert validator is not None
  381. try:
  382. errors = list(validator.iter_errors(arguments))
  383. except Exception as exc:
  384. return _ValidationResult(
  385. status=EventStatus.DEFINITION_ERROR,
  386. error=f"event argument validation failed: {exc}",
  387. )
  388. if not errors:
  389. return _ValidationResult()
  390. missing_errors = [error for error in errors if error.validator == "required"]
  391. error = missing_errors[0] if missing_errors else errors[0]
  392. return _ValidationResult(
  393. status=EventStatus.INVALID_ARGUMENTS,
  394. error=self._format_validation_error(error),
  395. missing_required=bool(missing_errors),
  396. )
  397. def _format_validation_error(self, error: ValidationError) -> str:
  398. if error.validator == "required":
  399. missing = [
  400. name
  401. for name in error.validator_value
  402. if name not in error.instance
  403. ]
  404. return f"missing required arguments: {', '.join(missing)}"
  405. if error.validator == "type" and error.path:
  406. return (
  407. f"invalid argument type for {error.path[-1]}: "
  408. f"expected {error.validator_value}"
  409. )
  410. return f"invalid event arguments: {error.message}"
  411. def _resolution_error(
  412. self,
  413. definition: EventDefinition,
  414. request: EventRequest,
  415. error: str,
  416. *,
  417. resolved: ResolvedEventArguments | None = None,
  418. used_fallback: bool = False,
  419. ) -> EventResult:
  420. return self._result(
  421. definition,
  422. request,
  423. EventStatus.RESOLUTION_ERROR,
  424. resolved=resolved
  425. or ResolvedEventArguments(
  426. event_name=request.name,
  427. arguments=dict(request.arguments),
  428. raw_arguments=request.raw_arguments,
  429. ),
  430. error=error,
  431. used_fallback=used_fallback,
  432. )
  433. def _result(
  434. self,
  435. definition: EventDefinition,
  436. request: EventRequest,
  437. status: EventStatus,
  438. *,
  439. resolved: ResolvedEventArguments,
  440. payload: dict[str, Any] | None = None,
  441. error: str | None = None,
  442. used_fallback: bool = False,
  443. ) -> EventResult:
  444. return EventResult(
  445. event_id=request.id,
  446. event_name=request.name,
  447. status=status,
  448. source=request.source,
  449. arguments=dict(resolved.arguments),
  450. raw_arguments=resolved.raw_arguments,
  451. payload=payload or {},
  452. error=error,
  453. used_fallback=used_fallback,
  454. result_policy=definition.result_policy,
  455. confirmation_policy=definition.confirmation_policy,
  456. risk_level=definition.risk_level,
  457. idempotency_key_fields=definition.idempotency_key_fields,
  458. concurrency_class=definition.concurrency_class,
  459. conflict_keys=definition.conflict_keys,
  460. timeout_seconds=definition.timeout_seconds,
  461. terminal=definition.terminal,
  462. )
  463. def _json_arguments(self, arguments: dict[str, Any]) -> str:
  464. return json.dumps(
  465. arguments,
  466. ensure_ascii=False,
  467. separators=(",", ":"),
  468. allow_nan=False,
  469. )
  470. def _reject_json_constant(value: str) -> None:
  471. raise ValueError(f"non-standard JSON constant: {value}")
  472. def _json_values_equal(left: Any, right: Any) -> bool:
  473. if type(left) is not type(right):
  474. return False
  475. if isinstance(left, dict):
  476. return left.keys() == right.keys() and all(
  477. _json_values_equal(left[key], right[key]) for key in left
  478. )
  479. if isinstance(left, list):
  480. return len(left) == len(right) and all(
  481. _json_values_equal(left_item, right_item)
  482. for left_item, right_item in zip(left, right, strict=True)
  483. )
  484. return left == right