ForcePilot/backend/package/yuxi/channel/hooks/lifecycle.py
Kris a81279cf61 feat(channel/hooks): 实现完整的钩子系统模块
新增了完整的钩子系统,包括元数据定义、扫描器、解析器、生命周期注册表和加载器,支持从目录扫描HOOK.md配置、验证依赖和系统兼容性,以及通过沙箱加载钩子脚本,实现了事件驱动的钩子回调机制。
2026-05-21 10:26:47 +08:00

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