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),
        )