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

282 lines
8.9 KiB
Python

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