sand-boxer/src/sandboxer/extensions/registry.py

74 lines
2.4 KiB
Python
Raw Normal View History

"""Extension discovery, validation, and handler resolution."""
from __future__ import annotations
import importlib
from pathlib import Path
from typing import Protocol
import yaml
from sandboxer.models import Extension, Profile
_REPO_ROOT = Path(__file__).resolve().parents[3]
_EXTENSIONS_DIR = _REPO_ROOT / "extensions"
_REQUIRED_CAPABILITY_FIELDS = ("isolation_levels", "pricing_model")
class ExtensionBackend(Protocol):
def provision(
self, profile: Profile, inputs: dict[str, str], host: str
) -> dict[str, str]: ...
def wait_ready(self, handle: dict[str, str]) -> dict[str, str]: ...
def teardown(self, handle: dict[str, str]) -> dict[str, str]: ...
def extensions_dir() -> Path:
return _EXTENSIONS_DIR
def _validate_extension_caps(ext: Extension) -> None:
caps = ext.capabilities
for field in _REQUIRED_CAPABILITY_FIELDS:
if not getattr(caps, field, None):
raise ValueError(f"Extension {ext.id} missing capability field: {field}")
if not ext.handler:
raise ValueError(f"Extension {ext.id} missing handler")
def load_extension(extension_id: str, *, extensions_root: Path | None = None) -> Extension:
root = extensions_root or _EXTENSIONS_DIR
path = root / f"{extension_id}.yaml"
if not path.exists():
raise FileNotFoundError(f"Extension not found: {extension_id} ({path})")
raw = yaml.safe_load(path.read_text())
ext = Extension.model_validate(raw)
if ext.id != extension_id:
raise ValueError(f"Extension id mismatch: file {extension_id}, content {ext.id}")
_validate_extension_caps(ext)
return ext
def load_all_extensions(*, extensions_root: Path | None = None) -> dict[str, Extension]:
root = extensions_root or _EXTENSIONS_DIR
extensions: dict[str, Extension] = {}
if not root.exists():
return extensions
for path in sorted(root.glob("*.yaml")):
ext = load_extension(path.stem, extensions_root=root)
if ext.id in extensions:
raise ValueError(f"Duplicate extension id: {ext.id}")
extensions[ext.id] = ext
return extensions
def resolve_backend(extension: Extension) -> ExtensionBackend:
module_path, _, attr = extension.handler.partition(":")
if not attr:
raise ValueError(f"Invalid handler for {extension.id}: {extension.handler}")
module = importlib.import_module(module_path)
cls = getattr(module, attr)
return cls(extension.config)