Spaces:
Running
Running
File size: 5,222 Bytes
3493993 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 | 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),
)
|