From a81279cf61356717b02beb3025a489f090b2dcd2 Mon Sep 17 00:00:00 2001 From: Kris <2893855659@qq.com> Date: Thu, 21 May 2026 10:26:47 +0800 Subject: [PATCH] =?UTF-8?q?feat(channel/hooks):=20=E5=AE=9E=E7=8E=B0?= =?UTF-8?q?=E5=AE=8C=E6=95=B4=E7=9A=84=E9=92=A9=E5=AD=90=E7=B3=BB=E7=BB=9F?= =?UTF-8?q?=E6=A8=A1=E5=9D=97?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 新增了完整的钩子系统,包括元数据定义、扫描器、解析器、生命周期注册表和加载器,支持从目录扫描HOOK.md配置、验证依赖和系统兼容性,以及通过沙箱加载钩子脚本,实现了事件驱动的钩子回调机制。 --- .../package/yuxi/channel/hooks/__init__.py | 43 +++ .../package/yuxi/channel/hooks/hook_loader.py | 281 ++++++++++++++++++ .../yuxi/channel/hooks/hook_metadata.py | 3 + .../package/yuxi/channel/hooks/hook_parser.py | 191 ++++++++++++ .../yuxi/channel/hooks/hook_scanner.py | 106 +++++++ .../package/yuxi/channel/hooks/lifecycle.py | 224 ++++++++++++++ backend/package/yuxi/channel/hooks/types.py | 53 ++++ 7 files changed, 901 insertions(+) create mode 100644 backend/package/yuxi/channel/hooks/__init__.py create mode 100644 backend/package/yuxi/channel/hooks/hook_loader.py create mode 100644 backend/package/yuxi/channel/hooks/hook_metadata.py create mode 100644 backend/package/yuxi/channel/hooks/hook_parser.py create mode 100644 backend/package/yuxi/channel/hooks/hook_scanner.py create mode 100644 backend/package/yuxi/channel/hooks/lifecycle.py create mode 100644 backend/package/yuxi/channel/hooks/types.py diff --git a/backend/package/yuxi/channel/hooks/__init__.py b/backend/package/yuxi/channel/hooks/__init__.py new file mode 100644 index 00000000..0a5224e4 --- /dev/null +++ b/backend/package/yuxi/channel/hooks/__init__.py @@ -0,0 +1,43 @@ +from yuxi.channel.hooks.hook_loader import ( + HookLoadEntry, + HookLoadResult, + load_all_hook_sources, + load_hooks_from_dir, + load_hooks_from_dirs, +) +from yuxi.channel.hooks.hook_parser import ( + HookParseError, + HookValidationError, + clear_bin_cache, + parse_hook_md, + validate_dependencies, + validate_os_compatibility, +) +from yuxi.channel.hooks.hook_scanner import DiscoveredHook, ScanResult, scan_directory +from yuxi.channel.hooks.lifecycle import HookCallback, HookEvent, HookRegistration, LifecycleHookRegistry +from yuxi.channel.hooks.types import HookConfigField, HookDependencies, HookMetadata, SkipReason + +__all__ = [ + "HookCallback", + "HookEvent", + "HookRegistration", + "LifecycleHookRegistry", + "HookMetadata", + "HookConfigField", + "HookDependencies", + "HookParseError", + "HookValidationError", + "clear_bin_cache", + "parse_hook_md", + "validate_os_compatibility", + "validate_dependencies", + "scan_directory", + "ScanResult", + "DiscoveredHook", + "load_hooks_from_dir", + "load_hooks_from_dirs", + "load_all_hook_sources", + "HookLoadResult", + "HookLoadEntry", + "SkipReason", +] diff --git a/backend/package/yuxi/channel/hooks/hook_loader.py b/backend/package/yuxi/channel/hooks/hook_loader.py new file mode 100644 index 00000000..fe91a135 --- /dev/null +++ b/backend/package/yuxi/channel/hooks/hook_loader.py @@ -0,0 +1,281 @@ +import logging +import sys as _sys +from dataclasses import dataclass, field + +from yuxi.channel.hooks.hook_parser import _current_os, validate_dependencies, validate_os_compatibility +from yuxi.channel.hooks.hook_scanner import DiscoveredHook, scan_directory +from yuxi.channel.hooks.lifecycle import HookEvent, LifecycleHookRegistry +from yuxi.channel.hooks.types import SkipReason +from yuxi.channel.plugins.sandbox import PluginSandbox, SandboxViolationError + +logger = logging.getLogger(__name__) + + +@dataclass +class HookLoadEntry: + hook_id: str + status: str + events_loaded: int = 0 + reason: str = "" + detail: dict = field(default_factory=dict) + + +@dataclass +class HookLoadResult: + loaded: list[HookLoadEntry] = field(default_factory=list) + skipped: list[HookLoadEntry] = field(default_factory=list) + failed: list[HookLoadEntry] = field(default_factory=list) + errors: list[str] = field(default_factory=list) + + @property + def total(self) -> int: + return len(self.loaded) + len(self.skipped) + len(self.failed) + + @property + def loaded_count(self) -> int: + return len(self.loaded) + + @property + def skipped_count(self) -> int: + return len(self.skipped) + + @property + def failed_count(self) -> int: + return len(self.failed) + + def to_dict(self) -> dict: + return { + "total": self.total, + "loaded": self.loaded_count, + "skipped": self.skipped_count, + "failed": self.failed_count, + "entries": { + "loaded": [{"hook_id": e.hook_id, "events": e.events_loaded} for e in self.loaded], + "skipped": [{"hook_id": e.hook_id, "reason": e.reason, "detail": e.detail} for e in self.skipped], + "failed": [{"hook_id": e.hook_id, "reason": e.reason} for e in self.failed], + }, + "errors": self.errors, + } + + +async def load_hooks_from_dir( + directory: str, + registry: LifecycleHookRegistry, + *, + callbacks: dict[str, dict] | None = None, +) -> HookLoadResult: + result = HookLoadResult() + + scan_result = scan_directory(directory) + result.errors.extend(scan_result.errors) + + if not scan_result.hooks: + logger.debug("No hooks discovered in %s", directory) + return result + + callback_map = callbacks or {} + + for discovered in scan_result.hooks: + await _load_single_hook(discovered, registry, callback_map, result) + + return result + + +async def _load_single_hook( + discovered: DiscoveredHook, + registry: LifecycleHookRegistry, + callback_map: dict[str, dict], + result: HookLoadResult, +) -> None: + metadata = discovered.metadata + hook_id = metadata.hook_id + + current_os = _current_os() + + if not validate_os_compatibility(metadata): + entry = HookLoadEntry( + hook_id=hook_id, + status=SkipReason.OS_INCOMPATIBLE, + reason=f"requires {metadata.os}, current: {current_os}", + detail={"required_os": metadata.os, "current_os": current_os}, + ) + result.skipped.append(entry) + logger.warning( + "Hook %s skipped: OS incompatible (requires %s, current: %s)", + hook_id, + metadata.os, + current_os, + ) + return + + deps_ok, dep_failures = validate_dependencies(metadata) + if not deps_ok: + entry = HookLoadEntry( + hook_id=hook_id, + status=SkipReason.DEPENDENCY_MISSING, + reason="; ".join(dep_failures), + detail={"failures": dep_failures}, + ) + result.skipped.append(entry) + logger.warning("Hook %s skipped: dependencies not met: %s", hook_id, dep_failures) + return + + callbacks_for_hook = callback_map.get(hook_id, {}) + + registered = 0 + for event_name in metadata.events: + callback = callbacks_for_hook.get(event_name) + if callback is None or not callable(callback): + entry = HookLoadEntry( + hook_id=hook_id, + status=SkipReason.NO_CALLBACK, + reason=f"no callback for event {event_name}", + ) + result.failed.append(entry) + logger.warning("Hook %s: no callback registered for event %s", hook_id, event_name) + continue + + await registry.register(event_name, callback, key=hook_id) + registered += 1 + + script_loaded = await _try_load_hook_script(discovered, registry, hook_id) + if registered > 0 or script_loaded > 0: + entry = HookLoadEntry(hook_id=hook_id, status="loaded", events_loaded=registered) + result.loaded.append(entry) + logger.info( + "Hook loaded: %s (%s) — %d events from %s", + hook_id, + metadata.name, + registered, + discovered.hook_dir, + ) + + +async def _try_load_hook_script(discovered: DiscoveredHook, registry: LifecycleHookRegistry, hook_id: str) -> int: + hook_py = discovered.hook_dir / "hook.py" + if not hook_py.is_file(): + return 0 + + module_name = f"_hook_{hook_id.replace('-', '_')}" + + sandbox = PluginSandbox() + try: + module = sandbox.load_module_from_file(hook_py, module_name) + except SandboxViolationError as e: + logger.error("Sandbox violation in hook script %s: %s", hook_py, e) + return 0 + except Exception: + logger.exception("Failed to load hook script from %s", hook_py) + return 0 + + _sys.modules[module_name] = module + + registered = 0 + valid_event_values = {e.value for e in HookEvent} + for attr_name in dir(module): + if attr_name not in valid_event_values: + continue + callback = getattr(module, attr_name) + if not callable(callback): + continue + await registry.register(attr_name, callback, key=hook_id) + registered += 1 + logger.debug("Auto-loaded callback '%s' from %s for event %s", attr_name, hook_py, attr_name) + + if registered > 0: + logger.info("Hook script loaded: %s — %d callbacks registered", hook_py, registered) + return registered + + +async def load_hooks_from_dirs( + directories: list[str], + registry: LifecycleHookRegistry, + *, + callbacks: dict[str, dict] | None = None, +) -> HookLoadResult: + combined = HookLoadResult() + + for directory in directories: + result = await load_hooks_from_dir(directory, registry, callbacks=callbacks) + combined.loaded.extend(result.loaded) + combined.skipped.extend(result.skipped) + combined.failed.extend(result.failed) + combined.errors.extend(result.errors) + + return combined + + +def _default_hook_sources(base_dir: str | None = None) -> list[tuple[str, str]]: + from pathlib import Path + + cwd = Path(base_dir) if base_dir else Path.cwd() + sources: list[tuple[str, str]] = [] + + bundled = Path(__file__).resolve().parent.parent / "bundled_hooks" + if bundled.is_dir(): + sources.append(("bundled", str(bundled))) + + managed = cwd / ".forcepilot" / "managed_hooks" + if managed.is_dir(): + sources.append(("managed", str(managed))) + + workspace = cwd / "hooks" + if workspace.is_dir(): + sources.append(("workspace", str(workspace))) + + return sources + + +async def load_all_hook_sources( + registry: LifecycleHookRegistry, + *, + callbacks: dict[str, dict] | None = None, + extra_sources: list[tuple[str, str]] | None = None, + base_dir: str | None = None, +) -> HookLoadResult: + sources = _default_hook_sources(base_dir) + if extra_sources: + priority_order = {"bundled": 0, "managed": 1, "workspace": 2} + existing = {s[0] for s in sources} + for label, path in extra_sources: + if label not in existing: + sources.append((label, path)) + sources.sort(key=lambda x: priority_order.get(x[0], 99)) + + combined = HookLoadResult() + seen_hook_ids: set[str] = set() + + for source_label, source_path in sources: + source_result = await load_hooks_from_dir(source_path, registry, callbacks=callbacks) + combined.errors.extend(source_result.errors) + + for entry in source_result.loaded: + if entry.hook_id in seen_hook_ids: + logger.debug( + "Hook %s from [%s] overridden by higher-priority source", + entry.hook_id, + source_label, + ) + continue + seen_hook_ids.add(entry.hook_id) + combined.loaded.append(entry) + + for entry in source_result.skipped: + if entry.hook_id not in seen_hook_ids: + seen_hook_ids.add(entry.hook_id) + combined.skipped.append(entry) + + for entry in source_result.failed: + if entry.hook_id not in seen_hook_ids: + seen_hook_ids.add(entry.hook_id) + combined.failed.append(entry) + + logger.info( + "Hook sources loaded: %d loaded, %d skipped, %d failed (sources: %s)", + combined.loaded_count, + combined.skipped_count, + combined.failed_count, + [s[0] for s in sources], + ) + + return combined diff --git a/backend/package/yuxi/channel/hooks/hook_metadata.py b/backend/package/yuxi/channel/hooks/hook_metadata.py new file mode 100644 index 00000000..6cf7a4d9 --- /dev/null +++ b/backend/package/yuxi/channel/hooks/hook_metadata.py @@ -0,0 +1,3 @@ +from yuxi.channel.hooks.types import HookConfigField, HookDependencies, HookMetadata + +__all__ = ["HookConfigField", "HookDependencies", "HookMetadata"] diff --git a/backend/package/yuxi/channel/hooks/hook_parser.py b/backend/package/yuxi/channel/hooks/hook_parser.py new file mode 100644 index 00000000..fe63427b --- /dev/null +++ b/backend/package/yuxi/channel/hooks/hook_parser.py @@ -0,0 +1,191 @@ +import logging +import os +import re +import sys +import time +from dataclasses import dataclass +from pathlib import Path + +import yaml + +from yuxi.channel.hooks.lifecycle import HookEvent +from yuxi.channel.hooks.types import HookConfigField, HookDependencies, HookMetadata + +logger = logging.getLogger(__name__) + +_FRONTMATTER_RE = re.compile(r"^---\s*\n(.*?)\n---\s*\n?(.*)", re.DOTALL) + +_REQUIRED_FIELDS = {"hook_id", "name", "events"} +_VALID_CONFIG_TYPES = {"str", "int", "float", "bool", "list", "dict"} + +_OS_ALIASES = { + "win32": "windows", + "cygwin": "windows", + "linux": "linux", + "darwin": "darwin", + "macos": "darwin", +} + + +def _normalize_os(name: str) -> str: + return _OS_ALIASES.get(name.lower(), name.lower()) + + +def _current_os() -> str: + return _normalize_os(sys.platform) + + +class HookParseError(Exception): + pass + + +class HookValidationError(Exception): + pass + + +def parse_hook_md(content: str, source_path: Path | str = "") -> HookMetadata: + m = _FRONTMATTER_RE.match(content) + if not m: + raise HookParseError(f"HOOK.md missing frontmatter: {source_path}") + + raw_yaml = m.group(1) + body = m.group(2).strip() + + try: + data = yaml.safe_load(raw_yaml) + except yaml.YAMLError as e: + raise HookParseError(f"HOOK.md invalid YAML frontmatter: {source_path}: {e}") from e + + if not isinstance(data, dict): + raise HookParseError(f"HOOK.md frontmatter is not a dict: {source_path}") + + return _build_metadata(data, body, str(source_path)) + + +def _build_metadata(data: dict, body: str, source_dir: str) -> HookMetadata: + missing = _REQUIRED_FIELDS - set(data.keys()) + if missing: + raise HookValidationError(f"HOOK.md missing required fields: {sorted(missing)} (source: {source_dir})") + + hook_id = data["hook_id"] + name = data["name"] + events = data["events"] + if not isinstance(hook_id, str): + raise HookValidationError(f"HOOK.md hook_id must be a string: {source_dir}") + if not isinstance(name, str): + raise HookValidationError(f"HOOK.md name must be a string: {source_dir}") + if not isinstance(events, list): + raise HookValidationError(f"HOOK.md events must be a list: {source_dir}") + + valid_event_values = {e.value for e in HookEvent} + invalid_events = [e for e in events if e not in valid_event_values] + if invalid_events: + raise HookValidationError( + f"HOOK.md contains invalid events: {invalid_events}. " + f"Valid events: {sorted(valid_event_values)}. (source: {source_dir})" + ) + + description = data.get("description", "") + if not isinstance(description, str): + description = str(description) + + os_list = data.get("os", []) + if isinstance(os_list, list): + os_list = [_normalize_os(o) for o in os_list if isinstance(o, str)] + else: + os_list = [] + + depends_raw = data.get("depends", {}) + if not isinstance(depends_raw, dict): + depends_raw = {} + depends = HookDependencies( + bin=list(depends_raw.get("bin", []) or []), + env=list(depends_raw.get("env", []) or []), + ) + + config_raw = data.get("config", {}) + if not isinstance(config_raw, dict): + config_raw = {} + config_fields: dict[str, HookConfigField] = {} + for key, val in config_raw.items(): + if not isinstance(val, dict): + continue + field_type = val.get("type", "str") + if field_type not in _VALID_CONFIG_TYPES: + field_type = "str" + config_fields[key] = HookConfigField( + type=field_type, + default=val.get("default"), + description=val.get("description", ""), + ) + + return HookMetadata( + hook_id=hook_id, + name=name, + events=events, + description=description, + os=os_list, + depends=depends, + config=config_fields, + body=body, + source_dir=source_dir, + ) + + +def validate_os_compatibility(metadata: HookMetadata) -> bool: + if not metadata.os: + return True + return _current_os() in metadata.os + + +def validate_dependencies(metadata: HookMetadata) -> tuple[bool, list[str]]: + failures: list[str] = [] + + for bin_name in metadata.depends.bin: + if not _find_bin(bin_name): + failures.append(f"binary not found: {bin_name}") + + for env_var in metadata.depends.env: + if env_var not in os.environ: + failures.append(f"env variable not set: {env_var}") + + return len(failures) == 0, failures + + +_CACHE_TTL = 300 + + +@dataclass +class _BinCacheEntry: + value: bool + timestamp: float + + +_find_bin_cache: dict[str, _BinCacheEntry] = {} + + +def clear_bin_cache() -> None: + _find_bin_cache.clear() + + +def _find_bin(name: str) -> bool: + entry = _find_bin_cache.get(name) + if entry is not None: + if time.monotonic() - entry.timestamp < _CACHE_TTL: + return entry.value + del _find_bin_cache[name] + + if os.name == "nt": + names = [name, f"{name}.exe", f"{name}.cmd", f"{name}.bat"] + else: + names = [name] + + for p in os.environ.get("PATH", "").split(os.pathsep): + for n in names: + candidate = os.path.join(p, n) + if os.path.isfile(candidate) and os.access(candidate, os.X_OK): + _find_bin_cache[name] = _BinCacheEntry(value=True, timestamp=time.monotonic()) + return True + + _find_bin_cache[name] = _BinCacheEntry(value=False, timestamp=time.monotonic()) + return False diff --git a/backend/package/yuxi/channel/hooks/hook_scanner.py b/backend/package/yuxi/channel/hooks/hook_scanner.py new file mode 100644 index 00000000..69296abb --- /dev/null +++ b/backend/package/yuxi/channel/hooks/hook_scanner.py @@ -0,0 +1,106 @@ +import logging +from dataclasses import dataclass, field +from pathlib import Path + +from yuxi.channel.hooks.hook_parser import HookParseError, HookValidationError, parse_hook_md +from yuxi.channel.hooks.types import HookMetadata + +logger = logging.getLogger(__name__) + +_SKIP_DIRS = frozenset( + { + "__pycache__", + ".git", + ".svn", + ".hg", + ".venv", + "venv", + ".tox", + ".eggs", + ".mypy_cache", + ".pytest_cache", + ".ruff_cache", + "node_modules", + ".idea", + ".vscode", + } +) + + +@dataclass +class DiscoveredHook: + hook_dir: Path + metadata: HookMetadata + + @property + def source(self) -> str: + return str(self.hook_dir) + + +@dataclass +class ScanResult: + hooks: list[DiscoveredHook] = field(default_factory=list) + errors: list[str] = field(default_factory=list) + + @property + def hook_count(self) -> int: + return len(self.hooks) + + @property + def succeeded(self) -> bool: + return len(self.errors) == 0 + + +def scan_directory(directory: Path | str, *, recursive: bool = True, max_depth: int = 10) -> ScanResult: + hook_dirs = _find_hook_dirs(Path(directory), recursive=recursive, max_depth=max_depth) + return _parse_hook_dirs(hook_dirs) + + +def _find_hook_dirs(root: Path, *, recursive: bool = True, max_depth: int = 10, _depth: int = 0) -> list[Path]: + root = root.resolve() + if not root.is_dir(): + return [] + if _depth > max_depth: + return [] + + result: list[Path] = [] + + candidate = root / "HOOK.md" + if candidate.is_file(): + result.append(root) + + try: + for entry in sorted(root.iterdir()): + if not entry.is_dir(): + continue + if entry.name in _SKIP_DIRS: + continue + candidate = entry / "HOOK.md" + if candidate.is_file(): + result.append(entry) + elif recursive: + result.extend(_find_hook_dirs(entry, recursive=recursive, max_depth=max_depth, _depth=_depth + 1)) + except PermissionError: + logger.warning("Permission denied scanning directory: %s", root) + + return result + + +def _parse_hook_dirs(hook_dirs: list[Path]) -> ScanResult: + result = ScanResult() + + for hook_dir in hook_dirs: + hook_md_path = hook_dir / "HOOK.md" + try: + content = hook_md_path.read_text(encoding="utf-8") + metadata = parse_hook_md(content, source_path=hook_md_path) + result.hooks.append(DiscoveredHook(hook_dir=hook_dir, metadata=metadata)) + logger.debug("Discovered hook: %s from %s", metadata.hook_id, hook_dir) + except (HookParseError, HookValidationError) as e: + result.errors.append(str(e)) + logger.warning("Hook discovery failed: %s", e) + except OSError as e: + result.errors.append(f"Cannot read {hook_md_path}: {e}") + logger.warning("Cannot read HOOK.md: %s", e) + + return result diff --git a/backend/package/yuxi/channel/hooks/lifecycle.py b/backend/package/yuxi/channel/hooks/lifecycle.py new file mode 100644 index 00000000..12bb0e7f --- /dev/null +++ b/backend/package/yuxi/channel/hooks/lifecycle.py @@ -0,0 +1,224 @@ +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()} diff --git a/backend/package/yuxi/channel/hooks/types.py b/backend/package/yuxi/channel/hooks/types.py new file mode 100644 index 00000000..d7a64961 --- /dev/null +++ b/backend/package/yuxi/channel/hooks/types.py @@ -0,0 +1,53 @@ +from dataclasses import dataclass, field +from enum import StrEnum +from typing import Any + + +class SkipReason(StrEnum): + OS_INCOMPATIBLE = "os_incompatible" + DEPENDENCY_MISSING = "dependency_missing" + NO_CALLBACK = "no_callback" + + +@dataclass +class HookConfigField: + type: str = "str" + default: Any = None + description: str = "" + + +@dataclass +class HookDependencies: + bin: list[str] = field(default_factory=list) + env: list[str] = field(default_factory=list) + + +@dataclass +class HookMetadata: + hook_id: str + name: str + events: list[str] + description: str = "" + os: list[str] = field(default_factory=list) + depends: HookDependencies = field(default_factory=HookDependencies) + config: dict[str, HookConfigField] = field(default_factory=dict) + body: str = "" + source_dir: str = "" + + def to_dict(self) -> dict[str, Any]: + return { + "hook_id": self.hook_id, + "name": self.name, + "description": self.description, + "events": self.events, + "os": self.os, + "depends": { + "bin": self.depends.bin, + "env": self.depends.env, + }, + "config": { + k: {"type": v.type, "default": v.default, "description": v.description} + for k, v in self.config.items() + }, + "source_dir": self.source_dir, + }