MediaRouter / app /generation /model_registry.py
basyx's picture
Upload 340 files
3493993 verified
Raw
History Blame Contribute Delete
5.22 kB
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
@model_validator(mode="after")
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),
)