| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511 |
- from __future__ import annotations
- import inspect
- import json
- from collections.abc import Awaitable, Callable, Iterable
- from dataclasses import dataclass, replace
- from typing import Any
- from jsonschema.exceptions import ValidationError
- from agent_lab.application.events.models import (
- EventDefinition,
- EventArgumentResolution,
- EventExecutionContext,
- EventRequest,
- EventResult,
- EventSource,
- EventStatus,
- ResolvedEventArguments,
- )
- from agent_lab.application.events.registry import EventRegistry
- ArgumentFallbackResult = ResolvedEventArguments | None
- ArgumentFallback = Callable[
- [EventDefinition, EventRequest, EventExecutionContext],
- ArgumentFallbackResult | Awaitable[ArgumentFallbackResult],
- ]
- @dataclass(frozen=True)
- class _ValidationResult:
- status: EventStatus | None = None
- error: str | None = None
- missing_required: bool = False
- class EventKernel:
- def __init__(
- self,
- registry: EventRegistry,
- argument_fallback: ArgumentFallback | None = None,
- ) -> None:
- self.registry = registry
- self.argument_fallback = argument_fallback
- async def execute(
- self,
- request: EventRequest,
- *,
- enabled_names: Iterable[str] | None = None,
- context: EventExecutionContext | None = None,
- ) -> EventResult:
- definition, early_result = self._lookup(request, enabled_names)
- if early_result is not None:
- return early_result
- assert definition is not None
- resolved_context = context or EventExecutionContext()
- resolved, resolution_complete, resolution_error = (
- self._resolve_deterministic(
- definition,
- request,
- resolved_context,
- )
- )
- if resolution_error is not None:
- return resolution_error
- assert resolved is not None
- validation = self._validate(definition, resolved.arguments)
- if not resolution_complete and validation.status is None:
- validation = _ValidationResult(
- status=EventStatus.INVALID_ARGUMENTS,
- error="event arguments incomplete",
- )
- used_fallback = False
- if (
- validation.status is not EventStatus.DEFINITION_ERROR
- and (validation.missing_required or not resolution_complete)
- and request.source is not EventSource.PROVIDER_RESOLVED
- and definition.fallback_allowed
- and self.argument_fallback is not None
- ):
- used_fallback = True
- fallback_request = replace(
- request,
- arguments=resolved.arguments,
- raw_arguments=resolved.raw_arguments,
- )
- try:
- fallback_value = self.argument_fallback(
- definition,
- fallback_request,
- resolved_context,
- )
- if inspect.isawaitable(fallback_value):
- fallback_value = await fallback_value
- except Exception as exc:
- return self._resolution_error(
- definition,
- request,
- f"event argument fallback failed: {exc}",
- resolved=resolved,
- used_fallback=True,
- )
- if fallback_value is not None:
- resolved, fallback_error = self._normalize_fallback(
- definition,
- request,
- fallback_value,
- )
- if fallback_error is not None:
- return fallback_error
- assert resolved is not None
- validation = self._validate(
- definition,
- resolved.arguments,
- )
- if validation.status is not None:
- return self._result(
- definition,
- request,
- validation.status,
- resolved=resolved,
- error=validation.error,
- used_fallback=used_fallback,
- )
- resolved_request = replace(
- request,
- arguments=resolved.arguments,
- raw_arguments=resolved.raw_arguments,
- )
- try:
- payload = definition.handler(resolved_request)
- if inspect.isawaitable(payload):
- payload = await payload
- except Exception as exc:
- return self._result(
- definition,
- request,
- EventStatus.HANDLER_ERROR,
- resolved=resolved,
- error=f"event handler failed: {exc}",
- used_fallback=used_fallback,
- )
- if not isinstance(payload, dict):
- return self._result(
- definition,
- request,
- EventStatus.HANDLER_ERROR,
- resolved=resolved,
- error="event handler returned non-object payload",
- used_fallback=used_fallback,
- )
- return self._result(
- definition,
- request,
- EventStatus.SUCCESS,
- resolved=resolved,
- payload=dict(payload),
- used_fallback=used_fallback,
- )
- def execute_sync(
- self,
- request: EventRequest,
- *,
- enabled_names: Iterable[str] | None = None,
- context: EventExecutionContext | None = None,
- ) -> EventResult:
- definition, early_result = self._lookup(request, enabled_names)
- if early_result is not None:
- return early_result
- assert definition is not None
- resolved, _, resolution_error = self._resolve_deterministic(
- definition,
- request,
- context or EventExecutionContext(),
- )
- if resolution_error is not None:
- return resolution_error
- assert resolved is not None
- validation = self._validate(definition, resolved.arguments)
- if validation.status is not None:
- return self._result(
- definition,
- request,
- validation.status,
- resolved=resolved,
- error=validation.error,
- )
- if inspect.iscoroutinefunction(definition.handler):
- return self._result(
- definition,
- request,
- EventStatus.HANDLER_ERROR,
- resolved=resolved,
- error="event handler failed: async event handlers require EventKernel.execute",
- )
- resolved_request = replace(
- request,
- arguments=resolved.arguments,
- raw_arguments=resolved.raw_arguments,
- )
- try:
- payload = definition.handler(resolved_request)
- except Exception as exc:
- return self._result(
- definition,
- request,
- EventStatus.HANDLER_ERROR,
- resolved=resolved,
- error=f"event handler failed: {exc}",
- )
- if inspect.isawaitable(payload):
- if inspect.iscoroutine(payload):
- payload.close()
- return self._result(
- definition,
- request,
- EventStatus.HANDLER_ERROR,
- resolved=resolved,
- error="event handler failed: async event handlers require EventKernel.execute",
- )
- if not isinstance(payload, dict):
- return self._result(
- definition,
- request,
- EventStatus.HANDLER_ERROR,
- resolved=resolved,
- error="event handler returned non-object payload",
- )
- return self._result(
- definition,
- request,
- EventStatus.SUCCESS,
- resolved=resolved,
- payload=dict(payload),
- )
- def _lookup(
- self,
- request: EventRequest,
- enabled_names: Iterable[str] | None,
- ) -> tuple[EventDefinition | None, EventResult | None]:
- definition = self.registry.definition(request.name)
- if definition is None:
- return None, EventResult(
- event_id=request.id,
- event_name=request.name,
- status=EventStatus.UNKNOWN,
- source=request.source,
- arguments=dict(request.arguments),
- raw_arguments=request.raw_arguments,
- error="unknown event",
- )
- if enabled_names is not None and request.name not in set(enabled_names):
- return definition, self._result(
- definition,
- request,
- EventStatus.DISABLED,
- resolved=ResolvedEventArguments(
- event_name=request.name,
- arguments=dict(request.arguments),
- raw_arguments=request.raw_arguments,
- ),
- error="event disabled",
- )
- return definition, None
- def _resolve_deterministic(
- self,
- definition: EventDefinition,
- request: EventRequest,
- context: EventExecutionContext,
- ) -> tuple[ResolvedEventArguments | None, bool, EventResult | None]:
- if request.source is EventSource.PROVIDER_RESOLVED:
- return (
- ResolvedEventArguments(
- event_name=request.name,
- arguments=dict(request.arguments),
- raw_arguments=request.raw_arguments,
- ),
- True,
- None,
- )
- try:
- value = (
- definition.resolver(request, context)
- if definition.resolver is not None
- else dict(request.arguments)
- )
- except Exception as exc:
- return None, False, self._resolution_error(
- definition,
- request,
- f"event argument resolver failed: {exc}",
- )
- if isinstance(value, EventArgumentResolution):
- arguments = value.arguments
- complete = value.complete
- else:
- arguments = value
- complete = True
- if not isinstance(arguments, dict) or not isinstance(complete, bool):
- return None, False, self._resolution_error(
- definition,
- request,
- "event argument resolver returned invalid payload",
- )
- try:
- raw_arguments = self._json_arguments(arguments)
- except (TypeError, ValueError) as exc:
- return None, False, self._resolution_error(
- definition,
- request,
- f"event argument resolver failed to serialize: {exc}",
- )
- return (
- ResolvedEventArguments(
- event_name=request.name,
- arguments=dict(arguments),
- raw_arguments=raw_arguments,
- ),
- complete,
- None,
- )
- def _normalize_fallback(
- self,
- definition: EventDefinition,
- request: EventRequest,
- value: Any,
- ) -> tuple[ResolvedEventArguments | None, EventResult | None]:
- if not isinstance(value, ResolvedEventArguments):
- return None, self._resolution_error(
- definition,
- request,
- "event argument fallback returned invalid payload",
- used_fallback=True,
- )
- if not isinstance(value.arguments, dict):
- return None, self._resolution_error(
- definition,
- request,
- "event argument fallback returned invalid payload",
- used_fallback=True,
- )
- if value.event_name != request.name:
- return None, self._resolution_error(
- definition,
- request,
- f"fallback returned tool {value.event_name} for {request.name}",
- resolved=value,
- used_fallback=True,
- )
- try:
- raw_arguments = json.loads(
- value.raw_arguments,
- parse_constant=_reject_json_constant,
- )
- except (TypeError, ValueError, json.JSONDecodeError):
- return None, self._resolution_error(
- definition,
- request,
- "fallback raw arguments are not valid JSON",
- resolved=value,
- used_fallback=True,
- )
- try:
- self._json_arguments(value.arguments)
- except (TypeError, ValueError):
- return None, self._resolution_error(
- definition,
- request,
- "fallback parsed arguments are not valid JSON",
- resolved=value,
- used_fallback=True,
- )
- if not _json_values_equal(raw_arguments, value.arguments):
- return None, self._resolution_error(
- definition,
- request,
- "fallback raw arguments do not match parsed arguments",
- resolved=value,
- used_fallback=True,
- )
- return value, None
- def _validate(
- self,
- definition: EventDefinition,
- arguments: dict[str, Any],
- ) -> _ValidationResult:
- validator = self.registry.validator(definition.name)
- assert validator is not None
- try:
- errors = list(validator.iter_errors(arguments))
- except Exception as exc:
- return _ValidationResult(
- status=EventStatus.DEFINITION_ERROR,
- error=f"event argument validation failed: {exc}",
- )
- if not errors:
- return _ValidationResult()
- missing_errors = [error for error in errors if error.validator == "required"]
- error = missing_errors[0] if missing_errors else errors[0]
- return _ValidationResult(
- status=EventStatus.INVALID_ARGUMENTS,
- error=self._format_validation_error(error),
- missing_required=bool(missing_errors),
- )
- def _format_validation_error(self, error: ValidationError) -> str:
- if error.validator == "required":
- missing = [
- name
- for name in error.validator_value
- if name not in error.instance
- ]
- return f"missing required arguments: {', '.join(missing)}"
- if error.validator == "type" and error.path:
- return (
- f"invalid argument type for {error.path[-1]}: "
- f"expected {error.validator_value}"
- )
- return f"invalid event arguments: {error.message}"
- def _resolution_error(
- self,
- definition: EventDefinition,
- request: EventRequest,
- error: str,
- *,
- resolved: ResolvedEventArguments | None = None,
- used_fallback: bool = False,
- ) -> EventResult:
- return self._result(
- definition,
- request,
- EventStatus.RESOLUTION_ERROR,
- resolved=resolved
- or ResolvedEventArguments(
- event_name=request.name,
- arguments=dict(request.arguments),
- raw_arguments=request.raw_arguments,
- ),
- error=error,
- used_fallback=used_fallback,
- )
- def _result(
- self,
- definition: EventDefinition,
- request: EventRequest,
- status: EventStatus,
- *,
- resolved: ResolvedEventArguments,
- payload: dict[str, Any] | None = None,
- error: str | None = None,
- used_fallback: bool = False,
- ) -> EventResult:
- return EventResult(
- event_id=request.id,
- event_name=request.name,
- status=status,
- source=request.source,
- arguments=dict(resolved.arguments),
- raw_arguments=resolved.raw_arguments,
- payload=payload or {},
- error=error,
- used_fallback=used_fallback,
- 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 _json_arguments(self, arguments: dict[str, Any]) -> str:
- return json.dumps(
- arguments,
- ensure_ascii=False,
- separators=(",", ":"),
- allow_nan=False,
- )
- def _reject_json_constant(value: str) -> None:
- raise ValueError(f"non-standard JSON constant: {value}")
- def _json_values_equal(left: Any, right: Any) -> bool:
- if type(left) is not type(right):
- return False
- if isinstance(left, dict):
- return left.keys() == right.keys() and all(
- _json_values_equal(left[key], right[key]) for key in left
- )
- if isinstance(left, list):
- return len(left) == len(right) and all(
- _json_values_equal(left_item, right_item)
- for left_item, right_item in zip(left, right, strict=True)
- )
- return left == right
|