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