import asyncio import bisect import logging from collections.abc import Callable, Coroutine from dataclasses import dataclass from enum import StrEnum from typing import Any, TYPE_CHECKING if TYPE_CHECKING: from yuxi.channel.events.bus import ChannelEventBus _DEFAULT_HOOK_TIMEOUT = 30.0 _DEFAULT_FIRE_CONCURRENCY = 10 logger = logging.getLogger(__name__) HookCallback = Callable[..., Coroutine[Any, Any, None]] class HookEvent(StrEnum): CHANNEL_STARTING = "on_channel_starting" CHANNEL_STOPPING = "on_channel_stopping" CHANNEL_RUNNING = "on_channel_running" CHANNEL_ERROR = "on_channel_error" CHANNEL_FAILED = "on_channel_failed" GATEWAY_START = "on_gateway_start" GATEWAY_STOP = "on_gateway_stop" MESSAGE_RECEIVED = "on_message_received" MESSAGE_SENDING = "on_message_sending" MESSAGE_SENT = "on_message_sent" BEFORE_AGENT_RUN = "on_before_agent_run" AFTER_AGENT_RUN = "on_after_agent_run" AGENT_BOOTSTRAP = "on_agent_bootstrap" SESSION_START = "on_session_start" SESSION_END = "on_session_end" @dataclass class HookRegistration: callback: HookCallback priority: int = 100 enabled: bool = True key: str = "" def __hash__(self) -> int: return id(self.callback) class LifecycleHookRegistry: def __init__( self, event_bus: "ChannelEventBus | None" = None, *, fire_concurrency: int = _DEFAULT_FIRE_CONCURRENCY, ) -> None: self._hooks: dict[str, list[HookRegistration]] = {} self._locks: dict[str, asyncio.Lock] = {} self._event_bus = event_bus self._fire_sem = asyncio.Semaphore(fire_concurrency) def _get_lock(self, event: HookEvent | str) -> asyncio.Lock: return self._locks.setdefault(event, asyncio.Lock()) async def register( self, event: HookEvent | str, callback: HookCallback, *, priority: int = 100, key: str = "", ) -> HookRegistration: reg = HookRegistration(callback=callback, priority=priority, key=key) async with self._get_lock(event): hooks = self._hooks.setdefault(event, []) idx = bisect.bisect_left([r.priority for r in hooks], priority) hooks.insert(idx, reg) return reg async def unregister(self, event: HookEvent | str, callback: HookCallback) -> None: async with self._get_lock(event): registrations = self._hooks.get(event) if registrations is None: return self._hooks[event] = [r for r in registrations if r.callback != callback] if not self._hooks[event]: del self._hooks[event] async def enable(self, event: HookEvent | str, callback: HookCallback) -> bool: async with self._get_lock(event): for reg in self._hooks.get(event, []): if reg.callback == callback: reg.enabled = True return True return False async def disable(self, event: HookEvent | str, callback: HookCallback) -> bool: async with self._get_lock(event): for reg in self._hooks.get(event, []): if reg.callback == callback: reg.enabled = False return True return False async def set_priority(self, event: HookEvent | str, callback: HookCallback, priority: int) -> bool: async with self._get_lock(event): for reg in self._hooks.get(event, []): if reg.callback == callback: reg.priority = priority self._hooks[event].sort(key=lambda r: r.priority) return True return False async def fire(self, event: HookEvent | str, *args: Any, timeout: float | None = None, **kwargs: Any) -> None: async with self._get_lock(event): registrations = self._hooks.get(event) if not registrations: return enabled = [r for r in registrations if r.enabled] if not enabled: return effective_timeout = timeout if timeout is not None else _DEFAULT_HOOK_TIMEOUT async def _safe_invoke(reg: HookRegistration) -> None: async with self._fire_sem: try: await asyncio.wait_for(reg.callback(*args, **kwargs), timeout=effective_timeout) except asyncio.TimeoutError: logger.error( "Hook callback timed out after %.1fs: event=%s key=%s", effective_timeout, event, reg.key, ) except Exception: logger.exception("Hook callback failed: event=%s key=%s", event, reg.key) await asyncio.gather(*(_safe_invoke(r) for r in enabled)) if self._event_bus is not None: await self._forward_to_bus(event, *args, **kwargs) async def fire_sequential(self, event: HookEvent | str, initial: Any = None, *args: Any, timeout: float | None = None, **kwargs: Any) -> Any: async with self._get_lock(event): registrations = list(self._hooks.get(event, [])) if not registrations: return initial effective_timeout = timeout if timeout is not None else _DEFAULT_HOOK_TIMEOUT result = initial for reg in registrations: if not reg.enabled: continue try: next_result = await asyncio.wait_for(reg.callback(result, *args, **kwargs), timeout=effective_timeout) if next_result is not None: result = next_result except asyncio.TimeoutError: logger.error( "Hook callback timed out after %.1fs: event=%s key=%s", effective_timeout, event, reg.key, ) except Exception: logger.exception("Hook callback failed: event=%s key=%s", event, reg.key) return result async def fire_first(self, event: HookEvent | str, *args: Any, timeout: float | None = None, **kwargs: Any) -> Any: async with self._get_lock(event): registrations = list(self._hooks.get(event, [])) if not registrations: return None effective_timeout = timeout if timeout is not None else _DEFAULT_HOOK_TIMEOUT for reg in registrations: if not reg.enabled: continue try: result = await asyncio.wait_for(reg.callback(*args, **kwargs), timeout=effective_timeout) if result is not None: return result except asyncio.TimeoutError: logger.error( "Hook callback timed out after %.1fs: event=%s key=%s", effective_timeout, event, reg.key, ) except Exception: logger.exception("Hook callback failed: event=%s key=%s", event, reg.key) return None async def clear(self) -> None: async with asyncio.Lock(): for lock in self._locks.values(): async with lock: pass self._hooks.clear() self._locks.clear() async def _forward_to_bus(self, event: str, *args: Any, **kwargs: Any) -> None: if self._event_bus is None: return from yuxi.channel.events.types import HOOK_EVENT_TO_TOPIC topic = HOOK_EVENT_TO_TOPIC.get(event) if topic is None: logger.debug("No event bus topic mapping for hook event: %s", event) return try: await self._event_bus.publish(topic, *args, **kwargs) except Exception: logger.exception("Event bus forward failed: hook=%s topic=%s", event, topic) def list_hooks(self, event: HookEvent | str | None = None) -> dict[str, list[HookRegistration]]: if event is not None: return {event: list(self._hooks.get(event, []))} return {k: list(v) for k, v in self._hooks.items()}