ForcePilot/backend/package/yuxi/channel/hooks/lifecycle.py

225 lines
8.1 KiB
Python
Raw Normal View History

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()}