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