Spaces:
Sleeping
Sleeping
| from __future__ import annotations | |
| from collections.abc import Iterable | |
| from pydantic import BaseModel, ConfigDict, Field, model_validator | |
| from app.generation.domain.capabilities import GenerationModelCapability | |
| from app.generation.domain.enums import WorkerReadinessStatus | |
| from app.generation.domain.errors import GenerationCapabilityUnsupportedError | |
| from app.generation.domain.runtime import WorkerInfo, WorkerReadiness, safe_worker_metadata | |
| class GenerationModelRegistration(BaseModel): | |
| """Trusted, server-owned model configuration metadata. | |
| ``configuration_reference`` names a server configuration record; it is | |
| deliberately not a URL, credential, or client-selectable worker value. | |
| """ | |
| model_config = ConfigDict(extra="forbid") | |
| provider_id: str = Field(pattern=r"^[a-z][a-z0-9_-]{0,63}$") | |
| model: GenerationModelCapability | |
| configuration_reference: str = Field( | |
| min_length=1, max_length=128, pattern=r"^[a-z][a-z0-9_.-]{0,127}$" | |
| ) | |
| metadata: dict[str, object] = Field(default_factory=dict) | |
| available: bool = False | |
| def no_initial_availability_claim(self) -> "GenerationModelRegistration": | |
| if self.available: | |
| raise ValueError("models may only become available after readiness verification") | |
| return self | |
| class GenerationModelView(BaseModel): | |
| model_config = ConfigDict(extra="forbid") | |
| provider_id: str | |
| model: GenerationModelCapability | |
| configuration_reference: str | |
| metadata: dict[str, object] = Field(default_factory=dict) | |
| available: bool = False | |
| class GenerationModelRegistry: | |
| """Provider-neutral registry with availability proven by worker readiness.""" | |
| def __init__(self, entries: Iterable[GenerationModelRegistration] | None = None) -> None: | |
| self._entries: dict[tuple[str, str], GenerationModelRegistration] = {} | |
| self._availability: dict[tuple[str, str], bool] = {} | |
| for entry in entries or (): | |
| self.register(entry) | |
| def register(self, entry: GenerationModelRegistration) -> None: | |
| key = (entry.provider_id, entry.model.id) | |
| if key in self._entries: | |
| raise ValueError( | |
| f"Duplicate generation model '{entry.model.id}' for '{entry.provider_id}'." | |
| ) | |
| # Never retain accidental credentials in configuration metadata. | |
| safe_metadata = safe_worker_metadata(entry.metadata) | |
| assert isinstance(safe_metadata, dict) | |
| self._entries[key] = entry.model_copy( | |
| update={"metadata": safe_metadata, "available": False} | |
| ) | |
| self._availability[key] = False | |
| def list(self, *, provider_id: str | None = None) -> list[GenerationModelView]: | |
| items = ( | |
| (key, entry) | |
| for key, entry in self._entries.items() | |
| if provider_id is None or key[0] == provider_id | |
| ) | |
| return [self._view(key, entry) for key, entry in sorted(items)] | |
| def get(self, provider_id: str, model_id: str) -> GenerationModelView: | |
| key = (provider_id, model_id) | |
| try: | |
| return self._view(key, self._entries[key]) | |
| except KeyError as exc: | |
| raise GenerationCapabilityUnsupportedError( | |
| f"{provider_id} does not support model '{model_id}'." | |
| ) from exc | |
| def verify_readiness( | |
| self, | |
| *, | |
| provider_id: str, | |
| worker_info: WorkerInfo, | |
| readiness: WorkerReadiness, | |
| provider_configured: bool, | |
| ) -> list[GenerationModelView]: | |
| """Update only models proven by matching info plus readiness. | |
| Liveness by itself is intentionally insufficient: the worker must be | |
| configured, report `ready`, report a loaded model, explicitly list the | |
| registered model ID in readiness, and discover that model with a | |
| matching output modality through `/v1/info`. | |
| """ | |
| available_ids = set(readiness.model_ids) | |
| discovered_models = {model.id: model for model in worker_info.models} | |
| for key, entry in self._entries.items(): | |
| if key[0] != provider_id: | |
| continue | |
| discovered = discovered_models.get(entry.model.id) | |
| self._availability[key] = bool( | |
| provider_configured | |
| and readiness.status is WorkerReadinessStatus.READY | |
| and readiness.model_loaded | |
| and entry.model.id in available_ids | |
| and discovered is not None | |
| and entry.model.modality in discovered.media_types | |
| ) | |
| return self.list(provider_id=provider_id) | |
| def mark_unavailable(self, provider_id: str) -> None: | |
| for key in self._availability: | |
| if key[0] == provider_id: | |
| self._availability[key] = False | |
| def _view( | |
| self, key: tuple[str, str], entry: GenerationModelRegistration | |
| ) -> GenerationModelView: | |
| return GenerationModelView( | |
| provider_id=entry.provider_id, | |
| model=entry.model, | |
| configuration_reference=entry.configuration_reference, | |
| metadata=entry.metadata, | |
| available=self._availability.get(key, False), | |
| ) | |