feat(model-registry): add delegation-JWT auth and config/snapshot HTTP API
- BFF service token (constant-time, plaintext or SHA-256 hash) plus X-Evo-Actor delegation JWT verification (ES256/RS256, iss/aud, <=60s lifetime, required claims, thread binding) with atomic jti anti-replay - Config API: GET/PUT /api/model-registry, credential rotation endpoint, GET /api/models selector; PUT runs the section 9.2 save-time checks inside the registry write transaction after credential writes - Snapshot API: create/bind/delete routes delegating to SnapshotService with thread/deployment binding checks and 9.5 unified error payloads - Platform security config loader (config.yaml fields), OpenAPI export (scripts/export_model_registry_schema.py -> model_registry/openapi.json) - Mount new routes in langgraph_dev/http.py; retire the legacy GET /api/models and POST /api/runtime-snapshots handlers - Declare PyJWT>=2.8 (previously transitive); extend the 9.5 error code table with the HTTP-layer codes (400/401/403/422/500)
This commit is contained in:
@@ -56,7 +56,6 @@ from EvoScientist.config.provider_profiles import (
|
||||
list_configured_model_entries,
|
||||
load_provider_profiles,
|
||||
provider_profiles_public,
|
||||
provider_profiles_public_revision,
|
||||
provider_profiles_revision,
|
||||
replace_provider_profiles,
|
||||
resolve_provider_profile_draft,
|
||||
@@ -76,7 +75,7 @@ from EvoScientist.llm.provider_operations import (
|
||||
discover_provider_models,
|
||||
test_provider_model,
|
||||
)
|
||||
from EvoScientist.llm.runtime_snapshots import create_run_runtime_snapshot
|
||||
from EvoScientist.model_registry.http_api import model_registry_routes
|
||||
from EvoScientist.sessions import (
|
||||
MAIN_THREAD_FILTER_PARAMS,
|
||||
MAIN_THREAD_FILTER_SQL,
|
||||
@@ -557,61 +556,6 @@ def _patch_llm_config(payload: Any) -> dict[str, Any]:
|
||||
return response
|
||||
|
||||
|
||||
async def get_models(_request: Request) -> JSONResponse:
|
||||
"""Return the model registry as ``{entries, default}``.
|
||||
|
||||
Managed built-in and custom entries come from ``providers.yaml``. Before
|
||||
that file exists, the legacy ``config.yaml`` model catalog/static registry
|
||||
remains available for backward compatibility.
|
||||
|
||||
``default`` is the configured pair when that pair remains available. If a
|
||||
WebUI registry has replaced the legacy catalog and the persisted pair is no
|
||||
longer present, the first enabled registry entry becomes the effective
|
||||
default so new WebUI runs cannot route through stale provider settings.
|
||||
|
||||
Uses ``get_effective_config()`` (not ``load_config()``) so env-var
|
||||
overrides like ``OLLAMA_BASE_URL`` from ``_ENV_MAPPINGS`` are
|
||||
honored — matching the deploy's actual model-building behavior.
|
||||
Offloaded to a thread because ``get_effective_config()`` calls
|
||||
``find_dotenv(usecwd=True)`` which invokes ``os.getcwd()`` — a
|
||||
blocking syscall that langgraph-dev's ``blockbuster`` middleware
|
||||
refuses to allow on the async event loop (would surface as a 500).
|
||||
"""
|
||||
cfg = await asyncio.to_thread(get_effective_config)
|
||||
entries = [
|
||||
{"name": name, "model_id": model_id, "provider": provider}
|
||||
for name, model_id, provider in await list_model_picker_entries(
|
||||
getattr(cfg, "ollama_base_url", None),
|
||||
include_custom_ollama=False,
|
||||
model_catalog=getattr(cfg, "model_catalog", None),
|
||||
)
|
||||
]
|
||||
default_entry = next(
|
||||
(
|
||||
{"name": entry["name"], "provider": entry["provider"]}
|
||||
for entry in entries
|
||||
if entry["name"] == cfg.model and entry["provider"] == cfg.provider
|
||||
),
|
||||
None,
|
||||
)
|
||||
if default_entry is None and entries:
|
||||
default_entry = {
|
||||
"name": entries[0]["name"],
|
||||
"provider": entries[0]["provider"],
|
||||
}
|
||||
try:
|
||||
profiles_revision = await asyncio.to_thread(provider_profiles_public_revision)
|
||||
except ProviderProfileError:
|
||||
profiles_revision = "invalid"
|
||||
return JSONResponse(
|
||||
{
|
||||
"entries": entries,
|
||||
"default": default_entry,
|
||||
"revision": profiles_revision,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _provider_admin_error(request: Request) -> JSONResponse | None:
|
||||
expected = get_provider_admin_token()
|
||||
if not expected:
|
||||
@@ -656,45 +600,6 @@ async def provider_profiles_endpoint(request: Request) -> JSONResponse:
|
||||
)
|
||||
|
||||
|
||||
async def runtime_snapshot_endpoint(request: Request) -> JSONResponse:
|
||||
"""Create an opaque server-side snapshot for one forthcoming graph run."""
|
||||
auth_error = _provider_admin_error(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
|
||||
try:
|
||||
payload = await request.json()
|
||||
if not isinstance(payload, dict):
|
||||
raise ProviderProfileError("Request body must be an object.")
|
||||
snapshot_id = payload.get("snapshot_id")
|
||||
model = payload.get("model")
|
||||
provider = payload.get("provider")
|
||||
if not isinstance(snapshot_id, str):
|
||||
raise ProviderProfileError("snapshot_id is required.")
|
||||
if not isinstance(model, str) or not isinstance(provider, str):
|
||||
raise ProviderProfileError("model and provider are required.")
|
||||
snapshot = await asyncio.to_thread(
|
||||
create_run_runtime_snapshot,
|
||||
snapshot_id,
|
||||
model=model,
|
||||
provider=provider,
|
||||
)
|
||||
return JSONResponse(
|
||||
{"snapshot": snapshot.public_payload() if snapshot is not None else None}
|
||||
)
|
||||
except (json.JSONDecodeError, UnicodeDecodeError):
|
||||
return JSONResponse(
|
||||
{"error": "Request body must be valid JSON."}, status_code=400
|
||||
)
|
||||
except ProviderProfileError as exc:
|
||||
return JSONResponse({"error": str(exc)}, status_code=400)
|
||||
except Exception:
|
||||
_logger.exception("Runtime snapshot request failed")
|
||||
return JSONResponse(
|
||||
{"error": "Runtime snapshot request failed."}, status_code=500
|
||||
)
|
||||
|
||||
|
||||
async def llm_config_endpoint(request: Request) -> JSONResponse:
|
||||
"""Read or patch the LLM-related fields persisted in config.yaml."""
|
||||
auth_error = _provider_admin_error(request)
|
||||
@@ -1670,17 +1575,16 @@ async def bind_workspace_turn(request: Request) -> JSONResponse:
|
||||
|
||||
app = Starlette(
|
||||
routes=[
|
||||
Route("/api/models", get_models, methods=["GET"]),
|
||||
# The unified model registry API (design doc section 9) owns
|
||||
# /api/models and /api/runtime-snapshots; the legacy handlers for
|
||||
# those paths were removed here and the remaining legacy provider
|
||||
# routes are deleted in Task 7.
|
||||
*model_registry_routes(),
|
||||
Route(
|
||||
"/api/provider-profiles",
|
||||
provider_profiles_endpoint,
|
||||
methods=["GET", "PUT"],
|
||||
),
|
||||
Route(
|
||||
"/api/runtime-snapshots",
|
||||
runtime_snapshot_endpoint,
|
||||
methods=["POST"],
|
||||
),
|
||||
Route(
|
||||
"/api/config",
|
||||
llm_config_endpoint,
|
||||
|
||||
@@ -19,10 +19,30 @@ from .adapters import (
|
||||
get_adapter,
|
||||
resolve_parameters,
|
||||
)
|
||||
from .auth import ActorContext, BffAuthenticator
|
||||
from .endpoint_policy import EndpointPolicy
|
||||
from .errors import ERROR_HTTP_STATUS, ErrorDetail, ErrorPayload, ModelRegistryError
|
||||
from .factory import build_chat_model
|
||||
from .hashing import configuration_hash
|
||||
from .http_api import (
|
||||
ApiServices,
|
||||
CredentialReplaceRequest,
|
||||
CredentialWriteResponse,
|
||||
GetModelRegistryResponse,
|
||||
GetSelectableModelsResponse,
|
||||
ModelRegistryHttpApi,
|
||||
PutModelRegistryRequest,
|
||||
SnapshotBindRequest,
|
||||
SnapshotPublicResponse,
|
||||
build_openapi_document,
|
||||
model_registry_routes,
|
||||
validate_registry_save,
|
||||
)
|
||||
from .platform import (
|
||||
PlatformConfigError,
|
||||
PlatformSecurityConfig,
|
||||
load_platform_security_config,
|
||||
)
|
||||
from .resolver import ModelRegistryResolver
|
||||
from .safe_transport import (
|
||||
AsyncSafeHttpTransport,
|
||||
@@ -71,32 +91,43 @@ __all__ = [
|
||||
"BOUND_RETENTION_SECONDS",
|
||||
"ERROR_HTTP_STATUS",
|
||||
"PREPARED_TTL_SECONDS",
|
||||
"ActorContext",
|
||||
"Adapter",
|
||||
"AdapterParameterSpec",
|
||||
"ApiServices",
|
||||
"AsyncSafeHttpTransport",
|
||||
"AsyncSafeNetworkBackend",
|
||||
"AuthConfig",
|
||||
"AuthRef",
|
||||
"AuthSpec",
|
||||
"BffAuthenticator",
|
||||
"BuiltRequest",
|
||||
"Capabilities",
|
||||
"CredentialReplaceRequest",
|
||||
"CredentialStatus",
|
||||
"CredentialWrite",
|
||||
"CredentialWriteResponse",
|
||||
"DevelopmentEndpoint",
|
||||
"EndpointPolicy",
|
||||
"EndpointPolicyPublic",
|
||||
"ErrorDetail",
|
||||
"ErrorPayload",
|
||||
"GetModelRegistryResponse",
|
||||
"GetSelectableModelsResponse",
|
||||
"ModelAvailability",
|
||||
"ModelConfig",
|
||||
"ModelRef",
|
||||
"ModelRegistryError",
|
||||
"ModelRegistryHttpApi",
|
||||
"ModelRegistryResolver",
|
||||
"ModelRuntimeConfig",
|
||||
"ModelRuntimeStore",
|
||||
"ParameterRule",
|
||||
"PlatformConfigError",
|
||||
"PlatformSecurityConfig",
|
||||
"ProviderConfig",
|
||||
"ProviderRuntimeConfig",
|
||||
"PutModelRegistryRequest",
|
||||
"RegistryV4",
|
||||
"ResolvedModelConfig",
|
||||
"ResolvedParameters",
|
||||
@@ -104,13 +135,16 @@ __all__ = [
|
||||
"SafeHttpTransport",
|
||||
"SafeNetworkBackend",
|
||||
"SharedStorageError",
|
||||
"SnapshotBindRequest",
|
||||
"SnapshotCreateRequest",
|
||||
"SnapshotCreation",
|
||||
"SnapshotPayload",
|
||||
"SnapshotPublicResponse",
|
||||
"SnapshotService",
|
||||
"VerificationInfo",
|
||||
"adapter_specs",
|
||||
"build_chat_model",
|
||||
"build_openapi_document",
|
||||
"build_safe_async_http_client",
|
||||
"build_safe_http_client",
|
||||
"compute_effective_capabilities",
|
||||
@@ -119,6 +153,9 @@ __all__ = [
|
||||
"configuration_hash",
|
||||
"find_adapter_spec",
|
||||
"get_adapter",
|
||||
"load_platform_security_config",
|
||||
"model_registry_routes",
|
||||
"public_snapshot_view",
|
||||
"resolve_parameters",
|
||||
"validate_registry_save",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,198 @@
|
||||
"""BFF service-token and delegation-JWT authentication (design doc 7.3).
|
||||
|
||||
Every BFF → EvoScientist request carries::
|
||||
|
||||
Authorization: Bearer <BFF service token>
|
||||
X-Evo-Actor: <short-lived signed delegation JWT>
|
||||
|
||||
The service token only proves the request comes from a trusted WebUI
|
||||
deployment; it is compared in constant time (against the configured
|
||||
plaintext token or the configured SHA-256 hash). The delegation JWT carries
|
||||
the end-user identity: it must be signed by a registered WebUI public key,
|
||||
have ``iss == "WebUI"`` and ``aud == "EvoScientist"``, live at most 60
|
||||
seconds, and carry ``sub``/``scopes``/``deployment_id``/``iat``/``exp``/
|
||||
``jti`` (plus ``thread_id`` for thread-level routes). Each ``jti`` is
|
||||
registered atomically in the ``delegation_jtis`` table; a live duplicate is
|
||||
rejected with ``401 DELEGATION_REPLAYED``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import secrets
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
|
||||
import jwt
|
||||
|
||||
from .errors import (
|
||||
DELEGATION_REPLAYED,
|
||||
FORBIDDEN,
|
||||
UNAUTHENTICATED,
|
||||
ModelRegistryError,
|
||||
)
|
||||
from .platform import DelegationPublicKey
|
||||
|
||||
DELEGATION_ISSUER = "WebUI"
|
||||
DELEGATION_AUDIENCE = "EvoScientist"
|
||||
DELEGATION_MAX_LIFETIME_SECONDS = 60
|
||||
DELEGATION_ALGORITHMS = ("ES256", "RS256")
|
||||
|
||||
_REQUIRED_CLAIMS = ("sub", "scopes", "deployment_id", "iat", "exp", "jti")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ActorContext:
|
||||
"""The verified end-user identity forwarded by the BFF."""
|
||||
|
||||
sub: str
|
||||
scopes: frozenset[str]
|
||||
deployment_id: str
|
||||
thread_id: str | None
|
||||
jti: str
|
||||
|
||||
|
||||
def _unauthenticated(message: str) -> ModelRegistryError:
|
||||
return ModelRegistryError(UNAUTHENTICATED, message)
|
||||
|
||||
|
||||
class BffAuthenticator:
|
||||
"""Verifies the fixed BFF authentication format on every request."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
service_token: str | None,
|
||||
service_token_hash: str | None,
|
||||
delegation_keys: tuple[DelegationPublicKey, ...],
|
||||
jti_store: object,
|
||||
) -> None:
|
||||
if not service_token and not service_token_hash:
|
||||
raise ValueError("A BFF service token or token hash must be configured.")
|
||||
if not delegation_keys:
|
||||
raise ValueError("At least one WebUI delegation key must be registered.")
|
||||
self._service_token = service_token
|
||||
self._service_token_hash = service_token_hash
|
||||
self._delegation_keys = {key.deployment_id: key for key in delegation_keys}
|
||||
self._jti_store = jti_store
|
||||
|
||||
# --- service token ---------------------------------------------------
|
||||
|
||||
def _check_service_token(self, headers: Mapping[str, str]) -> None:
|
||||
authorization = headers.get("authorization", "")
|
||||
scheme, _, supplied = authorization.partition(" ")
|
||||
if scheme.lower() != "bearer" or not supplied.strip():
|
||||
raise _unauthenticated(
|
||||
"The request must carry 'Authorization: Bearer <BFF service token>'."
|
||||
)
|
||||
supplied = supplied.strip()
|
||||
if self._service_token is not None:
|
||||
if secrets.compare_digest(supplied, self._service_token):
|
||||
return
|
||||
elif self._service_token_hash is not None:
|
||||
digest = hashlib.sha256(supplied.encode("utf-8")).hexdigest()
|
||||
if secrets.compare_digest(digest, self._service_token_hash.lower()):
|
||||
return
|
||||
raise _unauthenticated("The BFF service token is invalid.")
|
||||
|
||||
# --- delegation JWT ----------------------------------------------------
|
||||
|
||||
def _decode_delegation(self, headers: Mapping[str, str]) -> dict:
|
||||
token = headers.get("x-evo-actor", "")
|
||||
if not token.strip():
|
||||
raise _unauthenticated(
|
||||
"The request must carry an 'X-Evo-Actor' delegation JWT."
|
||||
)
|
||||
try:
|
||||
unverified = jwt.decode(token, options={"verify_signature": False})
|
||||
except jwt.PyJWTError as exc:
|
||||
raise _unauthenticated("The delegation JWT is malformed.") from exc
|
||||
deployment_id = unverified.get("deployment_id")
|
||||
key = (
|
||||
self._delegation_keys.get(deployment_id)
|
||||
if isinstance(deployment_id, str)
|
||||
else None
|
||||
)
|
||||
if key is None:
|
||||
raise _unauthenticated(
|
||||
"The delegation JWT references an unregistered deployment."
|
||||
)
|
||||
try:
|
||||
return jwt.decode(
|
||||
token,
|
||||
key.public_key,
|
||||
algorithms=list(DELEGATION_ALGORITHMS),
|
||||
audience=DELEGATION_AUDIENCE,
|
||||
issuer=DELEGATION_ISSUER,
|
||||
options={"require": list(_REQUIRED_CLAIMS)},
|
||||
)
|
||||
except jwt.PyJWTError as exc:
|
||||
raise _unauthenticated(
|
||||
"The delegation JWT failed signature, audience, issuer, or "
|
||||
"lifetime validation."
|
||||
) from exc
|
||||
|
||||
# --- combined check ------------------------------------------------------
|
||||
|
||||
def authenticate(
|
||||
self,
|
||||
headers: Mapping[str, str],
|
||||
*,
|
||||
required_scope: str,
|
||||
require_thread_id: bool,
|
||||
) -> ActorContext:
|
||||
"""Verify both credentials and return the actor, or raise 401/403.
|
||||
|
||||
``required_scope`` is the section 7.2 scope the route demands;
|
||||
thread-level routes additionally require a ``thread_id`` claim.
|
||||
"""
|
||||
self._check_service_token(headers)
|
||||
claims = self._decode_delegation(headers)
|
||||
|
||||
if not isinstance(claims["sub"], str) or not claims["sub"].strip():
|
||||
raise _unauthenticated("The delegation JWT 'sub' claim is invalid.")
|
||||
if not isinstance(claims["jti"], str) or not claims["jti"].strip():
|
||||
raise _unauthenticated("The delegation JWT 'jti' claim is invalid.")
|
||||
lifetime = int(claims["exp"]) - int(claims["iat"])
|
||||
if lifetime > DELEGATION_MAX_LIFETIME_SECONDS:
|
||||
raise _unauthenticated(
|
||||
"The delegation JWT lifetime exceeds "
|
||||
f"{DELEGATION_MAX_LIFETIME_SECONDS} seconds."
|
||||
)
|
||||
scopes = claims["scopes"]
|
||||
if not isinstance(scopes, list) or not all(
|
||||
isinstance(scope, str) for scope in scopes
|
||||
):
|
||||
raise _unauthenticated("The delegation JWT 'scopes' claim is invalid.")
|
||||
thread_id = claims.get("thread_id")
|
||||
if require_thread_id and (
|
||||
not isinstance(thread_id, str) or not thread_id.strip()
|
||||
):
|
||||
raise _unauthenticated(
|
||||
"This route requires a delegation JWT 'thread_id' claim."
|
||||
)
|
||||
if thread_id is not None and not isinstance(thread_id, str):
|
||||
raise _unauthenticated("The delegation JWT 'thread_id' claim is invalid.")
|
||||
|
||||
if required_scope not in scopes:
|
||||
raise ModelRegistryError(
|
||||
FORBIDDEN,
|
||||
f"The actor lacks the required scope {required_scope!r}.",
|
||||
)
|
||||
|
||||
registered = self._jti_store.register_delegation_jti(
|
||||
claims["jti"], int(claims["exp"])
|
||||
)
|
||||
if not registered:
|
||||
raise ModelRegistryError(
|
||||
DELEGATION_REPLAYED,
|
||||
"The delegation JWT has already been used.",
|
||||
)
|
||||
|
||||
return ActorContext(
|
||||
sub=claims["sub"].strip(),
|
||||
scopes=frozenset(scopes),
|
||||
deployment_id=claims["deployment_id"],
|
||||
thread_id=thread_id.strip() if isinstance(thread_id, str) else None,
|
||||
jti=claims["jti"],
|
||||
)
|
||||
@@ -15,6 +15,9 @@ from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
# 400 — malformed request bodies.
|
||||
INVALID_REQUEST = "INVALID_REQUEST"
|
||||
|
||||
# 409 — version or idempotency conflicts.
|
||||
REGISTRY_REVISION_CONFLICT = "REGISTRY_REVISION_CONFLICT"
|
||||
THREAD_MODEL_SELECTION_CONFLICT = "THREAD_MODEL_SELECTION_CONFLICT"
|
||||
@@ -25,6 +28,10 @@ MODEL_CONFIGURATION_CHANGED = "MODEL_CONFIGURATION_CHANGED"
|
||||
|
||||
# 401 — invalid identity.
|
||||
DELEGATION_REPLAYED = "DELEGATION_REPLAYED"
|
||||
UNAUTHENTICATED = "UNAUTHENTICATED"
|
||||
|
||||
# 403 — authenticated but insufficient scope.
|
||||
FORBIDDEN = "FORBIDDEN"
|
||||
|
||||
# 404 — missing resource.
|
||||
MODEL_NOT_FOUND = "MODEL_NOT_FOUND"
|
||||
@@ -48,7 +55,14 @@ MODEL_LIMITS_UNCONFIRMED = "MODEL_LIMITS_UNCONFIRMED"
|
||||
CONTEXT_BUDGET_UNSATISFIABLE = "CONTEXT_BUDGET_UNSATISFIABLE"
|
||||
PROVIDER_UNREACHABLE = "PROVIDER_UNREACHABLE"
|
||||
|
||||
# 422 — the request payload failed Pydantic schema validation.
|
||||
VALIDATION_FAILED = "VALIDATION_FAILED"
|
||||
|
||||
# 500 — the platform security configuration is missing or unusable.
|
||||
PLATFORM_CONFIG_MISSING = "PLATFORM_CONFIG_MISSING"
|
||||
|
||||
ERROR_HTTP_STATUS: dict[str, int] = {
|
||||
INVALID_REQUEST: 400,
|
||||
REGISTRY_REVISION_CONFLICT: 409,
|
||||
THREAD_MODEL_SELECTION_CONFLICT: 409,
|
||||
RUN_REQUEST_CONFLICT: 409,
|
||||
@@ -56,6 +70,8 @@ ERROR_HTTP_STATUS: dict[str, int] = {
|
||||
SNAPSHOT_EXPIRED: 409,
|
||||
MODEL_CONFIGURATION_CHANGED: 409,
|
||||
DELEGATION_REPLAYED: 401,
|
||||
UNAUTHENTICATED: 401,
|
||||
FORBIDDEN: 403,
|
||||
MODEL_NOT_FOUND: 404,
|
||||
SNAPSHOT_NOT_FOUND: 404,
|
||||
MODEL_REGISTRY_NOT_READY: 422,
|
||||
@@ -74,6 +90,8 @@ ERROR_HTTP_STATUS: dict[str, int] = {
|
||||
MODEL_LIMITS_UNCONFIRMED: 422,
|
||||
CONTEXT_BUDGET_UNSATISFIABLE: 422,
|
||||
PROVIDER_UNREACHABLE: 422,
|
||||
VALIDATION_FAILED: 422,
|
||||
PLATFORM_CONFIG_MISSING: 500,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,926 @@
|
||||
"""HTTP API for the unified model registry (design doc 7.3, 9.1-9.3, 9.5).
|
||||
|
||||
This module is the only backend entry point the WebUI BFF calls. It mounts:
|
||||
|
||||
- ``GET/PUT /api/model-registry`` and
|
||||
``PUT /api/model-registry/credentials/{credential_id}`` — the Config API,
|
||||
requiring ``model_config:read`` / ``model_config:write`` scopes.
|
||||
- ``GET /api/models`` — the user-facing selector, requiring ``model:select``.
|
||||
- ``POST /api/runtime-snapshots[/{snapshot_id}/bind]`` and
|
||||
``DELETE /api/runtime-snapshots/{snapshot_id}`` — thread-level snapshot
|
||||
routes requiring ``run:create`` plus a matching ``thread_id`` claim.
|
||||
|
||||
Every route passes the fixed BFF authentication format (service bearer token
|
||||
plus ``X-Evo-Actor`` delegation JWT, section 7.3) and answers failures with
|
||||
the unified section 9.5 error payload. The PUT route runs the section 9.2
|
||||
save-time checks inside the registry write transaction, after the
|
||||
``credential_writes`` have been applied, so the verification five-tuple sees
|
||||
the new credential revisions and any failure rolls back without partial
|
||||
updates.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import sqlite3
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Literal
|
||||
|
||||
from pydantic import BaseModel, Field, PositiveInt, TypeAdapter, ValidationError
|
||||
from pydantic.json_schema import models_json_schema
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import JSONResponse, Response
|
||||
from starlette.routing import Route
|
||||
|
||||
from .adapters import find_adapter_spec, resolve_parameters
|
||||
from .auth import ActorContext, BffAuthenticator
|
||||
from .endpoint_policy import EndpointPolicy
|
||||
from .errors import (
|
||||
CREDENTIAL_NOT_CONFIGURED,
|
||||
FORBIDDEN,
|
||||
INVALID_REQUEST,
|
||||
MODEL_LIMITS_UNCONFIRMED,
|
||||
MODEL_NOT_AVAILABLE,
|
||||
PLATFORM_CONFIG_MISSING,
|
||||
SNAPSHOT_NOT_FOUND,
|
||||
VALIDATION_FAILED,
|
||||
ErrorDetail,
|
||||
ErrorPayload,
|
||||
ModelRegistryError,
|
||||
)
|
||||
from .hashing import configuration_hash
|
||||
from .platform import (
|
||||
PlatformConfigError,
|
||||
PlatformSecurityConfig,
|
||||
load_platform_security_config,
|
||||
)
|
||||
from .resolver import ModelRegistryResolver
|
||||
from .schemas import (
|
||||
AdapterParameterSpec,
|
||||
Capabilities,
|
||||
CredentialId,
|
||||
CredentialStatus,
|
||||
CredentialWrite,
|
||||
EndpointPolicyPublic,
|
||||
ModelAvailability,
|
||||
ModelRef,
|
||||
NonEmptyString,
|
||||
RegistryV4,
|
||||
)
|
||||
from .snapshots import (
|
||||
SnapshotCreateRequest,
|
||||
SnapshotCreation,
|
||||
SnapshotService,
|
||||
public_snapshot_view,
|
||||
)
|
||||
from .store import ModelRuntimeStore
|
||||
|
||||
_logger = logging.getLogger(__name__)
|
||||
|
||||
_CREDENTIAL_ID_ADAPTER = TypeAdapter(CredentialId)
|
||||
|
||||
# --- API contract models (sections 9.1-9.3, 8.2); OpenAPI exports these ---
|
||||
|
||||
|
||||
class GetModelRegistryResponse(BaseModel):
|
||||
"""The section 9.1 registry read response."""
|
||||
|
||||
revision: PositiveInt
|
||||
registry: RegistryV4
|
||||
adapter_specs: list[AdapterParameterSpec]
|
||||
credential_status: list[CredentialStatus]
|
||||
model_status: list[ModelAvailability]
|
||||
endpoint_policy: EndpointPolicyPublic
|
||||
|
||||
|
||||
class PutModelRegistryRequest(BaseModel):
|
||||
"""The section 9.2 registry save request."""
|
||||
|
||||
expected_revision: PositiveInt
|
||||
registry: RegistryV4
|
||||
credential_writes: list[CredentialWrite] = Field(default_factory=list)
|
||||
|
||||
|
||||
class CredentialReplaceRequest(BaseModel):
|
||||
"""The section 9.3 credential rotation request."""
|
||||
|
||||
operation: Literal["replace"] = "replace"
|
||||
secret_value: NonEmptyString
|
||||
|
||||
|
||||
class CredentialWriteResponse(BaseModel):
|
||||
"""The section 9.3 credential rotation response; never the plaintext."""
|
||||
|
||||
credential_id: NonEmptyString
|
||||
revision: PositiveInt
|
||||
configured: bool
|
||||
hint: str | None = None
|
||||
updated_at: str | None = None
|
||||
|
||||
|
||||
class SelectableModel(BaseModel):
|
||||
"""One entry of the section 9.1 model selector response."""
|
||||
|
||||
model_ref: ModelRef
|
||||
name: str
|
||||
provider_name: str
|
||||
effective_capabilities: Capabilities
|
||||
|
||||
|
||||
class GetSelectableModelsResponse(BaseModel):
|
||||
"""The section 9.1 selector response: only selectable enabled models."""
|
||||
|
||||
models: list[SelectableModel]
|
||||
|
||||
|
||||
class SnapshotBindRequest(BaseModel):
|
||||
"""The section 9.5 bind request."""
|
||||
|
||||
langgraph_run_id: NonEmptyString
|
||||
|
||||
|
||||
class RuntimeOptionsPublic(BaseModel):
|
||||
"""The public runtime subset of a frozen model config (section 8.2)."""
|
||||
|
||||
max_output_tokens: PositiveInt
|
||||
temperature: float | None = None
|
||||
top_p: float | None = None
|
||||
timeout_seconds: int
|
||||
max_retries: int
|
||||
|
||||
|
||||
class ResolvedModelPublic(BaseModel):
|
||||
"""The public per-role diagnostic subset; never carries secrets."""
|
||||
|
||||
provider_id: str
|
||||
model_key: str
|
||||
adapter_spec_revision: PositiveInt
|
||||
runtime: RuntimeOptionsPublic
|
||||
|
||||
|
||||
class SnapshotPublicResponse(BaseModel):
|
||||
"""The section 8.2 public snapshot diagnostic subset."""
|
||||
|
||||
snapshot_id: str
|
||||
registry_revision: PositiveInt
|
||||
primary: ResolvedModelPublic
|
||||
auxiliary: ResolvedModelPublic | None = None
|
||||
|
||||
|
||||
# --- services bundle --------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ApiServices:
|
||||
"""Everything the HTTP layer needs; built once per deployment."""
|
||||
|
||||
store: ModelRuntimeStore
|
||||
resolver: ModelRegistryResolver
|
||||
snapshot_service: SnapshotService
|
||||
endpoint_policy: EndpointPolicy
|
||||
authenticator: BffAuthenticator
|
||||
|
||||
|
||||
def build_services(platform: PlatformSecurityConfig) -> ApiServices:
|
||||
"""Wire the store, resolver, snapshot service, policy, and authenticator."""
|
||||
store = (
|
||||
ModelRuntimeStore(database_path=platform.model_runtime_db)
|
||||
if platform.model_runtime_db is not None
|
||||
else ModelRuntimeStore()
|
||||
)
|
||||
policy = EndpointPolicy(platform.development_endpoints)
|
||||
resolver = ModelRegistryResolver(store)
|
||||
authenticator = BffAuthenticator(
|
||||
service_token=platform.bff_service_token,
|
||||
service_token_hash=platform.bff_service_token_hash,
|
||||
delegation_keys=platform.webui_delegation_public_keys,
|
||||
jti_store=store,
|
||||
)
|
||||
return ApiServices(
|
||||
store=store,
|
||||
resolver=resolver,
|
||||
snapshot_service=SnapshotService(store, resolver),
|
||||
endpoint_policy=policy,
|
||||
authenticator=authenticator,
|
||||
)
|
||||
|
||||
|
||||
_default_services: ApiServices | None = None
|
||||
|
||||
|
||||
def get_default_services() -> ApiServices:
|
||||
"""Lazily build the production services from ``config.yaml``.
|
||||
|
||||
A missing platform security configuration raises
|
||||
:class:`PlatformConfigError`; the routes map it to
|
||||
``500 PLATFORM_CONFIG_MISSING`` instead of serving unauthenticated.
|
||||
Only a successful build is cached.
|
||||
"""
|
||||
global _default_services
|
||||
if _default_services is None:
|
||||
try:
|
||||
_default_services = build_services(load_platform_security_config())
|
||||
except ValueError as exc:
|
||||
raise PlatformConfigError(str(exc)) from exc
|
||||
return _default_services
|
||||
|
||||
|
||||
# --- section 9.2 save-time validation ---------------------------------------
|
||||
|
||||
_ROOTED_PATH_PREFIXES = ("providers[", "defaults.", "expected_revision")
|
||||
|
||||
|
||||
def _scoped_error(exc: ModelRegistryError, prefix: str) -> ModelRegistryError:
|
||||
"""Re-raise ``exc`` with detail paths rooted at the offending entity."""
|
||||
details = [
|
||||
ErrorDetail(
|
||||
path=(
|
||||
detail.path
|
||||
if detail.path.startswith(_ROOTED_PATH_PREFIXES)
|
||||
else f"{prefix}.{detail.path}"
|
||||
),
|
||||
code=detail.code,
|
||||
)
|
||||
for detail in exc.details
|
||||
]
|
||||
if not details:
|
||||
details = [ErrorDetail(path=prefix, code=exc.code)]
|
||||
return ModelRegistryError(exc.code, exc.message, details=details)
|
||||
|
||||
|
||||
def _check_enabled_model(
|
||||
connection: sqlite3.Connection,
|
||||
provider: Any,
|
||||
model: Any,
|
||||
spec: AdapterParameterSpec,
|
||||
) -> None:
|
||||
"""Section 9.2 rules 4-5: confirmed limits and a passing five-tuple.
|
||||
|
||||
The credential pointer is read from the open transaction, so a
|
||||
``credential_writes`` rotation in the same request is already visible.
|
||||
"""
|
||||
if model.runtime.limits_status != "confirmed":
|
||||
raise ModelRegistryError(
|
||||
MODEL_LIMITS_UNCONFIRMED,
|
||||
"An enabled model requires confirmed context limits.",
|
||||
details=[
|
||||
{"path": "runtime.limits_status", "code": MODEL_LIMITS_UNCONFIRMED}
|
||||
],
|
||||
)
|
||||
auth_spec = spec.auth_specs[provider.auth.mode]
|
||||
credential_revision = 0
|
||||
if auth_spec.credential_required and provider.auth.credential_id is not None:
|
||||
row = connection.execute(
|
||||
"SELECT current_revision FROM credential_pointers WHERE credential_id = ?",
|
||||
(provider.auth.credential_id,),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
raise ModelRegistryError(
|
||||
CREDENTIAL_NOT_CONFIGURED,
|
||||
"The credential referenced by this provider is not configured.",
|
||||
details=[
|
||||
{
|
||||
"path": "auth.credential_id",
|
||||
"code": CREDENTIAL_NOT_CONFIGURED,
|
||||
}
|
||||
],
|
||||
)
|
||||
credential_revision = int(row[0])
|
||||
record = connection.execute(
|
||||
"SELECT result FROM model_verifications "
|
||||
"WHERE provider_id = ? AND model_key = ? "
|
||||
"AND configuration_hash = ? AND credential_revision = ? "
|
||||
"AND adapter_spec_revision = ?",
|
||||
(
|
||||
provider.id,
|
||||
model.key,
|
||||
configuration_hash(provider, model),
|
||||
credential_revision,
|
||||
spec.spec_revision,
|
||||
),
|
||||
).fetchone()
|
||||
if record is None or str(record[0]) != "passed":
|
||||
raise ModelRegistryError(
|
||||
MODEL_NOT_AVAILABLE,
|
||||
"An enabled model requires a passing verification for its "
|
||||
"current configuration, credential revision, and adapter contract.",
|
||||
details=[{"path": "model.verification", "code": MODEL_NOT_AVAILABLE}],
|
||||
)
|
||||
|
||||
|
||||
def validate_registry_save(
|
||||
connection: sqlite3.Connection,
|
||||
registry: RegistryV4,
|
||||
*,
|
||||
specs: tuple[AdapterParameterSpec, ...],
|
||||
policy: EndpointPolicy,
|
||||
) -> None:
|
||||
"""The section 9.2 atomic save-time checks, inside the write transaction.
|
||||
|
||||
Rule 1 (unique, legal IDs) is enforced by the Pydantic schema and rule 7
|
||||
(defaults reference enabled models) by ``ModelRuntimeStore``. Rule 8 has
|
||||
no role-override fields in RegistryV4; capability consistency between
|
||||
declared capabilities and the adapter protocol is checked here through
|
||||
``resolve_parameters``, and the resolve-time tightening rules are
|
||||
enforced by the resolver/adapter layer.
|
||||
"""
|
||||
for provider_index, provider in enumerate(registry.providers):
|
||||
provider_path = f"providers[{provider_index}]"
|
||||
try:
|
||||
policy.validate_base_url(provider.base_url)
|
||||
except ModelRegistryError as exc:
|
||||
raise _scoped_error(exc, f"{provider_path}.base_url") from exc
|
||||
for model_index, model in enumerate(provider.models):
|
||||
model_path = f"{provider_path}.models[{model_index}]"
|
||||
try:
|
||||
spec = find_adapter_spec(
|
||||
provider.adapter, model.upstream_model_id, specs=specs
|
||||
)
|
||||
if spec is None:
|
||||
# No matching contract: the model may only remain
|
||||
# 'configured' (resolver judgement), so rules 3-6 — which
|
||||
# all evaluate against a contract — do not apply.
|
||||
continue
|
||||
parameters = resolve_parameters(provider, model, spec)
|
||||
if provider.enabled and model.enabled:
|
||||
_check_enabled_model(connection, provider, model, spec)
|
||||
# Rule 6 holds whether or not the model is enabled.
|
||||
ModelRegistryResolver._resolve_budget(
|
||||
model, parameters.max_output_tokens
|
||||
)
|
||||
except ModelRegistryError as exc:
|
||||
raise _scoped_error(exc, model_path) from exc
|
||||
|
||||
|
||||
# --- HTTP layer ---------------------------------------------------------------
|
||||
|
||||
|
||||
def _error_response(exc: ModelRegistryError, request_id: str) -> JSONResponse:
|
||||
return JSONResponse(exc.payload(request_id=request_id), status_code=exc.http_status)
|
||||
|
||||
|
||||
def _validation_error_response(exc: ValidationError, request_id: str) -> JSONResponse:
|
||||
details = [
|
||||
ErrorDetail(
|
||||
path=".".join(str(part) for part in error["loc"]) or "(body)",
|
||||
code=str(error["type"]),
|
||||
)
|
||||
for error in exc.errors()
|
||||
]
|
||||
return _error_response(
|
||||
ModelRegistryError(
|
||||
VALIDATION_FAILED, "The request payload failed validation.", details=details
|
||||
),
|
||||
request_id,
|
||||
)
|
||||
|
||||
|
||||
class ModelRegistryHttpApi:
|
||||
"""The Starlette route handlers for the model registry HTTP API."""
|
||||
|
||||
def __init__(self, services_provider: Callable[[], ApiServices]) -> None:
|
||||
self._services_provider = services_provider
|
||||
|
||||
def routes(self) -> list[Route]:
|
||||
return [
|
||||
Route("/api/model-registry", self.get_model_registry, methods=["GET"]),
|
||||
Route("/api/model-registry", self.put_model_registry, methods=["PUT"]),
|
||||
Route(
|
||||
"/api/model-registry/credentials/{credential_id}",
|
||||
self.put_credential,
|
||||
methods=["PUT"],
|
||||
),
|
||||
Route("/api/models", self.get_selectable_models, methods=["GET"]),
|
||||
Route("/api/runtime-snapshots", self.create_snapshot, methods=["POST"]),
|
||||
Route(
|
||||
"/api/runtime-snapshots/{snapshot_id}/bind",
|
||||
self.bind_snapshot,
|
||||
methods=["POST"],
|
||||
),
|
||||
Route(
|
||||
"/api/runtime-snapshots/{snapshot_id}",
|
||||
self.delete_snapshot,
|
||||
methods=["DELETE"],
|
||||
),
|
||||
]
|
||||
|
||||
# --- helpers ----------------------------------------------------------
|
||||
|
||||
def _services(self) -> ApiServices:
|
||||
try:
|
||||
return self._services_provider()
|
||||
except PlatformConfigError:
|
||||
_logger.exception("Model registry platform configuration is unusable")
|
||||
raise ModelRegistryError(
|
||||
PLATFORM_CONFIG_MISSING,
|
||||
"The model registry platform security configuration is "
|
||||
"missing or invalid; the Config API is unavailable.",
|
||||
) from None
|
||||
|
||||
@staticmethod
|
||||
async def _body(request: Request) -> Any:
|
||||
try:
|
||||
return await request.json()
|
||||
except json.JSONDecodeError:
|
||||
raise ModelRegistryError(
|
||||
INVALID_REQUEST, "The request body must be valid JSON."
|
||||
) from None
|
||||
|
||||
def _authenticate(
|
||||
self, request: Request, *, required_scope: str, require_thread_id: bool
|
||||
) -> tuple[ApiServices, ActorContext]:
|
||||
services = self._services()
|
||||
actor = services.authenticator.authenticate(
|
||||
request.headers,
|
||||
required_scope=required_scope,
|
||||
require_thread_id=require_thread_id,
|
||||
)
|
||||
return services, actor
|
||||
|
||||
# --- Config API ----------------------------------------------------------
|
||||
|
||||
async def get_model_registry(self, request: Request) -> Response:
|
||||
request_id = uuid.uuid4().hex
|
||||
try:
|
||||
services, _actor = self._authenticate(
|
||||
request, required_scope="model_config:read", require_thread_id=False
|
||||
)
|
||||
response = await asyncio.to_thread(self._registry_response, services)
|
||||
except ModelRegistryError as exc:
|
||||
return _error_response(exc, request_id)
|
||||
return JSONResponse(response.model_dump(mode="json"))
|
||||
|
||||
async def put_model_registry(self, request: Request) -> Response:
|
||||
request_id = uuid.uuid4().hex
|
||||
try:
|
||||
services, _actor = self._authenticate(
|
||||
request, required_scope="model_config:write", require_thread_id=False
|
||||
)
|
||||
body = await self._body(request)
|
||||
response = await asyncio.to_thread(self._save_registry, services, body)
|
||||
except ModelRegistryError as exc:
|
||||
return _error_response(exc, request_id)
|
||||
except ValidationError as exc:
|
||||
return _validation_error_response(exc, request_id)
|
||||
return JSONResponse(response.model_dump(mode="json"))
|
||||
|
||||
async def put_credential(self, request: Request) -> Response:
|
||||
request_id = uuid.uuid4().hex
|
||||
try:
|
||||
services, _actor = self._authenticate(
|
||||
request, required_scope="model_config:write", require_thread_id=False
|
||||
)
|
||||
body = await self._body(request)
|
||||
credential_id = _CREDENTIAL_ID_ADAPTER.validate_python(
|
||||
request.path_params["credential_id"]
|
||||
)
|
||||
response = await asyncio.to_thread(
|
||||
self._replace_credential,
|
||||
services,
|
||||
credential_id,
|
||||
body,
|
||||
)
|
||||
except ModelRegistryError as exc:
|
||||
return _error_response(exc, request_id)
|
||||
except ValidationError as exc:
|
||||
return _validation_error_response(exc, request_id)
|
||||
return JSONResponse(response.model_dump(mode="json"))
|
||||
|
||||
async def get_selectable_models(self, request: Request) -> Response:
|
||||
request_id = uuid.uuid4().hex
|
||||
try:
|
||||
services, _actor = self._authenticate(
|
||||
request, required_scope="model:select", require_thread_id=False
|
||||
)
|
||||
response = await asyncio.to_thread(self._selectable_models, services)
|
||||
except ModelRegistryError as exc:
|
||||
return _error_response(exc, request_id)
|
||||
return JSONResponse(response.model_dump(mode="json"))
|
||||
|
||||
# --- snapshot API ----------------------------------------------------------
|
||||
|
||||
async def create_snapshot(self, request: Request) -> Response:
|
||||
request_id = uuid.uuid4().hex
|
||||
try:
|
||||
services, actor = self._authenticate(
|
||||
request, required_scope="run:create", require_thread_id=True
|
||||
)
|
||||
body = await self._body(request)
|
||||
creation = await asyncio.to_thread(
|
||||
self._create_snapshot, services, actor, body
|
||||
)
|
||||
except ModelRegistryError as exc:
|
||||
return _error_response(exc, request_id)
|
||||
except ValidationError as exc:
|
||||
return _validation_error_response(exc, request_id)
|
||||
status_code = 201 if creation.created else 200
|
||||
return JSONResponse(
|
||||
self._snapshot_response(creation.snapshot).model_dump(mode="json"),
|
||||
status_code=status_code,
|
||||
)
|
||||
|
||||
async def bind_snapshot(self, request: Request) -> Response:
|
||||
request_id = uuid.uuid4().hex
|
||||
try:
|
||||
services, actor = self._authenticate(
|
||||
request, required_scope="run:create", require_thread_id=True
|
||||
)
|
||||
body = await self._body(request)
|
||||
bound = await asyncio.to_thread(
|
||||
self._bind_snapshot,
|
||||
services,
|
||||
actor,
|
||||
request.path_params["snapshot_id"],
|
||||
body,
|
||||
)
|
||||
except ModelRegistryError as exc:
|
||||
return _error_response(exc, request_id)
|
||||
except ValidationError as exc:
|
||||
return _validation_error_response(exc, request_id)
|
||||
return JSONResponse(self._snapshot_response(bound).model_dump(mode="json"))
|
||||
|
||||
async def delete_snapshot(self, request: Request) -> Response:
|
||||
request_id = uuid.uuid4().hex
|
||||
try:
|
||||
services, actor = self._authenticate(
|
||||
request, required_scope="run:create", require_thread_id=True
|
||||
)
|
||||
await asyncio.to_thread(
|
||||
self._abort_snapshot,
|
||||
services,
|
||||
actor,
|
||||
request.path_params["snapshot_id"],
|
||||
)
|
||||
except ModelRegistryError as exc:
|
||||
return _error_response(exc, request_id)
|
||||
return Response(status_code=204)
|
||||
|
||||
# --- sync workers (run in a thread; SQLite and config I/O block) --------
|
||||
|
||||
@staticmethod
|
||||
def _registry_response(services: ApiServices) -> GetModelRegistryResponse:
|
||||
store = services.store
|
||||
registry = store.load_registry()
|
||||
credential_ids = sorted(
|
||||
{
|
||||
provider.auth.credential_id
|
||||
for provider in registry.providers
|
||||
if provider.auth.credential_id is not None
|
||||
}
|
||||
)
|
||||
return GetModelRegistryResponse(
|
||||
revision=registry.revision,
|
||||
registry=registry,
|
||||
adapter_specs=list(services.resolver.specs),
|
||||
credential_status=[
|
||||
store.credential_status(credential_id)
|
||||
for credential_id in credential_ids
|
||||
],
|
||||
# The full pointer mapping is mandatory: a default/empty mapping
|
||||
# misjudges verified models as stale after a credential rotation.
|
||||
model_status=services.resolver.compute_availability(
|
||||
registry,
|
||||
store.list_model_verifications(),
|
||||
credential_revisions=store.list_credential_revisions(),
|
||||
),
|
||||
endpoint_policy=services.endpoint_policy.public_view(),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _save_registry(services: ApiServices, body: Any) -> GetModelRegistryResponse:
|
||||
request = PutModelRegistryRequest.model_validate(body)
|
||||
|
||||
def validate(connection: sqlite3.Connection, registry: RegistryV4) -> None:
|
||||
validate_registry_save(
|
||||
connection,
|
||||
registry,
|
||||
specs=services.resolver.specs,
|
||||
policy=services.endpoint_policy,
|
||||
)
|
||||
|
||||
services.store.save_registry(
|
||||
expected_revision=request.expected_revision,
|
||||
registry=request.registry,
|
||||
credential_writes=request.credential_writes,
|
||||
validate=validate,
|
||||
)
|
||||
return ModelRegistryHttpApi._registry_response(services)
|
||||
|
||||
@staticmethod
|
||||
def _replace_credential(
|
||||
services: ApiServices, credential_id: str, body: Any
|
||||
) -> CredentialWriteResponse:
|
||||
request = CredentialReplaceRequest.model_validate(body)
|
||||
revision = services.store.write_credential_version(
|
||||
credential_id, request.secret_value
|
||||
)
|
||||
status = services.store.credential_status(credential_id)
|
||||
return CredentialWriteResponse(
|
||||
credential_id=credential_id,
|
||||
revision=revision,
|
||||
configured=status.configured,
|
||||
hint=status.hint,
|
||||
updated_at=status.updated_at,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _selectable_models(services: ApiServices) -> GetSelectableModelsResponse:
|
||||
store = services.store
|
||||
registry = store.load_registry()
|
||||
availability = services.resolver.compute_availability(
|
||||
registry,
|
||||
store.list_model_verifications(),
|
||||
credential_revisions=store.list_credential_revisions(),
|
||||
)
|
||||
models: list[SelectableModel] = []
|
||||
for item in availability:
|
||||
if not item.selectable:
|
||||
continue
|
||||
provider = registry.find_provider(item.model_ref.provider_id)
|
||||
model = (
|
||||
provider.find_model(item.model_ref.model_key)
|
||||
if provider is not None
|
||||
else None
|
||||
)
|
||||
if provider is None or model is None: # pragma: no cover - judged above
|
||||
continue
|
||||
models.append(
|
||||
SelectableModel(
|
||||
model_ref=item.model_ref,
|
||||
name=model.name,
|
||||
provider_name=provider.name,
|
||||
effective_capabilities=item.effective_capabilities,
|
||||
)
|
||||
)
|
||||
return GetSelectableModelsResponse(models=models)
|
||||
|
||||
@staticmethod
|
||||
def _create_snapshot(
|
||||
services: ApiServices, actor: ActorContext, body: Any
|
||||
) -> SnapshotCreation:
|
||||
request = SnapshotCreateRequest.model_validate(body)
|
||||
if request.thread_id != actor.thread_id:
|
||||
raise ModelRegistryError(
|
||||
FORBIDDEN,
|
||||
"The delegation JWT thread does not match the request thread.",
|
||||
)
|
||||
if request.deployment_id != actor.deployment_id:
|
||||
raise ModelRegistryError(
|
||||
FORBIDDEN,
|
||||
"The delegation JWT deployment does not match the request.",
|
||||
)
|
||||
return services.snapshot_service.create(request)
|
||||
|
||||
@staticmethod
|
||||
def _bound_row(services: ApiServices, actor: ActorContext, snapshot_id: str):
|
||||
row = services.store.get_run_snapshot(snapshot_id)
|
||||
if (
|
||||
row is None
|
||||
or row["thread_id"] != actor.thread_id
|
||||
or row["deployment_id"] != actor.deployment_id
|
||||
):
|
||||
raise ModelRegistryError(SNAPSHOT_NOT_FOUND, "The snapshot does not exist.")
|
||||
return row
|
||||
|
||||
@staticmethod
|
||||
def _bind_snapshot(
|
||||
services: ApiServices, actor: ActorContext, snapshot_id: str, body: Any
|
||||
):
|
||||
request = SnapshotBindRequest.model_validate(body)
|
||||
ModelRegistryHttpApi._bound_row(services, actor, snapshot_id)
|
||||
return services.snapshot_service.bind(snapshot_id, request.langgraph_run_id)
|
||||
|
||||
@staticmethod
|
||||
def _abort_snapshot(
|
||||
services: ApiServices, actor: ActorContext, snapshot_id: str
|
||||
) -> None:
|
||||
ModelRegistryHttpApi._bound_row(services, actor, snapshot_id)
|
||||
services.snapshot_service.abort(snapshot_id)
|
||||
|
||||
@staticmethod
|
||||
def _snapshot_response(snapshot: Any) -> SnapshotPublicResponse:
|
||||
return SnapshotPublicResponse.model_validate(public_snapshot_view(snapshot))
|
||||
|
||||
|
||||
def model_registry_routes(
|
||||
services_provider: Callable[[], ApiServices] = get_default_services,
|
||||
) -> list[Route]:
|
||||
"""Return the route table; services resolve lazily per request."""
|
||||
return ModelRegistryHttpApi(services_provider).routes()
|
||||
|
||||
|
||||
# --- OpenAPI / JSON Schema export (section 11 step 3) ------------------------
|
||||
|
||||
_OPENAPI_MODELS = (
|
||||
GetModelRegistryResponse,
|
||||
PutModelRegistryRequest,
|
||||
CredentialReplaceRequest,
|
||||
CredentialWriteResponse,
|
||||
GetSelectableModelsResponse,
|
||||
SnapshotCreateRequest,
|
||||
SnapshotBindRequest,
|
||||
SnapshotPublicResponse,
|
||||
ErrorPayload,
|
||||
)
|
||||
|
||||
|
||||
def _schema_ref(model: type[BaseModel]) -> dict[str, str]:
|
||||
return {"$ref": f"#/components/schemas/{model.__name__}"}
|
||||
|
||||
|
||||
def _error_responses(*status_codes: int) -> dict[str, Any]:
|
||||
return {
|
||||
str(status): {
|
||||
"description": "Unified error payload (section 9.5).",
|
||||
"content": {"application/json": {"schema": _schema_ref(ErrorPayload)}},
|
||||
}
|
||||
for status in status_codes
|
||||
}
|
||||
|
||||
|
||||
def _json_response(status: str, model: type[BaseModel], description: str) -> dict:
|
||||
return {
|
||||
status: {
|
||||
"description": description,
|
||||
"content": {"application/json": {"schema": _schema_ref(model)}},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def build_openapi_document() -> dict[str, Any]:
|
||||
"""Build the OpenAPI 3.1 contract from the same Pydantic models."""
|
||||
_, definitions = models_json_schema(
|
||||
[(model, "validation") for model in _OPENAPI_MODELS],
|
||||
ref_template="#/components/schemas/{model}",
|
||||
)
|
||||
auth_security = [{"BffServiceToken": [], "DelegationJwt": []}]
|
||||
paths: dict[str, Any] = {
|
||||
"/api/model-registry": {
|
||||
"get": {
|
||||
"operationId": "getModelRegistry",
|
||||
"security": auth_security,
|
||||
"responses": {
|
||||
**_json_response(
|
||||
"200", GetModelRegistryResponse, "The current registry."
|
||||
),
|
||||
**_error_responses(401, 403, 500),
|
||||
},
|
||||
},
|
||||
"put": {
|
||||
"operationId": "putModelRegistry",
|
||||
"security": auth_security,
|
||||
"requestBody": {
|
||||
"required": True,
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": _schema_ref(PutModelRegistryRequest)
|
||||
}
|
||||
},
|
||||
},
|
||||
"responses": {
|
||||
**_json_response(
|
||||
"200", GetModelRegistryResponse, "The saved registry."
|
||||
),
|
||||
**_error_responses(400, 401, 403, 409, 422, 500),
|
||||
},
|
||||
},
|
||||
},
|
||||
"/api/model-registry/credentials/{credential_id}": {
|
||||
"put": {
|
||||
"operationId": "replaceCredential",
|
||||
"security": auth_security,
|
||||
"parameters": [
|
||||
{
|
||||
"name": "credential_id",
|
||||
"in": "path",
|
||||
"required": True,
|
||||
"schema": {"type": "string"},
|
||||
}
|
||||
],
|
||||
"requestBody": {
|
||||
"required": True,
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": _schema_ref(CredentialReplaceRequest)
|
||||
}
|
||||
},
|
||||
},
|
||||
"responses": {
|
||||
**_json_response(
|
||||
"200", CredentialWriteResponse, "The rotated credential."
|
||||
),
|
||||
**_error_responses(400, 401, 403, 422, 500),
|
||||
},
|
||||
}
|
||||
},
|
||||
"/api/models": {
|
||||
"get": {
|
||||
"operationId": "getSelectableModels",
|
||||
"security": auth_security,
|
||||
"responses": {
|
||||
**_json_response(
|
||||
"200",
|
||||
GetSelectableModelsResponse,
|
||||
"Enabled models the user may select.",
|
||||
),
|
||||
**_error_responses(401, 403, 500),
|
||||
},
|
||||
}
|
||||
},
|
||||
"/api/runtime-snapshots": {
|
||||
"post": {
|
||||
"operationId": "createRuntimeSnapshot",
|
||||
"security": auth_security,
|
||||
"requestBody": {
|
||||
"required": True,
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": _schema_ref(SnapshotCreateRequest)
|
||||
}
|
||||
},
|
||||
},
|
||||
"responses": {
|
||||
**_json_response(
|
||||
"201", SnapshotPublicResponse, "The created snapshot."
|
||||
),
|
||||
**_json_response(
|
||||
"200",
|
||||
SnapshotPublicResponse,
|
||||
"Idempotent replay of the same run request.",
|
||||
),
|
||||
**_error_responses(400, 401, 403, 409, 422, 500),
|
||||
},
|
||||
}
|
||||
},
|
||||
"/api/runtime-snapshots/{snapshot_id}/bind": {
|
||||
"post": {
|
||||
"operationId": "bindRuntimeSnapshot",
|
||||
"security": auth_security,
|
||||
"parameters": [
|
||||
{
|
||||
"name": "snapshot_id",
|
||||
"in": "path",
|
||||
"required": True,
|
||||
"schema": {"type": "string"},
|
||||
}
|
||||
],
|
||||
"requestBody": {
|
||||
"required": True,
|
||||
"content": {
|
||||
"application/json": {"schema": _schema_ref(SnapshotBindRequest)}
|
||||
},
|
||||
},
|
||||
"responses": {
|
||||
**_json_response(
|
||||
"200", SnapshotPublicResponse, "The bound snapshot."
|
||||
),
|
||||
**_error_responses(400, 401, 403, 404, 409, 500),
|
||||
},
|
||||
}
|
||||
},
|
||||
"/api/runtime-snapshots/{snapshot_id}": {
|
||||
"delete": {
|
||||
"operationId": "abortRuntimeSnapshot",
|
||||
"security": auth_security,
|
||||
"parameters": [
|
||||
{
|
||||
"name": "snapshot_id",
|
||||
"in": "path",
|
||||
"required": True,
|
||||
"schema": {"type": "string"},
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"204": {"description": "The snapshot was aborted."},
|
||||
**_error_responses(401, 403, 404, 409, 500),
|
||||
},
|
||||
}
|
||||
},
|
||||
}
|
||||
return {
|
||||
"openapi": "3.1.0",
|
||||
"info": {
|
||||
"title": "EvoScientist Model Registry API",
|
||||
"version": "1.0.0",
|
||||
"description": (
|
||||
"Config, selector, and run snapshot API consumed by the "
|
||||
"WebUI BFF (design doc sections 7.3, 9.1-9.3, 9.5)."
|
||||
),
|
||||
},
|
||||
"paths": paths,
|
||||
"components": {
|
||||
"schemas": definitions["$defs"],
|
||||
"securitySchemes": {
|
||||
"BffServiceToken": {"type": "http", "scheme": "bearer"},
|
||||
"DelegationJwt": {
|
||||
"type": "apiKey",
|
||||
"in": "header",
|
||||
"name": "X-Evo-Actor",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,171 @@
|
||||
"""Platform security configuration from ``config.yaml`` (design doc 4.2).
|
||||
|
||||
The Config API cannot serve without an explicitly configured BFF service
|
||||
token and at least one registered WebUI delegation public key — silently
|
||||
falling back to defaults would mean accepting unauthenticated requests.
|
||||
A missing or incomplete security section raises :class:`PlatformConfigError`;
|
||||
the HTTP layer maps it to ``500 PLATFORM_CONFIG_MISSING`` with a safe
|
||||
message.
|
||||
|
||||
Recognized ``config.yaml`` fields (all other platform fields stay in
|
||||
``EvoScientistConfig`` and are untouched):
|
||||
|
||||
- ``bff_service_token``: the shared BFF → EvoScientist bearer token
|
||||
(plaintext; development convenience).
|
||||
- ``bff_service_token_hash``: hex SHA-256 of the token; used when the
|
||||
plaintext field is absent so production deployments never persist it.
|
||||
- ``webui_delegation_public_keys``: ``[{deployment_id, public_key}]`` entries
|
||||
registering each WebUI's delegation JWT signing key (PEM, ES256/RS256).
|
||||
- ``development_endpoints``: ``[{id, url, label}]`` EndpointPolicy entries.
|
||||
- ``local_deployment_id``: deployment ID for local (non-BFF) entry points;
|
||||
defaults to ``"local"``.
|
||||
- ``model_runtime_db``: optional explicit path of ``model-runtime.sqlite3``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import yaml
|
||||
|
||||
from EvoScientist.config.settings import get_config_path
|
||||
|
||||
from .schemas import DevelopmentEndpoint
|
||||
|
||||
DEFAULT_LOCAL_DEPLOYMENT_ID = "local"
|
||||
|
||||
|
||||
class PlatformConfigError(RuntimeError):
|
||||
"""Raised when the platform security configuration cannot be used."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DelegationPublicKey:
|
||||
"""One registered WebUI delegation signing key (section 7.3)."""
|
||||
|
||||
deployment_id: str
|
||||
public_key: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PlatformSecurityConfig:
|
||||
"""The parsed security-relevant platform fields of ``config.yaml``."""
|
||||
|
||||
bff_service_token: str | None = None
|
||||
bff_service_token_hash: str | None = None
|
||||
webui_delegation_public_keys: tuple[DelegationPublicKey, ...] = ()
|
||||
development_endpoints: tuple[DevelopmentEndpoint, ...] = ()
|
||||
local_deployment_id: str = DEFAULT_LOCAL_DEPLOYMENT_ID
|
||||
model_runtime_db: Path | None = None
|
||||
|
||||
|
||||
def _optional_string(data: dict[str, Any], key: str) -> str | None:
|
||||
value = data.get(key)
|
||||
if value is None:
|
||||
return None
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
raise PlatformConfigError(f"config.yaml field {key!r} must be a string.")
|
||||
return value.strip()
|
||||
|
||||
|
||||
def _parse_delegation_keys(raw: Any) -> tuple[DelegationPublicKey, ...]:
|
||||
if raw is None:
|
||||
return ()
|
||||
if not isinstance(raw, list):
|
||||
raise PlatformConfigError(
|
||||
"config.yaml field 'webui_delegation_public_keys' must be a list."
|
||||
)
|
||||
keys: list[DelegationPublicKey] = []
|
||||
seen: set[str] = set()
|
||||
for index, entry in enumerate(raw):
|
||||
if not isinstance(entry, dict):
|
||||
raise PlatformConfigError(
|
||||
f"webui_delegation_public_keys[{index}] must be an object."
|
||||
)
|
||||
deployment_id = entry.get("deployment_id")
|
||||
public_key = entry.get("public_key")
|
||||
if (
|
||||
not isinstance(deployment_id, str)
|
||||
or not deployment_id.strip()
|
||||
or not isinstance(public_key, str)
|
||||
or "PUBLIC KEY" not in public_key
|
||||
):
|
||||
raise PlatformConfigError(
|
||||
f"webui_delegation_public_keys[{index}] requires a non-empty "
|
||||
"'deployment_id' and a PEM 'public_key'."
|
||||
)
|
||||
deployment_id = deployment_id.strip()
|
||||
if deployment_id in seen:
|
||||
raise PlatformConfigError(
|
||||
f"Duplicate delegation deployment_id {deployment_id!r}."
|
||||
)
|
||||
seen.add(deployment_id)
|
||||
keys.append(
|
||||
DelegationPublicKey(deployment_id=deployment_id, public_key=public_key)
|
||||
)
|
||||
return tuple(keys)
|
||||
|
||||
|
||||
def _parse_development_endpoints(raw: Any) -> tuple[DevelopmentEndpoint, ...]:
|
||||
if raw is None:
|
||||
return ()
|
||||
if not isinstance(raw, list):
|
||||
raise PlatformConfigError(
|
||||
"config.yaml field 'development_endpoints' must be a list."
|
||||
)
|
||||
try:
|
||||
return tuple(DevelopmentEndpoint.model_validate(entry) for entry in raw)
|
||||
except ValueError as exc:
|
||||
raise PlatformConfigError(
|
||||
f"Invalid development_endpoints entry: {exc}."
|
||||
) from exc
|
||||
|
||||
|
||||
def load_platform_security_config(
|
||||
config_path: str | Path | None = None,
|
||||
) -> PlatformSecurityConfig:
|
||||
"""Load the security platform fields; raise when auth cannot be served."""
|
||||
path = Path(config_path) if config_path is not None else get_config_path()
|
||||
if not path.exists():
|
||||
raise PlatformConfigError(
|
||||
f"config.yaml not found at {path}; configure bff_service_token "
|
||||
"and webui_delegation_public_keys first."
|
||||
)
|
||||
try:
|
||||
with open(path, encoding="utf-8") as handle:
|
||||
data = yaml.safe_load(handle) or {}
|
||||
except yaml.YAMLError as exc:
|
||||
raise PlatformConfigError(f"config.yaml is not valid YAML: {exc}.") from exc
|
||||
if not isinstance(data, dict):
|
||||
raise PlatformConfigError("config.yaml must contain a mapping at the top.")
|
||||
|
||||
token = _optional_string(data, "bff_service_token")
|
||||
token_hash = _optional_string(data, "bff_service_token_hash")
|
||||
if token is None and token_hash is None:
|
||||
raise PlatformConfigError(
|
||||
"config.yaml must set 'bff_service_token' or "
|
||||
"'bff_service_token_hash'; the Config API refuses to serve "
|
||||
"without an explicitly configured BFF service token."
|
||||
)
|
||||
keys = _parse_delegation_keys(data.get("webui_delegation_public_keys"))
|
||||
if not keys:
|
||||
raise PlatformConfigError(
|
||||
"config.yaml must register at least one entry in "
|
||||
"'webui_delegation_public_keys'."
|
||||
)
|
||||
|
||||
raw_db = _optional_string(data, "model_runtime_db")
|
||||
return PlatformSecurityConfig(
|
||||
bff_service_token=token,
|
||||
bff_service_token_hash=token_hash,
|
||||
webui_delegation_public_keys=keys,
|
||||
development_endpoints=_parse_development_endpoints(
|
||||
data.get("development_endpoints")
|
||||
),
|
||||
local_deployment_id=(
|
||||
_optional_string(data, "local_deployment_id") or DEFAULT_LOCAL_DEPLOYMENT_ID
|
||||
),
|
||||
model_runtime_db=None if raw_db is None else Path(raw_db).expanduser(),
|
||||
)
|
||||
@@ -16,6 +16,7 @@ import re
|
||||
import sqlite3
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
@@ -37,6 +38,10 @@ _SNAPSHOT_STATUSES = ("prepared", "bound", "expired", "aborted")
|
||||
_VERIFICATION_RESULTS = ("passed", "failed")
|
||||
_ID_PATTERN = re.compile(r"[a-z0-9][a-z0-9._-]{0,63}")
|
||||
|
||||
# Save-time validation hook: runs inside the registry write transaction,
|
||||
# after credential writes, and raises ModelRegistryError on any violation.
|
||||
SaveValidator = Callable[[sqlite3.Connection, RegistryV4], None]
|
||||
|
||||
_DDL = """
|
||||
CREATE TABLE IF NOT EXISTS registry_state (
|
||||
id INTEGER PRIMARY KEY CHECK (id = 1),
|
||||
@@ -143,10 +148,22 @@ class ModelRuntimeStore:
|
||||
SQLite itself.
|
||||
"""
|
||||
|
||||
def __init__(self, config_dir: str | Path | None = None) -> None:
|
||||
self._config_dir = (
|
||||
Path(config_dir) if config_dir is not None else (DEFAULT_CONFIG_DIR)
|
||||
)
|
||||
def __init__(
|
||||
self,
|
||||
config_dir: str | Path | None = None,
|
||||
*,
|
||||
database_path: str | Path | None = None,
|
||||
) -> None:
|
||||
if database_path is not None:
|
||||
# The platform configuration may pin the database file location
|
||||
# (``model_runtime_db`` in config.yaml, design doc 4.2).
|
||||
self._database_path = Path(database_path)
|
||||
self._config_dir = self._database_path.parent
|
||||
else:
|
||||
self._database_path = None
|
||||
self._config_dir = (
|
||||
Path(config_dir) if config_dir is not None else (DEFAULT_CONFIG_DIR)
|
||||
)
|
||||
self._lock = threading.RLock()
|
||||
self._config_dir.mkdir(parents=True, exist_ok=True)
|
||||
try:
|
||||
@@ -167,6 +184,8 @@ class ModelRuntimeStore:
|
||||
|
||||
@property
|
||||
def database_path(self) -> Path:
|
||||
if self._database_path is not None:
|
||||
return self._database_path
|
||||
return self._config_dir / DATABASE_FILENAME
|
||||
|
||||
def _connect(self) -> sqlite3.Connection:
|
||||
@@ -227,11 +246,14 @@ class ModelRuntimeStore:
|
||||
expected_revision: int,
|
||||
registry: RegistryV4,
|
||||
credential_writes: list[CredentialWrite] | None = None,
|
||||
validate: SaveValidator | None = None,
|
||||
) -> RegistryV4:
|
||||
"""Compare-and-swap the registry inside one ``BEGIN IMMEDIATE`` transaction.
|
||||
|
||||
Credential versions are written first (immutable, with incrementing
|
||||
revisions), then
|
||||
revisions), then the optional ``validate`` hook runs the section 9.2
|
||||
save-time checks against the post-write state (so the verification
|
||||
five-tuple sees the new credential revisions), then
|
||||
the registry row is replaced with ``revision + 1``. A bootstrap
|
||||
document atomically turns ``active`` the first time it carries an
|
||||
enabled model, a valid primary, and a satisfied auth reference (4.3).
|
||||
@@ -262,6 +284,8 @@ class ModelRuntimeStore:
|
||||
self._write_credential_version(
|
||||
connection, write.credential_id, write.secret_value, now=now
|
||||
)
|
||||
if validate is not None:
|
||||
validate(connection, registry)
|
||||
self._validate_defaults(registry)
|
||||
new_state = self._resolve_state(connection, current_state, registry)
|
||||
stored = registry.model_copy(
|
||||
@@ -508,6 +532,19 @@ class ModelRuntimeStore:
|
||||
).fetchone()
|
||||
return None if row is None else int(row[0])
|
||||
|
||||
def list_credential_revisions(self) -> dict[str, int]:
|
||||
"""Return every credential pointer's current revision.
|
||||
|
||||
``compute_availability`` needs the complete mapping: a missing entry
|
||||
would compare verification records against revision 0 and misjudge
|
||||
verified models as stale after a credential rotation.
|
||||
"""
|
||||
with self._lock, self._connect() as connection:
|
||||
rows = connection.execute(
|
||||
"SELECT credential_id, current_revision FROM credential_pointers"
|
||||
).fetchall()
|
||||
return {str(row[0]): int(row[1]) for row in rows}
|
||||
|
||||
def credential_status(self, credential_id: str) -> CredentialStatus:
|
||||
"""Return the redacted browser-safe status; never the plaintext."""
|
||||
_validate_credential_id(credential_id)
|
||||
@@ -878,3 +915,37 @@ class ModelRuntimeStore:
|
||||
raise
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
# --- delegation JWT anti-replay (section 7.3) -------------------------
|
||||
|
||||
def register_delegation_jti(
|
||||
self, jti: str, expires_at: int, *, now: int | None = None
|
||||
) -> bool:
|
||||
"""Atomically register one delegation JWT ID; False means replay.
|
||||
|
||||
Expired rows are purged inside the same write transaction, so a
|
||||
``jti`` whose delegation window has passed may be registered again
|
||||
while a live duplicate is rejected (``401 DELEGATION_REPLAYED``).
|
||||
"""
|
||||
if not jti:
|
||||
raise ValueError("jti must not be empty.")
|
||||
now = int(time.time()) if now is None else now
|
||||
with self._lock:
|
||||
connection = self._connect()
|
||||
try:
|
||||
connection.execute("BEGIN IMMEDIATE")
|
||||
connection.execute(
|
||||
"DELETE FROM delegation_jtis WHERE expires_at <= ?", (now,)
|
||||
)
|
||||
cursor = connection.execute(
|
||||
"INSERT OR IGNORE INTO delegation_jtis (jti, expires_at) "
|
||||
"VALUES (?, ?)",
|
||||
(jti, expires_at),
|
||||
)
|
||||
connection.commit()
|
||||
return cursor.rowcount == 1
|
||||
except BaseException:
|
||||
connection.rollback()
|
||||
raise
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
@@ -35,6 +35,9 @@ dependencies = [
|
||||
"langgraph-checkpoint-sqlite>=3.0",
|
||||
"httpx>=0.28",
|
||||
"pydantic>=2.10",
|
||||
# Delegation JWT verification for the model registry HTTP API (7.3);
|
||||
# previously only a transitive dependency, now used directly.
|
||||
"PyJWT>=2.8",
|
||||
"psutil>=6.0",
|
||||
"filelock>=3.16",
|
||||
"lazy-loader>=0.5",
|
||||
|
||||
@@ -0,0 +1,36 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Export the model registry OpenAPI contract (design doc section 11 step 3).
|
||||
|
||||
Regenerates ``EvoScientist/model_registry/openapi.json`` from the same
|
||||
Pydantic schemas the HTTP API validates against, so contract tests and the
|
||||
WebUI type generation always track the backend::
|
||||
|
||||
.venv/bin/python scripts/export_model_registry_schema.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from EvoScientist.model_registry.http_api import build_openapi_document
|
||||
|
||||
OUTPUT_PATH = (
|
||||
Path(__file__).resolve().parent.parent
|
||||
/ "EvoScientist"
|
||||
/ "model_registry"
|
||||
/ "openapi.json"
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
document = build_openapi_document()
|
||||
OUTPUT_PATH.write_text(
|
||||
json.dumps(document, indent=2, sort_keys=True) + "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
print(f"Wrote {OUTPUT_PATH}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,259 @@
|
||||
"""Delegation JWT and BFF service-token authentication tests (design doc 7.3).
|
||||
|
||||
Covers the正反例 the brief requires: missing/invalid service token,
|
||||
malformed/expired/wrongly-signed delegation JWTs, missing claims, lifetime
|
||||
over 60 seconds, wrong issuer/audience, unregistered deployments, scope
|
||||
escalation, jti replay, and the thread-claim rules for config vs
|
||||
thread-level routes. ``GET /api/model-registry`` serves as the probe route.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import time
|
||||
import uuid
|
||||
|
||||
import jwt
|
||||
import pytest
|
||||
from cryptography.hazmat.primitives import serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import ec
|
||||
from starlette.applications import Starlette
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from EvoScientist.model_registry.auth import BffAuthenticator
|
||||
from EvoScientist.model_registry.endpoint_policy import EndpointPolicy
|
||||
from EvoScientist.model_registry.http_api import (
|
||||
ApiServices,
|
||||
model_registry_routes,
|
||||
)
|
||||
from EvoScientist.model_registry.platform import DelegationPublicKey
|
||||
from EvoScientist.model_registry.resolver import ModelRegistryResolver
|
||||
from EvoScientist.model_registry.snapshots import SnapshotService
|
||||
from EvoScientist.model_registry.store import ModelRuntimeStore
|
||||
|
||||
SERVICE_TOKEN = "bff-service-token"
|
||||
|
||||
_PRIVATE_KEY = ec.generate_private_key(ec.SECP256R1())
|
||||
_PRIVATE_PEM = _PRIVATE_KEY.private_bytes(
|
||||
serialization.Encoding.PEM,
|
||||
serialization.PrivateFormat.PKCS8,
|
||||
serialization.NoEncryption(),
|
||||
)
|
||||
_PUBLIC_PEM = _PRIVATE_KEY.public_key().public_bytes(
|
||||
serialization.Encoding.PEM,
|
||||
serialization.PublicFormat.SubjectPublicKeyInfo,
|
||||
)
|
||||
|
||||
_OTHER_KEY = ec.generate_private_key(ec.SECP256R1())
|
||||
_OTHER_PRIVATE_PEM = _OTHER_KEY.private_bytes(
|
||||
serialization.Encoding.PEM,
|
||||
serialization.PrivateFormat.PKCS8,
|
||||
serialization.NoEncryption(),
|
||||
)
|
||||
|
||||
|
||||
def _services(store: ModelRuntimeStore, **auth_overrides) -> ApiServices:
|
||||
resolver = ModelRegistryResolver(store)
|
||||
auth_kwargs = {
|
||||
"service_token": SERVICE_TOKEN,
|
||||
"service_token_hash": None,
|
||||
"delegation_keys": (
|
||||
DelegationPublicKey(deployment_id="webui-1", public_key=_PUBLIC_PEM),
|
||||
),
|
||||
"jti_store": store,
|
||||
}
|
||||
auth_kwargs.update(auth_overrides)
|
||||
return ApiServices(
|
||||
store=store,
|
||||
resolver=resolver,
|
||||
snapshot_service=SnapshotService(store, resolver),
|
||||
endpoint_policy=EndpointPolicy(),
|
||||
authenticator=BffAuthenticator(**auth_kwargs),
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def store(tmp_path):
|
||||
return ModelRuntimeStore(config_dir=tmp_path)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client(store):
|
||||
services = _services(store)
|
||||
app = Starlette(routes=model_registry_routes(lambda: services))
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def _claims(**overrides):
|
||||
now = int(time.time())
|
||||
claims = {
|
||||
"iss": "WebUI",
|
||||
"aud": "EvoScientist",
|
||||
"sub": "admin-1",
|
||||
"scopes": ["model_config:read"],
|
||||
"deployment_id": "webui-1",
|
||||
"iat": now,
|
||||
"exp": now + 30,
|
||||
"jti": uuid.uuid4().hex,
|
||||
}
|
||||
claims.update(overrides)
|
||||
return claims
|
||||
|
||||
|
||||
def _encode(claims, pem=_PRIVATE_PEM):
|
||||
return jwt.encode(claims, pem, algorithm="ES256")
|
||||
|
||||
|
||||
def _headers(claims=None, *, token=SERVICE_TOKEN, pem=_PRIVATE_PEM):
|
||||
headers = {}
|
||||
if token is not None:
|
||||
headers["Authorization"] = f"Bearer {token}"
|
||||
if claims is not None:
|
||||
headers["X-Evo-Actor"] = _encode(claims, pem)
|
||||
return headers
|
||||
|
||||
|
||||
def test_valid_token_and_delegation_pass(client):
|
||||
response = client.get("/api/model-registry", headers=_headers(_claims()))
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
def test_config_route_does_not_require_thread_claim(client):
|
||||
claims = _claims()
|
||||
assert "thread_id" not in claims
|
||||
response = client.get("/api/model-registry", headers=_headers(claims))
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
def test_missing_authorization_header_is_401(client):
|
||||
response = client.get("/api/model-registry", headers={"X-Evo-Actor": "x"})
|
||||
assert response.status_code == 401
|
||||
body = response.json()
|
||||
assert body["code"] == "UNAUTHENTICATED"
|
||||
assert set(body) == {"code", "message", "details", "request_id"}
|
||||
assert body["request_id"]
|
||||
|
||||
|
||||
def test_wrong_service_token_is_401(client):
|
||||
response = client.get(
|
||||
"/api/model-registry",
|
||||
headers=_headers(_claims(), token="not-the-token"),
|
||||
)
|
||||
assert response.status_code == 401
|
||||
assert response.json()["code"] == "UNAUTHENTICATED"
|
||||
|
||||
|
||||
def test_service_token_hash_configuration_passes(store):
|
||||
services = _services(
|
||||
store,
|
||||
service_token=None,
|
||||
service_token_hash=hashlib.sha256(SERVICE_TOKEN.encode()).hexdigest(),
|
||||
)
|
||||
app = Starlette(routes=model_registry_routes(lambda: services))
|
||||
hashed_client = TestClient(app)
|
||||
assert (
|
||||
hashed_client.get(
|
||||
"/api/model-registry", headers=_headers(_claims())
|
||||
).status_code
|
||||
== 200
|
||||
)
|
||||
assert (
|
||||
hashed_client.get(
|
||||
"/api/model-registry", headers=_headers(_claims(), token="wrong")
|
||||
).status_code
|
||||
== 401
|
||||
)
|
||||
|
||||
|
||||
def test_missing_delegation_jwt_is_401(client):
|
||||
response = client.get(
|
||||
"/api/model-registry",
|
||||
headers={"Authorization": f"Bearer {SERVICE_TOKEN}"},
|
||||
)
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
def test_malformed_delegation_jwt_is_401(client):
|
||||
response = client.get(
|
||||
"/api/model-registry",
|
||||
headers=_headers(None) | {"X-Evo-Actor": "not-a-jwt"},
|
||||
)
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"missing", ["sub", "scopes", "deployment_id", "iat", "exp", "jti"]
|
||||
)
|
||||
def test_missing_required_claim_is_401(client, missing):
|
||||
claims = _claims()
|
||||
del claims[missing]
|
||||
response = client.get("/api/model-registry", headers=_headers(claims))
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
def test_expired_delegation_jwt_is_401(client):
|
||||
now = int(time.time())
|
||||
claims = _claims(iat=now - 90, exp=now - 60)
|
||||
response = client.get("/api/model-registry", headers=_headers(claims))
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
def test_lifetime_over_60_seconds_is_401(client):
|
||||
now = int(time.time())
|
||||
claims = _claims(iat=now, exp=now + 120)
|
||||
response = client.get("/api/model-registry", headers=_headers(claims))
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
def test_wrong_signature_is_401(client):
|
||||
response = client.get(
|
||||
"/api/model-registry",
|
||||
headers=_headers(_claims(), pem=_OTHER_PRIVATE_PEM),
|
||||
)
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
def test_wrong_issuer_is_401(client):
|
||||
response = client.get(
|
||||
"/api/model-registry", headers=_headers(_claims(iss="not-webui"))
|
||||
)
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
def test_wrong_audience_is_401(client):
|
||||
response = client.get(
|
||||
"/api/model-registry", headers=_headers(_claims(aud="someone-else"))
|
||||
)
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
def test_unregistered_deployment_is_401(client):
|
||||
response = client.get(
|
||||
"/api/model-registry", headers=_headers(_claims(deployment_id="webui-x"))
|
||||
)
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
def test_insufficient_scope_is_403(client):
|
||||
claims = _claims(scopes=["model:select"])
|
||||
response = client.get("/api/model-registry", headers=_headers(claims))
|
||||
assert response.status_code == 403
|
||||
body = response.json()
|
||||
assert body["code"] == "FORBIDDEN"
|
||||
assert set(body) == {"code", "message", "details", "request_id"}
|
||||
|
||||
|
||||
def test_jti_replay_is_401_delegation_replayed(client):
|
||||
headers = _headers(_claims())
|
||||
first = client.get("/api/model-registry", headers=headers)
|
||||
assert first.status_code == 200
|
||||
second = client.get("/api/model-registry", headers=headers)
|
||||
assert second.status_code == 401
|
||||
assert second.json()["code"] == "DELEGATION_REPLAYED"
|
||||
|
||||
|
||||
def test_invalid_scopes_claim_is_401(client):
|
||||
response = client.get(
|
||||
"/api/model-registry", headers=_headers(_claims(scopes="model_config:read"))
|
||||
)
|
||||
assert response.status_code == 401
|
||||
@@ -1,6 +1,11 @@
|
||||
"""Smoke test for the /api/models route mounted via langgraph.json's
|
||||
``http`` field. We test the FastAPI app directly — no need to spin up
|
||||
"""Smoke tests for the legacy admin routes mounted via langgraph.json's
|
||||
``http`` field. We test the Starlette app directly — no need to spin up
|
||||
langgraph dev.
|
||||
|
||||
``GET /api/models`` and ``POST /api/runtime-snapshots`` now belong to the
|
||||
unified model registry API (``EvoScientist.model_registry.http_api``); their
|
||||
contract tests live in ``tests/test_model_registry_http.py`` and
|
||||
``tests/test_delegation_auth.py``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -14,7 +19,6 @@ from starlette.testclient import TestClient
|
||||
from EvoScientist.config import EvoScientistConfig, load_config, save_config
|
||||
from EvoScientist.config.provider_profiles import (
|
||||
load_provider_profiles,
|
||||
replace_provider_profiles,
|
||||
)
|
||||
from EvoScientist.langgraph_dev.http import app
|
||||
from EvoScientist.llm.provider_operations import (
|
||||
@@ -32,78 +36,6 @@ def isolate_xdg_config(tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path))
|
||||
|
||||
|
||||
def test_get_models_returns_entries_and_default():
|
||||
mock_cfg = EvoScientistConfig(
|
||||
model="claude-sonnet-4-6", provider="custom-anthropic"
|
||||
)
|
||||
with patch(
|
||||
"EvoScientist.langgraph_dev.http.get_effective_config", return_value=mock_cfg
|
||||
):
|
||||
resp = client.get("/api/models")
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert "entries" in body
|
||||
assert "default" in body
|
||||
assert body["default"] == {
|
||||
"name": "claude-sonnet-4-6",
|
||||
"provider": "custom-anthropic",
|
||||
}
|
||||
assert isinstance(body["entries"], list)
|
||||
assert len(body["entries"]) > 0
|
||||
# Every entry has the three required keys
|
||||
for entry in body["entries"]:
|
||||
assert set(entry.keys()) == {"name", "model_id", "provider"}
|
||||
assert isinstance(entry["name"], str)
|
||||
assert entry["name"]
|
||||
assert isinstance(entry["model_id"], str)
|
||||
assert entry["model_id"]
|
||||
assert isinstance(entry["provider"], str)
|
||||
assert entry["provider"]
|
||||
|
||||
|
||||
def test_get_models_only_returns_enabled_catalog_entries_when_configured():
|
||||
mock_cfg = EvoScientistConfig(
|
||||
model="chat-main",
|
||||
provider="openai",
|
||||
model_catalog=[
|
||||
{
|
||||
"provider": "openai",
|
||||
"id": "chat-main",
|
||||
"name": "Chat Main",
|
||||
"model_id": "gpt-upstream",
|
||||
"enabled": True,
|
||||
},
|
||||
{
|
||||
"provider": "anthropic",
|
||||
"id": "hidden-model",
|
||||
"name": "Hidden",
|
||||
"model_id": "claude-hidden",
|
||||
"enabled": False,
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
with patch(
|
||||
"EvoScientist.langgraph_dev.http.get_effective_config", return_value=mock_cfg
|
||||
):
|
||||
body = client.get("/api/models").json()
|
||||
|
||||
assert body["entries"] == [
|
||||
{"name": "chat-main", "model_id": "gpt-upstream", "provider": "openai"}
|
||||
]
|
||||
|
||||
|
||||
def test_get_models_empty_catalog_hides_all_builtin_models():
|
||||
mock_cfg = EvoScientistConfig(model_catalog=[])
|
||||
|
||||
with patch(
|
||||
"EvoScientist.langgraph_dev.http.get_effective_config", return_value=mock_cfg
|
||||
):
|
||||
body = client.get("/api/models").json()
|
||||
|
||||
assert body["entries"] == []
|
||||
|
||||
|
||||
def test_provider_profiles_api_requires_admin_token_header(tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path))
|
||||
monkeypatch.delenv("EVOSCIENTIST_PROVIDER_ADMIN_TOKEN", raising=False)
|
||||
@@ -156,41 +88,6 @@ def test_provider_profiles_api_round_trip_redacts_secret(tmp_path, monkeypatch):
|
||||
assert get_response.json()["providers"][0]["models"][0]["id"] == "lab-model"
|
||||
assert "provider-secret" not in get_response.text
|
||||
|
||||
snapshot_response = client.post(
|
||||
"/api/runtime-snapshots",
|
||||
headers=headers,
|
||||
json={
|
||||
"snapshot_id": "run-request-a",
|
||||
"model": "lab-model",
|
||||
"provider": "lab-openai",
|
||||
},
|
||||
)
|
||||
assert snapshot_response.status_code == 200
|
||||
snapshot = snapshot_response.json()["snapshot"]
|
||||
assert snapshot["runtime"]["max_input_tokens"] == 28_672
|
||||
assert "provider-secret" not in snapshot_response.text
|
||||
|
||||
|
||||
def test_runtime_snapshot_reports_missing_zhipu_key_before_creating_a_run(
|
||||
monkeypatch,
|
||||
):
|
||||
monkeypatch.setenv("EVOSCIENTIST_PROVIDER_ADMIN_TOKEN", "admin-secret")
|
||||
monkeypatch.delenv("ZHIPU_API_KEY", raising=False)
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "unrelated-openai-key")
|
||||
|
||||
response = client.post(
|
||||
"/api/runtime-snapshots",
|
||||
headers={"X-EvoScientist-Admin-Token": "admin-secret"},
|
||||
json={
|
||||
"snapshot_id": "run-request-zhipu",
|
||||
"model": "glm-5.2",
|
||||
"provider": "glm",
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 400
|
||||
assert response.json()["error"].startswith("ZHIPU_API_KEY_NOT_CONFIGURED")
|
||||
|
||||
|
||||
def test_llm_config_api_requires_admin_token_header(monkeypatch):
|
||||
monkeypatch.delenv("EVOSCIENTIST_PROVIDER_ADMIN_TOKEN", raising=False)
|
||||
@@ -300,11 +197,6 @@ def test_llm_config_api_saves_builtin_registry_without_rewriting_legacy_secret(
|
||||
assert saved_config.default_workdir == "/tmp/research"
|
||||
assert saved_config.model_catalog is None
|
||||
|
||||
models = client.get("/api/models").json()["entries"]
|
||||
assert models == [
|
||||
{"name": "chat-main", "model_id": "gpt-upstream", "provider": "openai"}
|
||||
]
|
||||
|
||||
|
||||
def test_first_builtin_registry_save_migrates_legacy_default_model(monkeypatch):
|
||||
monkeypatch.setenv("EVOSCIENTIST_PROVIDER_ADMIN_TOKEN", "admin-secret")
|
||||
@@ -794,158 +686,6 @@ def test_default_model_api_rejects_invalid_json(monkeypatch):
|
||||
assert response.json() == {"error": "Request body must be valid JSON."}
|
||||
|
||||
|
||||
def test_entries_preserve_registry_order():
|
||||
"""The picker uses position-in-list to rank providers per short name —
|
||||
the JSON must preserve the order returned by ``list_models_by_provider``.
|
||||
|
||||
Stubs ``get_effective_config`` to keep the assertion focused on
|
||||
registry order rather than implicitly depending on the ambient
|
||||
deploy config.
|
||||
"""
|
||||
from EvoScientist.llm.models import list_models_by_provider
|
||||
|
||||
expected = [
|
||||
{"name": n, "model_id": m, "provider": p}
|
||||
for n, m, p in list_models_by_provider()
|
||||
]
|
||||
mock_cfg = EvoScientistConfig()
|
||||
with patch(
|
||||
"EvoScientist.langgraph_dev.http.get_effective_config", return_value=mock_cfg
|
||||
):
|
||||
resp = client.get("/api/models")
|
||||
assert resp.json()["entries"] == expected
|
||||
|
||||
|
||||
def test_unavailable_default_falls_back_to_first_picker_entry():
|
||||
mock_cfg = EvoScientistConfig(model="some-retired-name", provider="some-provider")
|
||||
with patch(
|
||||
"EvoScientist.langgraph_dev.http.get_effective_config", return_value=mock_cfg
|
||||
):
|
||||
resp = client.get("/api/models")
|
||||
first = resp.json()["entries"][0]
|
||||
assert resp.json()["default"] == {
|
||||
"name": first["name"],
|
||||
"provider": first["provider"],
|
||||
}
|
||||
|
||||
|
||||
def test_custom_registry_default_does_not_reuse_stale_builtin_provider():
|
||||
replace_provider_profiles(
|
||||
{
|
||||
"providers": [
|
||||
{
|
||||
"id": "open",
|
||||
"name": "Open proxy",
|
||||
"adapter": "openai",
|
||||
"base_url": "https://proxy.example.test/v1",
|
||||
"api_key": "provider-secret",
|
||||
"enabled": True,
|
||||
"models": [
|
||||
{
|
||||
"id": "gpt-5.5",
|
||||
"name": "GPT 5.5",
|
||||
"model_id": "gpt-5.5",
|
||||
"enabled": True,
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
mock_cfg = EvoScientistConfig(model="gpt-5.5", provider="openai")
|
||||
|
||||
with patch(
|
||||
"EvoScientist.langgraph_dev.http.get_effective_config", return_value=mock_cfg
|
||||
):
|
||||
body = client.get("/api/models").json()
|
||||
|
||||
assert body["entries"] == [
|
||||
{"name": "gpt-5.5", "model_id": "gpt-5.5", "provider": "open"}
|
||||
]
|
||||
assert body["default"] == {"name": "gpt-5.5", "provider": "open"}
|
||||
|
||||
|
||||
def test_ollama_models_appended_when_base_url_configured():
|
||||
"""Mirrors the TUI ``/model`` picker: when ``ollama_base_url`` is set,
|
||||
locally-pulled Ollama models are appended after the static registry
|
||||
as ``provider: "ollama"`` entries.
|
||||
"""
|
||||
mock_cfg = EvoScientistConfig(
|
||||
model="claude-sonnet-4-6",
|
||||
provider="custom-anthropic",
|
||||
ollama_base_url="http://localhost:11434",
|
||||
)
|
||||
|
||||
async def fake_discover(_base_url, *, timeout):
|
||||
return ["llama3:8b", "mistral:7b"]
|
||||
|
||||
with (
|
||||
patch(
|
||||
"EvoScientist.langgraph_dev.http.get_effective_config",
|
||||
return_value=mock_cfg,
|
||||
),
|
||||
patch(
|
||||
"EvoScientist.llm.ollama_discovery.discover_ollama_models",
|
||||
new=fake_discover,
|
||||
),
|
||||
):
|
||||
body = client.get("/api/models").json()
|
||||
|
||||
# Assert the response is the static registry followed by the discovered
|
||||
# Ollama suffix — robust to future static Ollama entries in the registry.
|
||||
from EvoScientist.llm.models import list_models_by_provider
|
||||
|
||||
static_entries = [
|
||||
{"name": n, "model_id": m, "provider": p}
|
||||
for n, m, p in list_models_by_provider()
|
||||
]
|
||||
discovered_entries = [
|
||||
{"name": "llama3:8b", "model_id": "llama3:8b", "provider": "ollama"},
|
||||
{"name": "mistral:7b", "model_id": "mistral:7b", "provider": "ollama"},
|
||||
]
|
||||
assert body["entries"][: len(static_entries)] == static_entries
|
||||
assert body["entries"][len(static_entries) :] == discovered_entries
|
||||
# TUI's "Custom Ollama model…" sentinel is a widget-specific affordance —
|
||||
# it must not appear on the HTTP surface.
|
||||
assert not any(e["model_id"] == "__custom_ollama__" for e in body["entries"])
|
||||
|
||||
|
||||
def test_ollama_discovery_skipped_when_base_url_absent():
|
||||
"""No Ollama discovery should happen when ``ollama_base_url`` is unset —
|
||||
matches the ``/model`` picker's gating. The probe function should never
|
||||
be called in that case.
|
||||
"""
|
||||
mock_cfg = EvoScientistConfig(
|
||||
model="claude-sonnet-4-6", provider="custom-anthropic"
|
||||
)
|
||||
calls: list[str | None] = []
|
||||
|
||||
async def spy_discover(base_url, *, timeout):
|
||||
calls.append(base_url)
|
||||
return []
|
||||
|
||||
with (
|
||||
patch(
|
||||
"EvoScientist.langgraph_dev.http.get_effective_config",
|
||||
return_value=mock_cfg,
|
||||
),
|
||||
patch(
|
||||
"EvoScientist.llm.ollama_discovery.discover_ollama_models",
|
||||
new=spy_discover,
|
||||
),
|
||||
):
|
||||
body = client.get("/api/models").json()
|
||||
|
||||
assert calls == []
|
||||
# Response is exactly the static registry — no Ollama additions whatsoever.
|
||||
from EvoScientist.llm.models import list_models_by_provider
|
||||
|
||||
assert body["entries"] == [
|
||||
{"name": n, "model_id": m, "provider": p}
|
||||
for n, m, p in list_models_by_provider()
|
||||
]
|
||||
|
||||
|
||||
def test_final_answer_extracts_latest_ai_text_blocks():
|
||||
async def fake_metadata(_thread_id):
|
||||
return {"updated_at": "2026-07-06T14:14:53+00:00"}
|
||||
|
||||
@@ -0,0 +1,911 @@
|
||||
"""HTTP API contract tests for the model registry (design doc 9.1-9.3, 9.5).
|
||||
|
||||
Covers the Config API (GET/PUT registry, credential rotation, selector),
|
||||
the eight section 9.2 atomic save-time checks, snapshot creation/binding/
|
||||
deletion state machines, the unified error payload, and the OpenAPI export.
|
||||
Authentication itself is covered exhaustively in test_delegation_auth.py;
|
||||
these tests mint valid delegations per request.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
|
||||
import jwt
|
||||
import pytest
|
||||
from cryptography.hazmat.primitives import serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import ec
|
||||
from starlette.applications import Starlette
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from EvoScientist.model_registry.adapters import find_adapter_spec
|
||||
from EvoScientist.model_registry.auth import BffAuthenticator
|
||||
from EvoScientist.model_registry.endpoint_policy import EndpointPolicy
|
||||
from EvoScientist.model_registry.hashing import configuration_hash
|
||||
from EvoScientist.model_registry.http_api import (
|
||||
ApiServices,
|
||||
build_openapi_document,
|
||||
model_registry_routes,
|
||||
)
|
||||
from EvoScientist.model_registry.platform import (
|
||||
DelegationPublicKey,
|
||||
PlatformConfigError,
|
||||
)
|
||||
from EvoScientist.model_registry.resolver import ModelRegistryResolver
|
||||
from EvoScientist.model_registry.schemas import (
|
||||
CredentialWrite,
|
||||
DevelopmentEndpoint,
|
||||
RegistryV4,
|
||||
)
|
||||
from EvoScientist.model_registry.snapshots import SnapshotService
|
||||
from EvoScientist.model_registry.store import ModelRuntimeStore
|
||||
|
||||
SERVICE_TOKEN = "bff-service-token"
|
||||
SECRET = "sk-live-9876abcd"
|
||||
|
||||
_PRIVATE_KEY = ec.generate_private_key(ec.SECP256R1())
|
||||
_PRIVATE_PEM = _PRIVATE_KEY.private_bytes(
|
||||
serialization.Encoding.PEM,
|
||||
serialization.PrivateFormat.PKCS8,
|
||||
serialization.NoEncryption(),
|
||||
)
|
||||
_PUBLIC_PEM = _PRIVATE_KEY.public_key().public_bytes(
|
||||
serialization.Encoding.PEM,
|
||||
serialization.PublicFormat.SubjectPublicKeyInfo,
|
||||
)
|
||||
|
||||
OPENAPI_PATH = (
|
||||
Path(__file__).resolve().parent.parent
|
||||
/ "EvoScientist"
|
||||
/ "model_registry"
|
||||
/ "openapi.json"
|
||||
)
|
||||
|
||||
|
||||
# --- registry fixtures (mirrors test_snapshots.py) ---------------------------
|
||||
|
||||
|
||||
def _model_runtime(**overrides):
|
||||
payload = {
|
||||
"limit_mode": "combined",
|
||||
"context_window_tokens": 1048576,
|
||||
"max_input_tokens": None,
|
||||
"max_output_tokens": 32768,
|
||||
"min_effective_input_tokens": 8192,
|
||||
"fixed_system_reserve_tokens": 4096,
|
||||
"fixed_tools_reserve_tokens": 8192,
|
||||
"fixed_attachments_reserve_tokens": 4096,
|
||||
"limits_status": "confirmed",
|
||||
"limits_source": "provider",
|
||||
"temperature": None,
|
||||
"top_p": None,
|
||||
"reasoning_effort": "auto",
|
||||
"declared_capabilities": {
|
||||
"tools": True,
|
||||
"vision": False,
|
||||
"structured_output": True,
|
||||
},
|
||||
}
|
||||
payload.update(overrides)
|
||||
return payload
|
||||
|
||||
|
||||
def _zhipu_provider(**overrides):
|
||||
provider = {
|
||||
"id": "zhipu-glm",
|
||||
"name": "Zhipu GLM",
|
||||
"adapter": "openai-compatible",
|
||||
"base_url": "https://open.bigmodel.cn/api/paas/v4",
|
||||
"auth": {"mode": "api_key", "credential_id": "zhipu-primary"},
|
||||
"enabled": True,
|
||||
"runtime": {
|
||||
"timeout_seconds": 120,
|
||||
"max_retries": 2,
|
||||
"default_temperature": 0.7,
|
||||
"default_top_p": 0.95,
|
||||
"default_reasoning_effort": "auto",
|
||||
},
|
||||
"models": [
|
||||
{
|
||||
"key": "glm-5.2",
|
||||
"name": "GLM-5.2",
|
||||
"upstream_model_id": "glm-5.2",
|
||||
"enabled": True,
|
||||
"runtime": _model_runtime(),
|
||||
}
|
||||
],
|
||||
}
|
||||
provider.update(overrides)
|
||||
return provider
|
||||
|
||||
|
||||
def _ollama_provider(**overrides):
|
||||
provider = {
|
||||
"id": "local-ollama",
|
||||
"name": "Local Ollama",
|
||||
"adapter": "ollama",
|
||||
"base_url": "http://localhost:11434",
|
||||
"auth": {"mode": "none", "credential_id": None},
|
||||
"enabled": True,
|
||||
"runtime": {"timeout_seconds": 120, "max_retries": 2},
|
||||
"models": [
|
||||
{
|
||||
"key": "qwen3",
|
||||
"name": "Qwen3",
|
||||
"upstream_model_id": "qwen3",
|
||||
"enabled": True,
|
||||
"runtime": _model_runtime(),
|
||||
}
|
||||
],
|
||||
}
|
||||
provider.update(overrides)
|
||||
return provider
|
||||
|
||||
|
||||
def _registry_payload():
|
||||
return {
|
||||
"version": 4,
|
||||
"revision": 1,
|
||||
"state": "bootstrap",
|
||||
"defaults": {
|
||||
"primary": {"provider_id": "zhipu-glm", "model_key": "glm-5.2"},
|
||||
"auxiliary": {"provider_id": "local-ollama", "model_key": "qwen3"},
|
||||
},
|
||||
"providers": [_zhipu_provider(), _ollama_provider()],
|
||||
}
|
||||
|
||||
|
||||
def _verify(store, registry, provider_id, model_key):
|
||||
provider = registry.find_provider(provider_id)
|
||||
model = provider.find_model(model_key)
|
||||
spec = find_adapter_spec(provider.adapter, model.upstream_model_id)
|
||||
store.record_model_verification(
|
||||
provider_id=provider_id,
|
||||
model_key=model_key,
|
||||
configuration_hash=configuration_hash(provider, model),
|
||||
credential_revision=1 if provider.auth.credential_id else 0,
|
||||
adapter_spec_revision=spec.spec_revision,
|
||||
result="passed",
|
||||
verified_capabilities={
|
||||
"tools": True,
|
||||
"vision": False,
|
||||
"structured_output": True,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
# --- app fixtures --------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def store(tmp_path):
|
||||
return ModelRuntimeStore(config_dir=tmp_path)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def services(store):
|
||||
resolver = ModelRegistryResolver(store)
|
||||
return ApiServices(
|
||||
store=store,
|
||||
resolver=resolver,
|
||||
snapshot_service=SnapshotService(store, resolver),
|
||||
endpoint_policy=EndpointPolicy(
|
||||
[
|
||||
DevelopmentEndpoint(
|
||||
id="local-ollama",
|
||||
url="http://localhost:11434",
|
||||
label="Local Ollama",
|
||||
)
|
||||
]
|
||||
),
|
||||
authenticator=BffAuthenticator(
|
||||
service_token=SERVICE_TOKEN,
|
||||
service_token_hash=None,
|
||||
delegation_keys=(
|
||||
DelegationPublicKey(deployment_id="webui-1", public_key=_PUBLIC_PEM),
|
||||
),
|
||||
jti_store=store,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client(services):
|
||||
app = Starlette(routes=model_registry_routes(lambda: services))
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def active_store(store):
|
||||
registry = store.save_registry(
|
||||
expected_revision=1,
|
||||
registry=RegistryV4.model_validate(_registry_payload()),
|
||||
credential_writes=[
|
||||
CredentialWrite(credential_id="zhipu-primary", secret_value=SECRET)
|
||||
],
|
||||
)
|
||||
assert registry.state == "active"
|
||||
_verify(store, registry, "zhipu-glm", "glm-5.2")
|
||||
_verify(store, registry, "local-ollama", "qwen3")
|
||||
return store
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def active_client(client, active_store):
|
||||
return client
|
||||
|
||||
|
||||
# --- auth helpers --------------------------------------------------------------
|
||||
|
||||
|
||||
def _headers(scopes, *, thread_id=None):
|
||||
now = int(time.time())
|
||||
claims = {
|
||||
"iss": "WebUI",
|
||||
"aud": "EvoScientist",
|
||||
"sub": "user-1",
|
||||
"scopes": list(scopes),
|
||||
"deployment_id": "webui-1",
|
||||
"iat": now,
|
||||
"exp": now + 30,
|
||||
"jti": uuid.uuid4().hex,
|
||||
}
|
||||
if thread_id is not None:
|
||||
claims["thread_id"] = thread_id
|
||||
return {
|
||||
"Authorization": f"Bearer {SERVICE_TOKEN}",
|
||||
"X-Evo-Actor": jwt.encode(claims, _PRIVATE_PEM, algorithm="ES256"),
|
||||
}
|
||||
|
||||
|
||||
def _admin_headers(**kwargs):
|
||||
return _headers(["model_config:read", "model_config:write"], **kwargs)
|
||||
|
||||
|
||||
def _run_headers(thread_id="thread-1"):
|
||||
return _headers(["run:create"], thread_id=thread_id)
|
||||
|
||||
|
||||
# --- GET /api/model-registry ----------------------------------------------------
|
||||
|
||||
|
||||
def test_get_registry_response_structure(active_client):
|
||||
response = active_client.get("/api/model-registry", headers=_admin_headers())
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert set(body) == {
|
||||
"revision",
|
||||
"registry",
|
||||
"adapter_specs",
|
||||
"credential_status",
|
||||
"model_status",
|
||||
"endpoint_policy",
|
||||
}
|
||||
assert body["revision"] == 2
|
||||
assert body["registry"]["state"] == "active"
|
||||
assert len(body["adapter_specs"]) == 6
|
||||
status = body["credential_status"]
|
||||
assert [entry["credential_id"] for entry in status] == ["zhipu-primary"]
|
||||
assert status[0]["configured"] is True
|
||||
assert status[0]["hint"] == "...abcd"
|
||||
assert status[0]["updated_at"]
|
||||
assert SECRET not in response.text
|
||||
states = {item["model_ref"]["model_key"]: item for item in body["model_status"]}
|
||||
assert states["glm-5.2"]["state"] == "enabled"
|
||||
assert states["glm-5.2"]["selectable"] is True
|
||||
assert states["qwen3"]["state"] == "enabled"
|
||||
policy = body["endpoint_policy"]
|
||||
assert policy["public_https_allowed"] is True
|
||||
assert policy["development_endpoints"] == [
|
||||
{"id": "local-ollama", "url": "http://localhost:11434", "label": "Local Ollama"}
|
||||
]
|
||||
|
||||
|
||||
def test_get_registry_marks_model_stale_after_credential_rotation(active_client):
|
||||
rotate = active_client.put(
|
||||
"/api/model-registry/credentials/zhipu-primary",
|
||||
headers=_admin_headers(),
|
||||
json={"operation": "replace", "secret_value": "rotated-secret-9999"},
|
||||
)
|
||||
assert rotate.status_code == 200
|
||||
|
||||
body = active_client.get("/api/model-registry", headers=_admin_headers()).json()
|
||||
states = {item["model_ref"]["model_key"]: item for item in body["model_status"]}
|
||||
# The full credential_revisions mapping must reach compute_availability:
|
||||
# after rotation the glm model compares against revision 2 and goes stale.
|
||||
assert states["glm-5.2"]["state"] == "verification_stale"
|
||||
assert states["glm-5.2"]["selectable"] is False
|
||||
assert states["glm-5.2"]["verification"]["status"] == "stale"
|
||||
# mode=none models are unaffected by credential rotation.
|
||||
assert states["qwen3"]["state"] == "enabled"
|
||||
|
||||
|
||||
# --- PUT /api/model-registry -----------------------------------------------------
|
||||
|
||||
|
||||
def _current_registry(active_client):
|
||||
return active_client.get("/api/model-registry", headers=_admin_headers()).json()
|
||||
|
||||
|
||||
def _put(active_client, registry, *, expected_revision=2, credential_writes=None):
|
||||
payload = {
|
||||
"expected_revision": expected_revision,
|
||||
"registry": registry,
|
||||
"credential_writes": credential_writes or [],
|
||||
}
|
||||
return active_client.put(
|
||||
"/api/model-registry", headers=_admin_headers(), json=payload
|
||||
)
|
||||
|
||||
|
||||
def test_put_registry_success_returns_full_response(active_client):
|
||||
current = _current_registry(active_client)
|
||||
registry = current["registry"]
|
||||
registry["providers"][0]["models"][0]["name"] = "GLM-5.2 Turbo"
|
||||
|
||||
response = _put(active_client, registry)
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["revision"] == current["revision"] + 1
|
||||
assert body["registry"]["providers"][0]["models"][0]["name"] == "GLM-5.2 Turbo"
|
||||
states = {item["model_ref"]["model_key"]: item for item in body["model_status"]}
|
||||
# The model name is not part of the configuration hash: still enabled.
|
||||
assert states["glm-5.2"]["state"] == "enabled"
|
||||
|
||||
|
||||
def test_put_registry_revision_conflict(active_client):
|
||||
current = _current_registry(active_client)
|
||||
response = _put(active_client, current["registry"], expected_revision=99)
|
||||
assert response.status_code == 409
|
||||
body = response.json()
|
||||
assert body["code"] == "REGISTRY_REVISION_CONFLICT"
|
||||
assert set(body) == {"code", "message", "details", "request_id"}
|
||||
|
||||
|
||||
def test_put_registry_rejects_ssrf_endpoint(active_client):
|
||||
registry = _current_registry(active_client)["registry"]
|
||||
registry["providers"][0]["base_url"] = "http://127.0.0.1:9000/v1"
|
||||
|
||||
response = _put(active_client, registry)
|
||||
assert response.status_code == 422
|
||||
body = response.json()
|
||||
assert body["code"] == "ENDPOINT_NOT_ALLOWED"
|
||||
assert body["details"][0]["path"].startswith("providers[0].base_url")
|
||||
|
||||
|
||||
def test_put_registry_rejects_unsupported_parameter(active_client):
|
||||
registry = _current_registry(active_client)["registry"]
|
||||
# The glm-5.2 contract caps temperature at 1.
|
||||
registry["providers"][0]["models"][0]["runtime"]["temperature"] = 1.5
|
||||
|
||||
response = _put(active_client, registry)
|
||||
assert response.status_code == 422
|
||||
body = response.json()
|
||||
assert body["code"] == "UNSUPPORTED_RUNTIME_PARAMETER"
|
||||
assert "providers[0].models[0]" in body["details"][0]["path"]
|
||||
|
||||
|
||||
def test_put_registry_rejects_unconfirmed_limits(active_client):
|
||||
registry = _current_registry(active_client)["registry"]
|
||||
registry["providers"][0]["models"][0]["runtime"]["limits_status"] = (
|
||||
"needs_confirmation"
|
||||
)
|
||||
|
||||
response = _put(active_client, registry)
|
||||
assert response.status_code == 422
|
||||
assert response.json()["code"] == "MODEL_LIMITS_UNCONFIRMED"
|
||||
|
||||
|
||||
def test_put_registry_rejects_enabled_model_without_verification(active_client):
|
||||
registry = _current_registry(active_client)["registry"]
|
||||
registry["providers"][0]["models"].append(
|
||||
{
|
||||
"key": "glm-5.3",
|
||||
"name": "GLM-5.3",
|
||||
"upstream_model_id": "glm-5.3",
|
||||
"enabled": True,
|
||||
"runtime": _model_runtime(),
|
||||
}
|
||||
)
|
||||
|
||||
response = _put(active_client, registry)
|
||||
assert response.status_code == 422
|
||||
assert response.json()["code"] == "MODEL_NOT_AVAILABLE"
|
||||
|
||||
|
||||
def test_put_registry_rejects_unsatisfiable_budget(active_client):
|
||||
registry = _current_registry(active_client)["registry"]
|
||||
# A disabled draft model still must satisfy the four-mode budget rule.
|
||||
registry["providers"][0]["models"].append(
|
||||
{
|
||||
"key": "glm-draft",
|
||||
"name": "GLM Draft",
|
||||
"upstream_model_id": "glm-draft",
|
||||
"enabled": False,
|
||||
"runtime": _model_runtime(min_effective_input_tokens=1048576),
|
||||
}
|
||||
)
|
||||
|
||||
response = _put(active_client, registry)
|
||||
assert response.status_code == 422
|
||||
assert response.json()["code"] == "CONTEXT_BUDGET_UNSATISFIABLE"
|
||||
|
||||
|
||||
def test_put_registry_rejects_unsupported_auth_mode(active_client):
|
||||
registry = _current_registry(active_client)["registry"]
|
||||
registry["providers"][1]["auth"] = {"mode": "api_key", "credential_id": "some-key"}
|
||||
|
||||
response = _put(active_client, registry)
|
||||
assert response.status_code == 422
|
||||
assert response.json()["code"] == "AUTH_MODE_UNSUPPORTED"
|
||||
|
||||
|
||||
def test_put_registry_rejects_defaults_to_disabled_model(active_client):
|
||||
# With a bootstrap-state payload the schema-level check does not fire, so
|
||||
# the store's rule-7 guard (enforced on every save) is what rejects it.
|
||||
registry = _current_registry(active_client)["registry"]
|
||||
registry["state"] = "bootstrap"
|
||||
registry["providers"][1]["models"][0]["enabled"] = False
|
||||
|
||||
response = _put(active_client, registry)
|
||||
assert response.status_code == 422
|
||||
assert response.json()["code"] == "MODEL_DISABLED"
|
||||
|
||||
|
||||
def test_put_registry_active_state_validates_defaults_in_schema(active_client):
|
||||
# The same rule fires earlier (Pydantic) when the payload claims active.
|
||||
registry = _current_registry(active_client)["registry"]
|
||||
registry["providers"][1]["models"][0]["enabled"] = False
|
||||
|
||||
response = _put(active_client, registry)
|
||||
assert response.status_code == 422
|
||||
assert response.json()["code"] == "VALIDATION_FAILED"
|
||||
|
||||
|
||||
def test_put_registry_atomic_rollback_with_credential_writes(active_client, store):
|
||||
registry = _current_registry(active_client)["registry"]
|
||||
registry["providers"][0]["base_url"] = "http://127.0.0.1:9000/v1"
|
||||
|
||||
response = _put(
|
||||
active_client,
|
||||
registry,
|
||||
credential_writes=[
|
||||
{
|
||||
"credential_id": "zhipu-primary",
|
||||
"operation": "replace",
|
||||
"secret_value": "should-not-persist",
|
||||
}
|
||||
],
|
||||
)
|
||||
assert response.status_code == 422
|
||||
# Neither the credential write nor the registry update may survive.
|
||||
assert store.current_credential_revision("zhipu-primary") == 1
|
||||
assert store.load_registry().revision == 2
|
||||
|
||||
|
||||
def test_put_registry_schema_validation_details(active_client):
|
||||
registry = _current_registry(active_client)["registry"]
|
||||
registry["providers"][0]["id"] = "INVALID ID!"
|
||||
|
||||
response = _put(active_client, registry)
|
||||
assert response.status_code == 422
|
||||
body = response.json()
|
||||
assert body["code"] == "VALIDATION_FAILED"
|
||||
assert body["details"]
|
||||
assert all("path" in detail and "code" in detail for detail in body["details"])
|
||||
|
||||
|
||||
def test_put_registry_bootstrap_flow_configure_test_enable(client, store):
|
||||
# Step 1: save a draft registry — models disabled, defaults null, and the
|
||||
# credential written in the same transaction (section 10 steps 5-7).
|
||||
draft = _registry_payload()
|
||||
draft["defaults"] = {"primary": None, "auxiliary": None}
|
||||
for provider in draft["providers"]:
|
||||
for model in provider["models"]:
|
||||
model["enabled"] = False
|
||||
response = client.put(
|
||||
"/api/model-registry",
|
||||
headers=_admin_headers(),
|
||||
json={
|
||||
"expected_revision": 1,
|
||||
"registry": draft,
|
||||
"credential_writes": [
|
||||
{
|
||||
"credential_id": "zhipu-primary",
|
||||
"operation": "replace",
|
||||
"secret_value": "bootstrap-secret",
|
||||
}
|
||||
],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["registry"]["state"] == "bootstrap"
|
||||
assert body["revision"] == 2
|
||||
assert "bootstrap-secret" not in response.text
|
||||
states = {item["model_ref"]["model_key"]: item for item in body["model_status"]}
|
||||
assert states["glm-5.2"]["state"] == "configured"
|
||||
|
||||
# Step 2: provider tests pass (Task 5b writes these records via the test
|
||||
# endpoint; here they are seeded directly).
|
||||
registry = store.load_registry()
|
||||
_verify(store, registry, "zhipu-glm", "glm-5.2")
|
||||
_verify(store, registry, "local-ollama", "qwen3")
|
||||
|
||||
# Step 3: enabling the verified models with defaults turns the registry
|
||||
# active in the same commit.
|
||||
enabled = body["registry"]
|
||||
enabled["defaults"] = {
|
||||
"primary": {"provider_id": "zhipu-glm", "model_key": "glm-5.2"},
|
||||
"auxiliary": {"provider_id": "local-ollama", "model_key": "qwen3"},
|
||||
}
|
||||
for provider in enabled["providers"]:
|
||||
for model in provider["models"]:
|
||||
model["enabled"] = True
|
||||
response = client.put(
|
||||
"/api/model-registry",
|
||||
headers=_admin_headers(),
|
||||
json={"expected_revision": 2, "registry": enabled, "credential_writes": []},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["registry"]["state"] == "active"
|
||||
states = {item["model_ref"]["model_key"]: item for item in body["model_status"]}
|
||||
assert states["glm-5.2"]["state"] == "enabled"
|
||||
|
||||
|
||||
def test_put_registry_malformed_json_is_400(active_client):
|
||||
response = active_client.put(
|
||||
"/api/model-registry",
|
||||
headers={**_admin_headers(), "Content-Type": "application/json"},
|
||||
content="{not json",
|
||||
)
|
||||
assert response.status_code == 400
|
||||
assert response.json()["code"] == "INVALID_REQUEST"
|
||||
|
||||
|
||||
# --- PUT /api/model-registry/credentials/{credential_id} ------------------------
|
||||
|
||||
|
||||
def test_credential_rotation_creates_new_masked_revision(active_client):
|
||||
response = active_client.put(
|
||||
"/api/model-registry/credentials/zhipu-primary",
|
||||
headers=_admin_headers(),
|
||||
json={"operation": "replace", "secret_value": "rotated-secret-9999"},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["credential_id"] == "zhipu-primary"
|
||||
assert body["revision"] == 2
|
||||
assert body["configured"] is True
|
||||
assert body["hint"] == "...9999"
|
||||
assert body["updated_at"]
|
||||
assert "rotated-secret-9999" not in response.text
|
||||
|
||||
|
||||
def test_credential_rotation_rejects_invalid_credential_id(active_client):
|
||||
response = active_client.put(
|
||||
"/api/model-registry/credentials/INVALID!!",
|
||||
headers=_admin_headers(),
|
||||
json={"operation": "replace", "secret_value": "whatever"},
|
||||
)
|
||||
assert response.status_code == 422
|
||||
|
||||
|
||||
def test_credential_rotation_requires_write_scope(active_client):
|
||||
response = active_client.put(
|
||||
"/api/model-registry/credentials/zhipu-primary",
|
||||
headers=_headers(["model_config:read"]),
|
||||
json={"operation": "replace", "secret_value": "whatever"},
|
||||
)
|
||||
assert response.status_code == 403
|
||||
|
||||
|
||||
# --- GET /api/models -------------------------------------------------------------
|
||||
|
||||
|
||||
def test_get_models_returns_only_selectable(active_client, store):
|
||||
response = active_client.get("/api/models", headers=_headers(["model:select"]))
|
||||
assert response.status_code == 200
|
||||
models = response.json()["models"]
|
||||
refs = {
|
||||
(m["model_ref"]["provider_id"], m["model_ref"]["model_key"]) for m in models
|
||||
}
|
||||
assert refs == {("zhipu-glm", "glm-5.2"), ("local-ollama", "qwen3")}
|
||||
glm = next(m for m in models if m["model_ref"]["model_key"] == "glm-5.2")
|
||||
assert glm["name"] == "GLM-5.2"
|
||||
assert glm["provider_name"] == "Zhipu GLM"
|
||||
assert glm["effective_capabilities"] == {
|
||||
"tools": True,
|
||||
"vision": False,
|
||||
"structured_output": True,
|
||||
}
|
||||
|
||||
# Disable qwen3 (and drop the auxiliary default): it leaves the selector.
|
||||
registry = store.load_registry()
|
||||
payload = registry.model_dump(mode="json")
|
||||
payload["defaults"]["auxiliary"] = None
|
||||
payload["providers"][1]["models"][0]["enabled"] = False
|
||||
store.save_registry(
|
||||
expected_revision=registry.revision,
|
||||
registry=RegistryV4.model_validate(payload),
|
||||
)
|
||||
models = active_client.get(
|
||||
"/api/models", headers=_headers(["model:select"])
|
||||
).json()["models"]
|
||||
assert [m["model_ref"]["model_key"] for m in models] == ["glm-5.2"]
|
||||
|
||||
|
||||
def test_get_models_requires_select_scope(active_client):
|
||||
response = active_client.get("/api/models", headers=_headers(["run:create"]))
|
||||
assert response.status_code == 403
|
||||
|
||||
|
||||
# --- snapshot API -----------------------------------------------------------------
|
||||
|
||||
|
||||
def _snapshot_body(**overrides):
|
||||
body = {
|
||||
"run_request_id": "req-1",
|
||||
"thread_id": "thread-1",
|
||||
"deployment_id": "webui-1",
|
||||
"model_selection_revision": 4,
|
||||
"primary": None,
|
||||
"auxiliary": None,
|
||||
}
|
||||
body.update(overrides)
|
||||
return body
|
||||
|
||||
|
||||
def test_snapshot_create_201_and_idempotent_200(active_client):
|
||||
first = active_client.post(
|
||||
"/api/runtime-snapshots", headers=_run_headers(), json=_snapshot_body()
|
||||
)
|
||||
assert first.status_code == 201
|
||||
body = first.json()
|
||||
assert body["snapshot_id"]
|
||||
assert body["registry_revision"] == 2
|
||||
assert body["primary"]["provider_id"] == "zhipu-glm"
|
||||
assert body["primary"]["adapter_spec_revision"] == 1
|
||||
assert body["primary"]["runtime"]["max_output_tokens"] == 32768
|
||||
assert body["auxiliary"]["model_key"] == "qwen3"
|
||||
assert SECRET not in first.text
|
||||
|
||||
replay = active_client.post(
|
||||
"/api/runtime-snapshots", headers=_run_headers(), json=_snapshot_body()
|
||||
)
|
||||
assert replay.status_code == 200
|
||||
assert replay.json()["snapshot_id"] == body["snapshot_id"]
|
||||
|
||||
|
||||
def test_snapshot_create_conflicting_selection_is_409(active_client):
|
||||
active_client.post(
|
||||
"/api/runtime-snapshots", headers=_run_headers(), json=_snapshot_body()
|
||||
)
|
||||
conflict = active_client.post(
|
||||
"/api/runtime-snapshots",
|
||||
headers=_run_headers(),
|
||||
json=_snapshot_body(
|
||||
primary={"provider_id": "zhipu-glm", "model_key": "glm-5.2"}
|
||||
),
|
||||
)
|
||||
assert conflict.status_code == 409
|
||||
assert conflict.json()["code"] == "RUN_REQUEST_CONFLICT"
|
||||
|
||||
|
||||
def test_snapshot_create_requires_active_registry(client):
|
||||
response = client.post(
|
||||
"/api/runtime-snapshots", headers=_run_headers(), json=_snapshot_body()
|
||||
)
|
||||
assert response.status_code == 422
|
||||
assert response.json()["code"] == "MODEL_REGISTRY_NOT_READY"
|
||||
|
||||
|
||||
def test_snapshot_create_thread_mismatch_is_403(active_client):
|
||||
response = active_client.post(
|
||||
"/api/runtime-snapshots",
|
||||
headers=_run_headers("thread-1"),
|
||||
json=_snapshot_body(thread_id="thread-2"),
|
||||
)
|
||||
assert response.status_code == 403
|
||||
assert response.json()["code"] == "FORBIDDEN"
|
||||
|
||||
|
||||
def test_snapshot_create_requires_thread_claim(active_client):
|
||||
headers = _headers(["run:create"]) # no thread_id claim
|
||||
response = active_client.post(
|
||||
"/api/runtime-snapshots", headers=headers, json=_snapshot_body()
|
||||
)
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
def test_snapshot_create_unknown_model_is_404(active_client):
|
||||
response = active_client.post(
|
||||
"/api/runtime-snapshots",
|
||||
headers=_run_headers(),
|
||||
json=_snapshot_body(primary={"provider_id": "zhipu-glm", "model_key": "nope"}),
|
||||
)
|
||||
assert response.status_code == 404
|
||||
assert response.json()["code"] == "MODEL_NOT_FOUND"
|
||||
|
||||
|
||||
def _create_snapshot(active_client, **overrides):
|
||||
response = active_client.post(
|
||||
"/api/runtime-snapshots",
|
||||
headers=_run_headers(),
|
||||
json=_snapshot_body(**overrides),
|
||||
)
|
||||
assert response.status_code == 201
|
||||
return response.json()["snapshot_id"]
|
||||
|
||||
|
||||
def test_snapshot_bind_state_machine(active_client):
|
||||
snapshot_id = _create_snapshot(active_client)
|
||||
bound = active_client.post(
|
||||
f"/api/runtime-snapshots/{snapshot_id}/bind",
|
||||
headers=_run_headers(),
|
||||
json={"langgraph_run_id": "run-1"},
|
||||
)
|
||||
assert bound.status_code == 200
|
||||
assert bound.json()["snapshot_id"] == snapshot_id
|
||||
|
||||
again = active_client.post(
|
||||
f"/api/runtime-snapshots/{snapshot_id}/bind",
|
||||
headers=_run_headers(),
|
||||
json={"langgraph_run_id": "run-1"},
|
||||
)
|
||||
assert again.status_code == 200
|
||||
|
||||
conflict = active_client.post(
|
||||
f"/api/runtime-snapshots/{snapshot_id}/bind",
|
||||
headers=_run_headers(),
|
||||
json={"langgraph_run_id": "run-2"},
|
||||
)
|
||||
assert conflict.status_code == 409
|
||||
assert conflict.json()["code"] == "SNAPSHOT_ALREADY_BOUND"
|
||||
|
||||
|
||||
def test_snapshot_bind_unknown_and_cross_thread_are_404(active_client):
|
||||
snapshot_id = _create_snapshot(active_client)
|
||||
missing = active_client.post(
|
||||
"/api/runtime-snapshots/snap-nope/bind",
|
||||
headers=_run_headers(),
|
||||
json={"langgraph_run_id": "run-1"},
|
||||
)
|
||||
assert missing.status_code == 404
|
||||
assert missing.json()["code"] == "SNAPSHOT_NOT_FOUND"
|
||||
|
||||
cross_thread = active_client.post(
|
||||
f"/api/runtime-snapshots/{snapshot_id}/bind",
|
||||
headers=_run_headers("thread-2"),
|
||||
json={"langgraph_run_id": "run-1"},
|
||||
)
|
||||
assert cross_thread.status_code == 404
|
||||
assert cross_thread.json()["code"] == "SNAPSHOT_NOT_FOUND"
|
||||
|
||||
|
||||
def test_snapshot_delete_state_machine(active_client):
|
||||
snapshot_id = _create_snapshot(active_client)
|
||||
deleted = active_client.delete(
|
||||
f"/api/runtime-snapshots/{snapshot_id}", headers=_run_headers()
|
||||
)
|
||||
assert deleted.status_code == 204
|
||||
assert deleted.content == b""
|
||||
|
||||
bound_id = _create_snapshot(active_client, run_request_id="req-2")
|
||||
active_client.post(
|
||||
f"/api/runtime-snapshots/{bound_id}/bind",
|
||||
headers=_run_headers(),
|
||||
json={"langgraph_run_id": "run-1"},
|
||||
)
|
||||
delete_bound = active_client.delete(
|
||||
f"/api/runtime-snapshots/{bound_id}", headers=_run_headers()
|
||||
)
|
||||
assert delete_bound.status_code == 409
|
||||
assert delete_bound.json()["code"] == "SNAPSHOT_ALREADY_BOUND"
|
||||
|
||||
|
||||
def test_snapshot_expired_rejects_bind_and_delete(active_client, services):
|
||||
snapshot_id = _create_snapshot(active_client)
|
||||
row = services.store.get_run_snapshot(snapshot_id)
|
||||
services.store.insert_run_snapshot(
|
||||
snapshot_id="snap-expired",
|
||||
deployment_id="webui-1",
|
||||
thread_id="thread-1",
|
||||
run_request_id="req-expired",
|
||||
selection_hash=row["selection_hash"],
|
||||
payload=row["payload"],
|
||||
expires_at=int(time.time()) - 1,
|
||||
)
|
||||
services.snapshot_service.cleanup_expired()
|
||||
|
||||
bind = active_client.post(
|
||||
"/api/runtime-snapshots/snap-expired/bind",
|
||||
headers=_run_headers(),
|
||||
json={"langgraph_run_id": "run-1"},
|
||||
)
|
||||
assert bind.status_code == 409
|
||||
assert bind.json()["code"] == "SNAPSHOT_EXPIRED"
|
||||
|
||||
delete = active_client.delete(
|
||||
"/api/runtime-snapshots/snap-expired", headers=_run_headers()
|
||||
)
|
||||
assert delete.status_code == 409
|
||||
assert delete.json()["code"] == "SNAPSHOT_EXPIRED"
|
||||
|
||||
# The expired triplet is free again: recreation returns 201.
|
||||
recreated = active_client.post(
|
||||
"/api/runtime-snapshots",
|
||||
headers=_run_headers(),
|
||||
json=_snapshot_body(run_request_id="req-expired"),
|
||||
)
|
||||
assert recreated.status_code == 201
|
||||
|
||||
|
||||
def test_snapshot_error_payload_structure(active_client):
|
||||
response = active_client.post(
|
||||
"/api/runtime-snapshots",
|
||||
headers=_run_headers(),
|
||||
json=_snapshot_body(primary={"provider_id": "zhipu-glm", "model_key": "nope"}),
|
||||
)
|
||||
body = response.json()
|
||||
assert set(body) == {"code", "message", "details", "request_id"}
|
||||
assert isinstance(body["details"], list)
|
||||
assert body["request_id"]
|
||||
|
||||
|
||||
# --- platform configuration ------------------------------------------------
|
||||
|
||||
|
||||
def test_missing_platform_config_refuses_service():
|
||||
def _unconfigured():
|
||||
raise PlatformConfigError("missing bff_service_token")
|
||||
|
||||
app = Starlette(routes=model_registry_routes(_unconfigured))
|
||||
response = TestClient(app).get("/api/model-registry")
|
||||
assert response.status_code == 500
|
||||
body = response.json()
|
||||
assert body["code"] == "PLATFORM_CONFIG_MISSING"
|
||||
assert set(body) == {"code", "message", "details", "request_id"}
|
||||
|
||||
|
||||
# --- OpenAPI export -----------------------------------------------------------------
|
||||
|
||||
|
||||
def test_openapi_export_file_matches_pydantic_schemas():
|
||||
assert OPENAPI_PATH.exists()
|
||||
exported = json.loads(OPENAPI_PATH.read_text(encoding="utf-8"))
|
||||
# The checked-in file is exactly what the export script regenerates.
|
||||
assert exported == build_openapi_document()
|
||||
|
||||
for path in (
|
||||
"/api/model-registry",
|
||||
"/api/model-registry/credentials/{credential_id}",
|
||||
"/api/models",
|
||||
"/api/runtime-snapshots",
|
||||
"/api/runtime-snapshots/{snapshot_id}/bind",
|
||||
"/api/runtime-snapshots/{snapshot_id}",
|
||||
):
|
||||
assert path in exported["paths"]
|
||||
|
||||
schema = exported["components"]["schemas"]["GetModelRegistryResponse"]
|
||||
assert set(schema["properties"]) == {
|
||||
"revision",
|
||||
"registry",
|
||||
"adapter_specs",
|
||||
"credential_status",
|
||||
"model_status",
|
||||
"endpoint_policy",
|
||||
}
|
||||
registry_schema = exported["components"]["schemas"]["RegistryV4"]
|
||||
assert "providers" in registry_schema["properties"]
|
||||
error_schema = exported["components"]["schemas"]["ErrorPayload"]
|
||||
assert set(error_schema["properties"]) == {
|
||||
"code",
|
||||
"message",
|
||||
"details",
|
||||
"request_id",
|
||||
}
|
||||
@@ -42,6 +42,7 @@ from EvoScientist.model_registry.schemas import (
|
||||
)
|
||||
|
||||
ALL_ERROR_CODES = {
|
||||
"INVALID_REQUEST": 400,
|
||||
"REGISTRY_REVISION_CONFLICT": 409,
|
||||
"THREAD_MODEL_SELECTION_CONFLICT": 409,
|
||||
"RUN_REQUEST_CONFLICT": 409,
|
||||
@@ -49,6 +50,8 @@ ALL_ERROR_CODES = {
|
||||
"SNAPSHOT_EXPIRED": 409,
|
||||
"MODEL_CONFIGURATION_CHANGED": 409,
|
||||
"DELEGATION_REPLAYED": 401,
|
||||
"UNAUTHENTICATED": 401,
|
||||
"FORBIDDEN": 403,
|
||||
"MODEL_NOT_FOUND": 404,
|
||||
"SNAPSHOT_NOT_FOUND": 404,
|
||||
"MODEL_REGISTRY_NOT_READY": 422,
|
||||
@@ -67,6 +70,8 @@ ALL_ERROR_CODES = {
|
||||
"MODEL_LIMITS_UNCONFIRMED": 422,
|
||||
"CONTEXT_BUDGET_UNSATISFIABLE": 422,
|
||||
"PROVIDER_UNREACHABLE": 422,
|
||||
"VALIDATION_FAILED": 422,
|
||||
"PLATFORM_CONFIG_MISSING": 500,
|
||||
}
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user