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