新增了完整的钩子系统,包括元数据定义、扫描器、解析器、生命周期注册表和加载器,支持从目录扫描HOOK.md配置、验证依赖和系统兼容性,以及通过沙箱加载钩子脚本,实现了事件驱动的钩子回调机制。
225 lines
8.1 KiB
Python
225 lines
8.1 KiB
Python
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()}
|