first commit
This commit is contained in:
199
core/module_loader.py
Normal file
199
core/module_loader.py
Normal file
@@ -0,0 +1,199 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import logging
|
||||
from importlib.metadata import entry_points
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from core.base import BaseWorker
|
||||
from core.plugin_types import HYDROGEN_API_VERSION, EmptyPluginConfig
|
||||
from core.schemas import ModuleManifest
|
||||
from core.schemas.config import Config, ResolvedConfig
|
||||
from core.schemas.modules import LoadedModule, LoadedWorker, ModuleLoadResult
|
||||
from reporting.models import PluginRuntimeError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def load_module(package_name: str) -> LoadedModule | None:
|
||||
module = importlib.import_module(package_name)
|
||||
raw_manifest = getattr(module, "MANIFEST", None)
|
||||
|
||||
if raw_manifest is None:
|
||||
logger.warning("module package %s does not export MANIFEST", package_name)
|
||||
return None
|
||||
|
||||
manifest = ModuleManifest.model_validate(raw_manifest)
|
||||
if manifest.api_version != HYDROGEN_API_VERSION:
|
||||
raise ValueError(
|
||||
f"module package {package_name} targets api_version={manifest.api_version}, "
|
||||
f"expected {HYDROGEN_API_VERSION}"
|
||||
)
|
||||
|
||||
builder = getattr(module, "build_worker", None)
|
||||
|
||||
if builder is None:
|
||||
logger.warning("module package %s does not export build_worker", package_name)
|
||||
return None
|
||||
|
||||
if not callable(builder):
|
||||
logger.warning("build_worker in package %s is not callable", package_name)
|
||||
return None
|
||||
|
||||
config_model = getattr(module, "CONFIG_MODEL", None)
|
||||
if config_model is None:
|
||||
config_model = getattr(getattr(builder, "__self__", None), "config_model", None)
|
||||
if config_model is None:
|
||||
config_model = EmptyPluginConfig
|
||||
if not isinstance(config_model, type) or not issubclass(config_model, BaseModel):
|
||||
raise TypeError(f"module package {package_name} exports invalid CONFIG_MODEL")
|
||||
|
||||
return LoadedModule(
|
||||
package_name=package_name,
|
||||
manifest=manifest,
|
||||
build_worker=builder,
|
||||
config_model=config_model,
|
||||
)
|
||||
|
||||
|
||||
def build_worker(
|
||||
loaded_module: LoadedModule, config: BaseModel | dict[str, Any]
|
||||
) -> LoadedWorker | None:
|
||||
worker = loaded_module.build_worker(config)
|
||||
if not isinstance(worker, BaseWorker):
|
||||
logger.warning(
|
||||
"build_worker in package %s did not return BaseWorker", loaded_module.package_name
|
||||
)
|
||||
return None
|
||||
|
||||
return LoadedWorker(worker=worker, manifest=loaded_module.manifest, config=config)
|
||||
|
||||
|
||||
def discover_modules(path: Path, config: Config, package_prefix: str = "modules") -> list[str]:
|
||||
package_names: list[str] = []
|
||||
|
||||
if path.exists():
|
||||
for entry in path.iterdir():
|
||||
if not entry.is_dir() or entry.name.startswith("__"):
|
||||
continue
|
||||
if not (entry / "__init__.py").exists():
|
||||
continue
|
||||
package_names.append(f"{package_prefix}.{entry.name}")
|
||||
|
||||
package_names.extend(config.plugin_packages.modules)
|
||||
|
||||
try:
|
||||
module_entry_points = entry_points(group="hydrogen.modules")
|
||||
except TypeError:
|
||||
module_entry_points = entry_points().select(group="hydrogen.modules")
|
||||
|
||||
for module_entry_point in module_entry_points:
|
||||
package_names.append(module_entry_point.value.partition(":")[0])
|
||||
|
||||
return list(dict.fromkeys(package_names))
|
||||
|
||||
|
||||
def load_modules(path: Path, config: Config, package_prefix: str = "modules") -> ModuleLoadResult:
|
||||
workers: list[LoadedWorker] = []
|
||||
errors: list[PluginRuntimeError] = []
|
||||
|
||||
for package_name in discover_modules(path, config, package_prefix):
|
||||
try:
|
||||
module = load_module(package_name)
|
||||
if module is None:
|
||||
continue
|
||||
|
||||
if module.manifest.category in config.exclude_categories:
|
||||
logger.info(
|
||||
"module %s was skipped due to config exclusion", module.manifest.identifier
|
||||
)
|
||||
continue
|
||||
|
||||
module_config = config.modules.get(module.manifest.identifier, {})
|
||||
validated_config = module.config_model.model_validate(module_config)
|
||||
worker = build_worker(module, validated_config)
|
||||
if worker is not None:
|
||||
workers.append(worker)
|
||||
logger.info("module %s has been successfully loaded!", module.manifest.identifier)
|
||||
except Exception as exc:
|
||||
logger.exception("failed to load module package %s", package_name)
|
||||
errors.append(
|
||||
PluginRuntimeError(
|
||||
plugin_kind="module",
|
||||
plugin_name=package_name,
|
||||
stage="load",
|
||||
message=str(exc),
|
||||
)
|
||||
)
|
||||
|
||||
return ModuleLoadResult(loaded_workers=workers, errors=errors)
|
||||
|
||||
|
||||
def load_prevalidated_workers(
|
||||
loaded_modules: list[LoadedModule],
|
||||
resolved_config: ResolvedConfig,
|
||||
) -> ModuleLoadResult:
|
||||
workers: list[LoadedWorker] = []
|
||||
errors: list[PluginRuntimeError] = []
|
||||
|
||||
for loaded_module in loaded_modules:
|
||||
if loaded_module.manifest.category in resolved_config.exclude_categories:
|
||||
continue
|
||||
|
||||
module_config = resolved_config.modules.get(loaded_module.manifest.identifier)
|
||||
if module_config is None:
|
||||
continue
|
||||
|
||||
try:
|
||||
worker = build_worker(loaded_module, module_config)
|
||||
if worker is None:
|
||||
errors.append(
|
||||
PluginRuntimeError(
|
||||
plugin_kind="module",
|
||||
plugin_name=loaded_module.manifest.identifier,
|
||||
stage="build_worker",
|
||||
message="build_worker returned an invalid worker instance",
|
||||
)
|
||||
)
|
||||
continue
|
||||
workers.append(worker)
|
||||
except Exception as exc:
|
||||
logger.exception("failed to build worker for module %s", loaded_module.package_name)
|
||||
errors.append(
|
||||
PluginRuntimeError(
|
||||
plugin_kind="module",
|
||||
plugin_name=loaded_module.manifest.identifier,
|
||||
stage="build_worker",
|
||||
message=str(exc),
|
||||
)
|
||||
)
|
||||
|
||||
return ModuleLoadResult(loaded_workers=workers, errors=errors)
|
||||
|
||||
|
||||
def load_module_descriptors(
|
||||
path: Path, config: Config, package_prefix: str = "modules"
|
||||
) -> tuple[list[LoadedModule], list[PluginRuntimeError]]:
|
||||
modules: list[LoadedModule] = []
|
||||
errors: list[PluginRuntimeError] = []
|
||||
|
||||
for package_name in discover_modules(path, config, package_prefix):
|
||||
try:
|
||||
module = load_module(package_name)
|
||||
if module is not None:
|
||||
modules.append(module)
|
||||
except Exception as exc:
|
||||
logger.exception("failed to discover module package %s", package_name)
|
||||
errors.append(
|
||||
PluginRuntimeError(
|
||||
plugin_kind="module",
|
||||
plugin_name=package_name,
|
||||
stage="discovery",
|
||||
message=str(exc),
|
||||
)
|
||||
)
|
||||
|
||||
return modules, errors
|
||||
Reference in New Issue
Block a user