Back to home page

EIC code displayed by LXR

 
 

    


File indexing completed on 2026-08-12 08:24:56

0001 """Registry for runner-specific scheduler config models.
0002 
0003 Follows the optimizer _registry pattern: register/get/list_registered. The
0004 registry maps runner type names (e.g., "JobLibRunner") to Pydantic config
0005 classes (e.g., JobLibRunnerConfig) used to validate runner-specific parameters.
0006 """
0007 
0008 from typing import Callable, Dict, Type, Optional
0009 from pydantic import BaseModel
0010 
0011 _runner_config_registry: Dict[str, Type[BaseModel]] = {}
0012 _runner_config_loaders: Dict[str, Callable[[], Type[BaseModel]]] = {
0013     "JobLibRunner": lambda: __import__(
0014         "aid2e.schedulers.JobLib", fromlist=["JobLibRunnerConfig"]
0015     ).JobLibRunnerConfig,
0016     "PanDAiDDSRunner": lambda: __import__(
0017         "aid2e.schedulers.PanDAiDDS", fromlist=["PanDAiDDSRunnerConfig"]
0018     ).PanDAiDDSRunnerConfig,
0019     "SlurmRunner": lambda: __import__(
0020         "aid2e.schedulers.Slurm", fromlist=["SlurmRunnerConfig"]
0021     ).SlurmRunnerConfig,
0022 }
0023 
0024 
0025 def register(name: str, model: Type[BaseModel]) -> None:
0026     """Register a Pydantic model for a scheduler runner type.
0027 
0028     Args:
0029         name: Runner type identifier (e.g., "JobLibRunner", "SlurmRunner", "PanDAiDDSRunner").
0030         model: Pydantic model class that validates runner-specific params.
0031     """
0032     name_key = name
0033     _runner_config_registry[name_key] = model
0034 
0035 
0036 def get(name: str) -> Optional[Type[BaseModel]]:
0037     """Retrieve a registered runner config model by name.
0038 
0039     Args:
0040         name: Runner type identifier.
0041 
0042     Returns:
0043         The registered Pydantic model class, or None if not registered.
0044     """
0045     if name in _runner_config_registry:
0046         return _runner_config_registry[name]
0047     if name in _runner_config_loaders:
0048         model = _runner_config_loaders[name]()
0049         _runner_config_registry[name] = model
0050         return model
0051     return None
0052 
0053 
0054 def list_registered() -> Dict[str, Type[BaseModel]]:
0055     """Get all registered runner config models.
0056 
0057     Returns:
0058         Dict mapping runner type names to Pydantic config classes.
0059     """
0060     for name in list(_runner_config_loaders.keys()):
0061         if name not in _runner_config_registry:
0062             try:
0063                 _runner_config_registry[name] = _runner_config_loaders[name]()
0064             except Exception:
0065                 pass
0066     return _runner_config_registry.copy()