feat(model-registry): add provider test API and guarded verification recording
POST /api/model-registry/test (model_config:test) runs the section 9.4 flow: resolve_for_test, per-test credential resolution, build_chat_model with both safe clients, one minimal chat call, and per-capability probes (tools/structured_output/vision) whose failures only mark that capability unverified. Results upsert the model_verifications five-tuple inside a BEGIN IMMEDIATE transaction that re-checks the registry revision and configuration hash, returning 409 MODEL_CONFIGURATION_CHANGED on any concurrent change. resolve_for_test now also relaxes the passing- verification gate, which the provider test itself produces. effective_request_options reuses the redacted adapter.build_request output; the OpenAPI contract and checked-in openapi.json are updated.
This commit is contained in:
@@ -31,6 +31,8 @@ from .http_api import (
|
||||
GetModelRegistryResponse,
|
||||
GetSelectableModelsResponse,
|
||||
ModelRegistryHttpApi,
|
||||
ProviderTestRequest,
|
||||
ProviderTestResponse,
|
||||
PutModelRegistryRequest,
|
||||
SnapshotBindRequest,
|
||||
SnapshotPublicResponse,
|
||||
@@ -43,6 +45,7 @@ from .platform import (
|
||||
PlatformSecurityConfig,
|
||||
load_platform_security_config,
|
||||
)
|
||||
from .provider_test import ProviderTester, ProviderTestResult
|
||||
from .resolver import ModelRegistryResolver
|
||||
from .safe_transport import (
|
||||
AsyncSafeHttpTransport,
|
||||
@@ -127,6 +130,10 @@ __all__ = [
|
||||
"PlatformSecurityConfig",
|
||||
"ProviderConfig",
|
||||
"ProviderRuntimeConfig",
|
||||
"ProviderTestRequest",
|
||||
"ProviderTestResponse",
|
||||
"ProviderTestResult",
|
||||
"ProviderTester",
|
||||
"PutModelRegistryRequest",
|
||||
"RegistryV4",
|
||||
"ResolvedModelConfig",
|
||||
|
||||
@@ -1,10 +1,12 @@
|
||||
"""HTTP API for the unified model registry (design doc 7.3, 9.1-9.3, 9.5).
|
||||
"""HTTP API for the unified model registry (design doc 7.3, 9.1-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.
|
||||
- ``POST /api/model-registry/test`` — the section 9.4 provider test, the
|
||||
only model verification entry point, requiring ``model_config:test``.
|
||||
- ``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
|
||||
@@ -30,7 +32,14 @@ from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Literal
|
||||
|
||||
from pydantic import BaseModel, Field, PositiveInt, TypeAdapter, ValidationError
|
||||
from pydantic import (
|
||||
BaseModel,
|
||||
Field,
|
||||
NonNegativeInt,
|
||||
PositiveInt,
|
||||
TypeAdapter,
|
||||
ValidationError,
|
||||
)
|
||||
from pydantic.json_schema import models_json_schema
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import JSONResponse, Response
|
||||
@@ -58,6 +67,7 @@ from .platform import (
|
||||
PlatformSecurityConfig,
|
||||
load_platform_security_config,
|
||||
)
|
||||
from .provider_test import ProviderTester
|
||||
from .resolver import ModelRegistryResolver
|
||||
from .schemas import (
|
||||
AdapterParameterSpec,
|
||||
@@ -112,6 +122,25 @@ class CredentialReplaceRequest(BaseModel):
|
||||
secret_value: NonEmptyString
|
||||
|
||||
|
||||
class ProviderTestRequest(BaseModel):
|
||||
"""The section 9.4 provider test request."""
|
||||
|
||||
expected_registry_revision: PositiveInt
|
||||
model_ref: ModelRef
|
||||
|
||||
|
||||
class ProviderTestResponse(BaseModel):
|
||||
"""The section 9.4 provider test outcome; never carries secrets."""
|
||||
|
||||
ok: bool
|
||||
registry_revision: PositiveInt
|
||||
model_ref: ModelRef
|
||||
adapter_spec_revision: PositiveInt
|
||||
effective_request_options: dict[str, Any]
|
||||
model_status: ModelAvailability
|
||||
latency_ms: NonNegativeInt
|
||||
|
||||
|
||||
class CredentialWriteResponse(BaseModel):
|
||||
"""The section 9.3 credential rotation response; never the plaintext."""
|
||||
|
||||
@@ -183,6 +212,9 @@ class ApiServices:
|
||||
snapshot_service: SnapshotService
|
||||
endpoint_policy: EndpointPolicy
|
||||
authenticator: BffAuthenticator
|
||||
# Optional so pre-Task-5b service bundles keep working; the test route
|
||||
# lazily builds a default tester when the field is unset.
|
||||
provider_tester: ProviderTester | None = None
|
||||
|
||||
|
||||
def build_services(platform: PlatformSecurityConfig) -> ApiServices:
|
||||
@@ -206,6 +238,7 @@ def build_services(platform: PlatformSecurityConfig) -> ApiServices:
|
||||
snapshot_service=SnapshotService(store, resolver),
|
||||
endpoint_policy=policy,
|
||||
authenticator=authenticator,
|
||||
provider_tester=ProviderTester(store, resolver, policy),
|
||||
)
|
||||
|
||||
|
||||
@@ -407,6 +440,11 @@ class ModelRegistryHttpApi:
|
||||
self.put_credential,
|
||||
methods=["PUT"],
|
||||
),
|
||||
Route(
|
||||
"/api/model-registry/test",
|
||||
self.test_provider,
|
||||
methods=["POST"],
|
||||
),
|
||||
Route("/api/models", self.get_selectable_models, methods=["GET"]),
|
||||
Route("/api/runtime-snapshots", self.create_snapshot, methods=["POST"]),
|
||||
Route(
|
||||
@@ -503,6 +541,20 @@ class ModelRegistryHttpApi:
|
||||
return _validation_error_response(exc, request_id)
|
||||
return JSONResponse(response.model_dump(mode="json"))
|
||||
|
||||
async def test_provider(self, request: Request) -> Response:
|
||||
request_id = uuid.uuid4().hex
|
||||
try:
|
||||
services, _actor = self._authenticate(
|
||||
request, required_scope="model_config:test", require_thread_id=False
|
||||
)
|
||||
body = await self._body(request)
|
||||
response = await asyncio.to_thread(self._run_provider_test, 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 get_selectable_models(self, request: Request) -> Response:
|
||||
request_id = uuid.uuid4().hex
|
||||
try:
|
||||
@@ -640,6 +692,26 @@ class ModelRegistryHttpApi:
|
||||
updated_at=status.updated_at,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _run_provider_test(services: ApiServices, body: Any) -> ProviderTestResponse:
|
||||
request = ProviderTestRequest.model_validate(body)
|
||||
tester = services.provider_tester or ProviderTester(
|
||||
services.store, services.resolver, services.endpoint_policy
|
||||
)
|
||||
result = tester.run(
|
||||
expected_registry_revision=request.expected_registry_revision,
|
||||
model_ref=request.model_ref,
|
||||
)
|
||||
return ProviderTestResponse(
|
||||
ok=result.ok,
|
||||
registry_revision=result.registry_revision,
|
||||
model_ref=result.model_ref,
|
||||
adapter_spec_revision=result.adapter_spec_revision,
|
||||
effective_request_options=result.effective_request_options,
|
||||
model_status=result.model_status,
|
||||
latency_ms=result.latency_ms,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _selectable_models(services: ApiServices) -> GetSelectableModelsResponse:
|
||||
store = services.store
|
||||
@@ -733,6 +805,8 @@ _OPENAPI_MODELS = (
|
||||
PutModelRegistryRequest,
|
||||
CredentialReplaceRequest,
|
||||
CredentialWriteResponse,
|
||||
ProviderTestRequest,
|
||||
ProviderTestResponse,
|
||||
GetSelectableModelsResponse,
|
||||
SnapshotCreateRequest,
|
||||
SnapshotBindRequest,
|
||||
@@ -830,6 +904,24 @@ def build_openapi_document() -> dict[str, Any]:
|
||||
},
|
||||
}
|
||||
},
|
||||
"/api/model-registry/test": {
|
||||
"post": {
|
||||
"operationId": "testProviderModel",
|
||||
"security": auth_security,
|
||||
"requestBody": {
|
||||
"required": True,
|
||||
"content": {
|
||||
"application/json": {"schema": _schema_ref(ProviderTestRequest)}
|
||||
},
|
||||
},
|
||||
"responses": {
|
||||
**_json_response(
|
||||
"200", ProviderTestResponse, "The provider test outcome."
|
||||
),
|
||||
**_error_responses(400, 401, 403, 404, 409, 422, 500),
|
||||
},
|
||||
}
|
||||
},
|
||||
"/api/models": {
|
||||
"get": {
|
||||
"operationId": "getSelectableModels",
|
||||
|
||||
@@ -941,6 +941,71 @@
|
||||
"title": "ProviderRuntimeConfig",
|
||||
"type": "object"
|
||||
},
|
||||
"ProviderTestRequest": {
|
||||
"description": "The section 9.4 provider test request.",
|
||||
"properties": {
|
||||
"expected_registry_revision": {
|
||||
"exclusiveMinimum": 0,
|
||||
"title": "Expected Registry Revision",
|
||||
"type": "integer"
|
||||
},
|
||||
"model_ref": {
|
||||
"$ref": "#/components/schemas/ModelRef"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"expected_registry_revision",
|
||||
"model_ref"
|
||||
],
|
||||
"title": "ProviderTestRequest",
|
||||
"type": "object"
|
||||
},
|
||||
"ProviderTestResponse": {
|
||||
"description": "The section 9.4 provider test outcome; never carries secrets.",
|
||||
"properties": {
|
||||
"adapter_spec_revision": {
|
||||
"exclusiveMinimum": 0,
|
||||
"title": "Adapter Spec Revision",
|
||||
"type": "integer"
|
||||
},
|
||||
"effective_request_options": {
|
||||
"additionalProperties": true,
|
||||
"title": "Effective Request Options",
|
||||
"type": "object"
|
||||
},
|
||||
"latency_ms": {
|
||||
"minimum": 0,
|
||||
"title": "Latency Ms",
|
||||
"type": "integer"
|
||||
},
|
||||
"model_ref": {
|
||||
"$ref": "#/components/schemas/ModelRef"
|
||||
},
|
||||
"model_status": {
|
||||
"$ref": "#/components/schemas/ModelAvailability"
|
||||
},
|
||||
"ok": {
|
||||
"title": "Ok",
|
||||
"type": "boolean"
|
||||
},
|
||||
"registry_revision": {
|
||||
"exclusiveMinimum": 0,
|
||||
"title": "Registry Revision",
|
||||
"type": "integer"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"ok",
|
||||
"registry_revision",
|
||||
"model_ref",
|
||||
"adapter_spec_revision",
|
||||
"effective_request_options",
|
||||
"model_status",
|
||||
"latency_ms"
|
||||
],
|
||||
"title": "ProviderTestResponse",
|
||||
"type": "object"
|
||||
},
|
||||
"PutModelRegistryRequest": {
|
||||
"description": "The section 9.2 registry save request.",
|
||||
"properties": {
|
||||
@@ -1544,6 +1609,109 @@
|
||||
]
|
||||
}
|
||||
},
|
||||
"/api/model-registry/test": {
|
||||
"post": {
|
||||
"operationId": "testProviderModel",
|
||||
"requestBody": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/ProviderTestRequest"
|
||||
}
|
||||
}
|
||||
},
|
||||
"required": true
|
||||
},
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/ProviderTestResponse"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "The provider test outcome."
|
||||
},
|
||||
"400": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/ErrorPayload"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Unified error payload (section 9.5)."
|
||||
},
|
||||
"401": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/ErrorPayload"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Unified error payload (section 9.5)."
|
||||
},
|
||||
"403": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/ErrorPayload"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Unified error payload (section 9.5)."
|
||||
},
|
||||
"404": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/ErrorPayload"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Unified error payload (section 9.5)."
|
||||
},
|
||||
"409": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/ErrorPayload"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Unified error payload (section 9.5)."
|
||||
},
|
||||
"422": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/ErrorPayload"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Unified error payload (section 9.5)."
|
||||
},
|
||||
"500": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/ErrorPayload"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Unified error payload (section 9.5)."
|
||||
}
|
||||
},
|
||||
"security": [
|
||||
{
|
||||
"BffServiceToken": [],
|
||||
"DelegationJwt": []
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
"/api/models": {
|
||||
"get": {
|
||||
"operationId": "getSelectableModels",
|
||||
|
||||
@@ -0,0 +1,460 @@
|
||||
"""Provider test execution: the only model verification entry point (9.4).
|
||||
|
||||
``ProviderTester.run`` implements the section 9.4 flow end to end:
|
||||
|
||||
1. The ``ModelRef`` must already be saved in the registry (browser free
|
||||
drafts are rejected with ``MODEL_NOT_FOUND``) and its base URL must still
|
||||
pass the EndpointPolicy — the save-time check is re-applied because the
|
||||
test is a network operation (SSRF defense, section 4.3).
|
||||
2. ``resolve_for_test`` freezes a temporary ``ResolvedModelConfig``; the
|
||||
credential secret is resolved from the store against the frozen
|
||||
``auth_ref`` on every test — never from a process cache (section 5.2).
|
||||
3. ``build_chat_model`` constructs the LangChain model with both safe HTTP
|
||||
clients, so all egress shares the SafeHttpTransport defenses; adapters
|
||||
whose contract keeps retries out of the request (empty ``target_name``,
|
||||
e.g. ollama) get them through the safe client builder instead.
|
||||
4. One minimal real chat call runs, followed by one minimal probe per
|
||||
declared capability (section 6.2). A failing probe only marks that
|
||||
capability unverified; it never blocks the others.
|
||||
5. The ``model_verifications`` five-tuple record is upserted inside a
|
||||
transaction that re-checks ``expected_registry_revision`` and the model's
|
||||
``configuration_hash``; a concurrent change aborts with
|
||||
``409 MODEL_CONFIGURATION_CHANGED`` and records nothing (section 9.4).
|
||||
|
||||
Errors are classified into the stable section 9.5 codes; messages are
|
||||
static strings and never carry secrets, stack traces, or raw provider
|
||||
responses. ``effective_request_options`` reuses ``adapter.build_request``
|
||||
output with the connection and credential fields stripped, proving the
|
||||
configuration, the snapshot, and the actual request agree (section 6.4).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from collections.abc import Callable, Iterator
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from langchain_core.language_models.chat_models import BaseChatModel
|
||||
from langchain_core.messages import HumanMessage
|
||||
from langchain_core.tools import tool
|
||||
|
||||
from .adapters import Adapter, BuiltRequest, get_adapter
|
||||
from .endpoint_policy import EndpointPolicy
|
||||
from .errors import (
|
||||
CREDENTIAL_REJECTED,
|
||||
MODEL_CONFIGURATION_CHANGED,
|
||||
MODEL_NOT_AVAILABLE,
|
||||
MODEL_NOT_FOUND,
|
||||
PROVIDER_UNREACHABLE,
|
||||
ModelRegistryError,
|
||||
)
|
||||
from .factory import build_chat_model
|
||||
from .hashing import configuration_hash
|
||||
from .resolver import ModelRegistryResolver
|
||||
from .safe_transport import build_safe_async_http_client, build_safe_http_client
|
||||
from .schemas import (
|
||||
Capabilities,
|
||||
ModelAvailability,
|
||||
ModelRef,
|
||||
ResolvedModelConfig,
|
||||
)
|
||||
from .store import ModelRuntimeStore
|
||||
|
||||
# Provider tests must fail fast; the resolved contract timeout (up to 600s)
|
||||
# is capped at this configurable, deliberately short default (section 9.4).
|
||||
DEFAULT_TEST_TIMEOUT_SECONDS = 30
|
||||
|
||||
SyncClientBuilder = Callable[[float, int], httpx.Client]
|
||||
AsyncClientBuilder = Callable[[float, int], httpx.AsyncClient]
|
||||
|
||||
_CAPABILITY_NAMES = ("tools", "vision", "structured_output")
|
||||
|
||||
# A 1x1 transparent PNG for the vision capability probe (section 6.2).
|
||||
_PIXEL_PNG_B64 = (
|
||||
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8"
|
||||
"z8BQDwAEhQGAhKmMIQAAAABJRU5ErkJggg=="
|
||||
)
|
||||
|
||||
_STRUCTURED_PROBE_SCHEMA = {
|
||||
"title": "ProviderTestProbe",
|
||||
"description": "Trivial structured-output capability probe.",
|
||||
"type": "object",
|
||||
"properties": {"answer": {"type": "string"}},
|
||||
"required": ["answer"],
|
||||
}
|
||||
|
||||
|
||||
@tool
|
||||
def _ping_tool(text: str) -> str:
|
||||
"""Echo the input text; exists only as a trivial tools probe."""
|
||||
return text
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ProviderTestResult:
|
||||
"""The successful section 9.4 test outcome; carries no secrets."""
|
||||
|
||||
ok: bool
|
||||
registry_revision: int
|
||||
model_ref: ModelRef
|
||||
adapter_spec_revision: int
|
||||
effective_request_options: dict[str, Any]
|
||||
model_status: ModelAvailability
|
||||
latency_ms: int
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ExecutionOutcome:
|
||||
"""The network phase result: per-capability verdicts and timing."""
|
||||
|
||||
verified_capabilities: dict[str, bool]
|
||||
latency_ms: int
|
||||
|
||||
|
||||
def _walk_causes(exc: BaseException) -> Iterator[BaseException]:
|
||||
"""Yield ``exc`` and its ``__cause__``/``__context__`` chain, cycle-safe."""
|
||||
seen: set[int] = set()
|
||||
current: BaseException | None = exc
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
yield current
|
||||
current = current.__cause__ or current.__context__
|
||||
|
||||
|
||||
def _status_code(exc: BaseException) -> int | None:
|
||||
"""Find an HTTP status anywhere in the exception chain.
|
||||
|
||||
Both the OpenAI and Anthropic SDKs expose ``status_code`` on their
|
||||
status errors; plain httpx surfaces it via ``response``.
|
||||
"""
|
||||
for err in _walk_causes(exc):
|
||||
status = getattr(err, "status_code", None)
|
||||
if isinstance(status, int):
|
||||
return status
|
||||
response = getattr(err, "response", None)
|
||||
if isinstance(response, httpx.Response):
|
||||
return response.status_code
|
||||
return None
|
||||
|
||||
|
||||
def _classify_provider_error(exc: BaseException) -> ModelRegistryError:
|
||||
"""Map any provider-call failure onto the stable section 9.5 codes.
|
||||
|
||||
SDKs wrap transport failures (including the EndpointPolicy rejection
|
||||
raised by SafeHttpTransport) into their own connection errors, so the
|
||||
``ModelRegistryError`` search must walk the whole cause chain.
|
||||
"""
|
||||
for err in _walk_causes(exc):
|
||||
if isinstance(err, ModelRegistryError):
|
||||
return err
|
||||
status = _status_code(exc)
|
||||
if status in (401, 403):
|
||||
return ModelRegistryError(
|
||||
CREDENTIAL_REJECTED, "The provider rejected the configured credential."
|
||||
)
|
||||
if status == 404:
|
||||
return ModelRegistryError(
|
||||
MODEL_NOT_FOUND, "The provider reports that the model does not exist."
|
||||
)
|
||||
for err in _walk_causes(exc):
|
||||
if isinstance(err, httpx.TransportError):
|
||||
return ModelRegistryError(
|
||||
PROVIDER_UNREACHABLE, "The provider endpoint could not be reached."
|
||||
)
|
||||
if status is not None and status >= 500:
|
||||
return ModelRegistryError(
|
||||
PROVIDER_UNREACHABLE, "The provider endpoint could not be reached."
|
||||
)
|
||||
return ModelRegistryError(MODEL_NOT_AVAILABLE, "The provider test request failed.")
|
||||
|
||||
|
||||
def _probe_tools(chat_model: BaseChatModel) -> None:
|
||||
"""Bind a trivial tool; the request must carry it and not error (6.2)."""
|
||||
bound = chat_model.bind_tools([_ping_tool])
|
||||
messages = [HumanMessage(content="Call the ping tool with text 'probe'.")]
|
||||
request_payload = getattr(bound, "_get_request_payload", None)
|
||||
if callable(request_payload):
|
||||
# The bound kwargs hold the formatted tools; ``_get_request_payload``
|
||||
# only folds them into the payload when they are passed explicitly.
|
||||
bound_kwargs = getattr(bound, "kwargs", None)
|
||||
payload = request_payload(messages, **(bound_kwargs or {}))
|
||||
if "tools" not in payload:
|
||||
raise RuntimeError("The bound request payload does not contain tools.")
|
||||
bound.invoke(messages)
|
||||
|
||||
|
||||
def _probe_structured_output(chat_model: BaseChatModel) -> None:
|
||||
"""Request a JSON-schema structured response (section 6.2)."""
|
||||
try:
|
||||
structured = chat_model.with_structured_output(
|
||||
_STRUCTURED_PROBE_SCHEMA, method="json_schema"
|
||||
)
|
||||
except (TypeError, ValueError):
|
||||
# Adapters without a json_schema method fall back to their native
|
||||
# structured-output mechanism.
|
||||
structured = chat_model.with_structured_output(_STRUCTURED_PROBE_SCHEMA)
|
||||
structured.invoke([HumanMessage(content='Answer with {"answer": "ok"}.')])
|
||||
|
||||
|
||||
def _probe_vision(chat_model: BaseChatModel) -> None:
|
||||
"""Send a 1x1 test image (section 6.2)."""
|
||||
chat_model.invoke(
|
||||
[
|
||||
HumanMessage(
|
||||
content=[
|
||||
{"type": "text", "text": "Describe this image in one word."},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": f"data:image/png;base64,{_PIXEL_PNG_B64}"},
|
||||
},
|
||||
]
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
_CAPABILITY_PROBES = {
|
||||
"tools": _probe_tools,
|
||||
"structured_output": _probe_structured_output,
|
||||
"vision": _probe_vision,
|
||||
}
|
||||
|
||||
|
||||
def _redacted_options(adapter: Adapter, built: BuiltRequest) -> dict[str, Any]:
|
||||
"""The effective request parameters with connection/secret fields removed.
|
||||
|
||||
Reuses the exact ``adapter.build_request`` output the model was built
|
||||
from, minus the connection identity (``model``/``base_url``) and every
|
||||
auth target the contract declares (``api_key``/``default_headers``).
|
||||
"""
|
||||
options = dict(built.client_options)
|
||||
connection = adapter.spec.connection
|
||||
if connection is not None:
|
||||
options.pop(connection.model_field, None)
|
||||
options.pop(connection.base_url_field, None)
|
||||
for auth_spec in adapter.spec.auth_specs.values():
|
||||
if auth_spec.target == "client_option" and auth_spec.target_name:
|
||||
options.pop(auth_spec.target_name, None)
|
||||
elif auth_spec.target == "request_header":
|
||||
options.pop("default_headers", None)
|
||||
options.update(built.request_options)
|
||||
return options
|
||||
|
||||
|
||||
class ProviderTester:
|
||||
"""Runs provider tests and records their verdicts (section 9.4)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
store: ModelRuntimeStore,
|
||||
resolver: ModelRegistryResolver,
|
||||
endpoint_policy: EndpointPolicy,
|
||||
*,
|
||||
timeout_seconds: int = DEFAULT_TEST_TIMEOUT_SECONDS,
|
||||
sync_client_builder: SyncClientBuilder | None = None,
|
||||
async_client_builder: AsyncClientBuilder | None = None,
|
||||
) -> None:
|
||||
if timeout_seconds <= 0:
|
||||
raise ValueError("timeout_seconds must be positive.")
|
||||
self._store = store
|
||||
self._resolver = resolver
|
||||
self._endpoint_policy = endpoint_policy
|
||||
self._timeout_seconds = timeout_seconds
|
||||
self._sync_client_builder = sync_client_builder or (
|
||||
lambda timeout, retries: build_safe_http_client(
|
||||
endpoint_policy, timeout=timeout, retries=retries
|
||||
)
|
||||
)
|
||||
self._async_client_builder = async_client_builder or (
|
||||
lambda timeout, retries: build_safe_async_http_client(
|
||||
endpoint_policy, timeout=timeout, retries=retries
|
||||
)
|
||||
)
|
||||
|
||||
def run(
|
||||
self, *, expected_registry_revision: int, model_ref: ModelRef
|
||||
) -> ProviderTestResult:
|
||||
"""Test one saved model and record the verdict.
|
||||
|
||||
Provider-call failures upsert a ``failed`` record (when the registry
|
||||
is unchanged) and then raise the classified section 9.5 error;
|
||||
precondition failures raise before any network I/O or record write.
|
||||
"""
|
||||
registry = self._store.load_registry()
|
||||
provider = registry.find_provider(model_ref.provider_id)
|
||||
model = provider.find_model(model_ref.model_key) if provider else None
|
||||
if provider is None or model is None:
|
||||
raise ModelRegistryError(
|
||||
MODEL_NOT_FOUND,
|
||||
"The model reference does not exist in the saved registry; "
|
||||
"only saved models can be tested.",
|
||||
details=[{"path": "model_ref", "code": MODEL_NOT_FOUND}],
|
||||
)
|
||||
# Cheap pre-check: a stale caller view fails before any network I/O.
|
||||
# The transactional re-check at record time stays authoritative.
|
||||
if registry.revision != expected_registry_revision:
|
||||
raise ModelRegistryError(
|
||||
MODEL_CONFIGURATION_CHANGED,
|
||||
"The registry changed since the test was requested; reload and retry.",
|
||||
details=[
|
||||
{
|
||||
"path": "expected_registry_revision",
|
||||
"code": MODEL_CONFIGURATION_CHANGED,
|
||||
}
|
||||
],
|
||||
)
|
||||
# Re-apply the save-time URL check: the test is itself an SSRF
|
||||
# channel if it trusted a stored URL blindly (section 4.3).
|
||||
self._endpoint_policy.validate_base_url(provider.base_url)
|
||||
resolved = self._resolver.resolve_for_test(model_ref)
|
||||
credential = self._resolve_credential(resolved)
|
||||
adapter = get_adapter(
|
||||
resolved.adapter_id,
|
||||
resolved.upstream_model_id,
|
||||
spec_revision=resolved.adapter_spec_revision,
|
||||
)
|
||||
built = adapter.build_request(resolved, credential=credential)
|
||||
effective_request_options = _redacted_options(adapter, built)
|
||||
config_hash = configuration_hash(provider, model)
|
||||
|
||||
try:
|
||||
outcome = self._execute(
|
||||
resolved, adapter, credential, model.runtime.declared_capabilities
|
||||
)
|
||||
except ModelRegistryError as exc:
|
||||
self._record(
|
||||
expected_registry_revision,
|
||||
resolved,
|
||||
config_hash,
|
||||
result="failed",
|
||||
verified_capabilities=dict.fromkeys(_CAPABILITY_NAMES, False),
|
||||
error_code=exc.code,
|
||||
)
|
||||
raise
|
||||
self._record(
|
||||
expected_registry_revision,
|
||||
resolved,
|
||||
config_hash,
|
||||
result="passed",
|
||||
verified_capabilities=outcome.verified_capabilities,
|
||||
error_code=None,
|
||||
)
|
||||
return ProviderTestResult(
|
||||
ok=True,
|
||||
registry_revision=expected_registry_revision,
|
||||
model_ref=model_ref,
|
||||
adapter_spec_revision=resolved.adapter_spec_revision,
|
||||
effective_request_options=effective_request_options,
|
||||
model_status=self._model_status(model_ref),
|
||||
latency_ms=outcome.latency_ms,
|
||||
)
|
||||
|
||||
def _resolve_credential(self, resolved: ResolvedModelConfig) -> str | None:
|
||||
"""Resolve the frozen credential version per test; never cached (5.2)."""
|
||||
auth_ref = resolved.auth_ref
|
||||
if auth_ref.mode == "none":
|
||||
return None
|
||||
if auth_ref.credential_id is None or auth_ref.credential_revision is None:
|
||||
return None
|
||||
return self._store.resolve_credential(
|
||||
auth_ref.credential_id, auth_ref.credential_revision
|
||||
)
|
||||
|
||||
def _execute(
|
||||
self,
|
||||
resolved: ResolvedModelConfig,
|
||||
adapter: Adapter,
|
||||
credential: str | None,
|
||||
declared: Capabilities,
|
||||
) -> _ExecutionOutcome:
|
||||
"""Build the model and run the minimal call plus capability probes."""
|
||||
# Contracts that validate max_retries without sending it (empty
|
||||
# target_name, e.g. ollama) have them enforced by the safe transport.
|
||||
retries = (
|
||||
resolved.client_options.max_retries
|
||||
if adapter.spec.parameters["max_retries"].target_name == ""
|
||||
else 0
|
||||
)
|
||||
timeout = float(
|
||||
min(resolved.client_options.timeout_seconds, self._timeout_seconds)
|
||||
)
|
||||
sync_client = self._sync_client_builder(timeout, retries)
|
||||
async_client = self._async_client_builder(timeout, retries)
|
||||
try:
|
||||
chat_model = build_chat_model(
|
||||
resolved, sync_client, async_client, credential=credential
|
||||
)
|
||||
started = time.monotonic()
|
||||
self._minimal_chat(chat_model)
|
||||
verified = dict.fromkeys(_CAPABILITY_NAMES, False)
|
||||
for name in _CAPABILITY_NAMES:
|
||||
if getattr(declared, name):
|
||||
verified[name] = self._probe(name, chat_model)
|
||||
latency_ms = int((time.monotonic() - started) * 1000)
|
||||
return _ExecutionOutcome(
|
||||
verified_capabilities=verified, latency_ms=latency_ms
|
||||
)
|
||||
finally:
|
||||
sync_client.close()
|
||||
# The tester runs in a worker thread without a running loop.
|
||||
asyncio.run(async_client.aclose())
|
||||
|
||||
@staticmethod
|
||||
def _minimal_chat(chat_model: BaseChatModel) -> None:
|
||||
"""One low-cost chat request; failures classify into section 9.5 codes."""
|
||||
try:
|
||||
chat_model.invoke([HumanMessage(content="Reply with exactly: ok")])
|
||||
except ModelRegistryError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
raise _classify_provider_error(exc) from exc
|
||||
|
||||
@staticmethod
|
||||
def _probe(name: str, chat_model: BaseChatModel) -> bool:
|
||||
"""Run one capability probe; any failure marks only that capability."""
|
||||
try:
|
||||
_CAPABILITY_PROBES[name](chat_model)
|
||||
except Exception:
|
||||
return False
|
||||
return True
|
||||
|
||||
def _record(
|
||||
self,
|
||||
expected_registry_revision: int,
|
||||
resolved: ResolvedModelConfig,
|
||||
config_hash: str,
|
||||
*,
|
||||
result: str,
|
||||
verified_capabilities: dict[str, bool],
|
||||
error_code: str | None,
|
||||
) -> None:
|
||||
"""The guarded five-tuple upsert; raises 409 when the config changed."""
|
||||
self._store.record_model_verification_if_current(
|
||||
expected_registry_revision=expected_registry_revision,
|
||||
provider_id=resolved.model_ref.provider_id,
|
||||
model_key=resolved.model_ref.model_key,
|
||||
configuration_hash=config_hash,
|
||||
credential_revision=resolved.auth_ref.credential_revision or 0,
|
||||
adapter_spec_revision=resolved.adapter_spec_revision,
|
||||
result=result,
|
||||
verified_capabilities=verified_capabilities,
|
||||
error_code=error_code,
|
||||
)
|
||||
|
||||
def _model_status(self, model_ref: ModelRef) -> ModelAvailability:
|
||||
"""Recompute the backend ModelAvailability after the record write."""
|
||||
registry = self._store.load_registry()
|
||||
availability = self._resolver.compute_availability(
|
||||
registry,
|
||||
self._store.list_model_verifications(),
|
||||
credential_revisions=self._store.list_credential_revisions(),
|
||||
)
|
||||
for item in availability:
|
||||
if item.model_ref == model_ref:
|
||||
return item
|
||||
raise ModelRegistryError( # pragma: no cover - judged from same registry
|
||||
MODEL_NOT_FOUND, "The model reference does not exist in the registry."
|
||||
)
|
||||
@@ -107,10 +107,18 @@ class ModelRegistryResolver:
|
||||
def resolve_for_test(self, model_ref: ModelRef) -> ResolvedModelConfig:
|
||||
"""Resolve for a provider test (section 9.4).
|
||||
|
||||
Only the "model must be enabled" visibility check is relaxed; auth,
|
||||
The run-time visibility gates — the model must be enabled and a
|
||||
passing verification must already exist — are relaxed, because the
|
||||
provider test is exactly what produces that verification. Auth,
|
||||
parameter, capability, limits, and budget validation all still run.
|
||||
"""
|
||||
return self._resolve(model_ref, "primary", registry=None, require_enabled=False)
|
||||
return self._resolve(
|
||||
model_ref,
|
||||
"primary",
|
||||
registry=None,
|
||||
require_enabled=False,
|
||||
require_verified=False,
|
||||
)
|
||||
|
||||
def _resolve(
|
||||
self,
|
||||
@@ -119,6 +127,7 @@ class ModelRegistryResolver:
|
||||
*,
|
||||
registry: RegistryV4 | None,
|
||||
require_enabled: bool,
|
||||
require_verified: bool = True,
|
||||
) -> ResolvedModelConfig:
|
||||
_check_role(role)
|
||||
if registry is None:
|
||||
@@ -190,8 +199,9 @@ class ModelRegistryResolver:
|
||||
credential_revision=credential_revision,
|
||||
)
|
||||
|
||||
# The current verification five-tuple must exist and have passed;
|
||||
# any configuration or credential change invalidates it (4.3).
|
||||
# Run gate (relaxed for provider tests): the current verification
|
||||
# five-tuple must exist and have passed; any configuration or
|
||||
# credential change invalidates it (4.3).
|
||||
record = self._store.get_model_verification(
|
||||
provider_id=provider.id,
|
||||
model_key=model.key,
|
||||
@@ -199,7 +209,7 @@ class ModelRegistryResolver:
|
||||
credential_revision=credential_revision or 0,
|
||||
adapter_spec_revision=spec.spec_revision,
|
||||
)
|
||||
if record is None or record["result"] != "passed":
|
||||
if require_verified and (record is None or record["result"] != "passed"):
|
||||
raise ModelRegistryError(
|
||||
MODEL_NOT_AVAILABLE,
|
||||
"The model has no passing verification for its current configuration.",
|
||||
@@ -216,7 +226,9 @@ class ModelRegistryResolver:
|
||||
)
|
||||
|
||||
budget = self._resolve_budget(model, parameters.max_output_tokens)
|
||||
verified = record["verified_capabilities"]
|
||||
# Without a passing record (provider tests) every capability is
|
||||
# unverified until the test's own probes prove it (section 6.2).
|
||||
verified = record["verified_capabilities"] if record is not None else {}
|
||||
verified_capabilities = Capabilities(
|
||||
**{name: bool(verified.get(name, False)) for name in _CAPABILITY_NAMES}
|
||||
)
|
||||
|
||||
@@ -23,11 +23,13 @@ from typing import Any
|
||||
|
||||
from .errors import (
|
||||
CREDENTIAL_NOT_CONFIGURED,
|
||||
MODEL_CONFIGURATION_CHANGED,
|
||||
MODEL_DISABLED,
|
||||
REGISTRY_REVISION_CONFLICT,
|
||||
RUN_CREDENTIAL_REVISION_UNAVAILABLE,
|
||||
ModelRegistryError,
|
||||
)
|
||||
from .hashing import configuration_hash as _configuration_hash
|
||||
from .schemas import CredentialStatus, CredentialWrite, RegistryV4
|
||||
|
||||
DEFAULT_CONFIG_DIR = Path.home() / ".config" / "evoscientist"
|
||||
@@ -576,6 +578,51 @@ class ModelRuntimeStore:
|
||||
|
||||
# --- model verifications ---------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def _upsert_model_verification(
|
||||
connection: sqlite3.Connection,
|
||||
*,
|
||||
provider_id: str,
|
||||
model_key: str,
|
||||
configuration_hash: str,
|
||||
credential_revision: int,
|
||||
adapter_spec_revision: int,
|
||||
result: str,
|
||||
verified_capabilities: dict[str, bool],
|
||||
error_code: str | None,
|
||||
now: int,
|
||||
) -> None:
|
||||
"""Insert-or-replace the latest test result for one five-tuple."""
|
||||
connection.execute(
|
||||
"INSERT INTO model_verifications "
|
||||
"(provider_id, model_key, configuration_hash, "
|
||||
" credential_revision, adapter_spec_revision, result, "
|
||||
" verified_capabilities_json, verified_at, error_code) "
|
||||
"VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) "
|
||||
"ON CONFLICT (provider_id, model_key, configuration_hash, "
|
||||
" credential_revision, adapter_spec_revision) DO UPDATE SET "
|
||||
"result = excluded.result, "
|
||||
"verified_capabilities_json = "
|
||||
"excluded.verified_capabilities_json, "
|
||||
"verified_at = excluded.verified_at, "
|
||||
"error_code = excluded.error_code",
|
||||
(
|
||||
provider_id,
|
||||
model_key,
|
||||
configuration_hash,
|
||||
credential_revision,
|
||||
adapter_spec_revision,
|
||||
result,
|
||||
json.dumps(
|
||||
verified_capabilities,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
),
|
||||
now,
|
||||
error_code,
|
||||
),
|
||||
)
|
||||
|
||||
def record_model_verification(
|
||||
self,
|
||||
*,
|
||||
@@ -596,34 +643,91 @@ class ModelRuntimeStore:
|
||||
connection = self._connect()
|
||||
try:
|
||||
connection.execute("BEGIN IMMEDIATE")
|
||||
connection.execute(
|
||||
"INSERT INTO model_verifications "
|
||||
"(provider_id, model_key, configuration_hash, "
|
||||
" credential_revision, adapter_spec_revision, result, "
|
||||
" verified_capabilities_json, verified_at, error_code) "
|
||||
"VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) "
|
||||
"ON CONFLICT (provider_id, model_key, configuration_hash, "
|
||||
" credential_revision, adapter_spec_revision) DO UPDATE SET "
|
||||
"result = excluded.result, "
|
||||
"verified_capabilities_json = "
|
||||
"excluded.verified_capabilities_json, "
|
||||
"verified_at = excluded.verified_at, "
|
||||
"error_code = excluded.error_code",
|
||||
(
|
||||
provider_id,
|
||||
model_key,
|
||||
configuration_hash,
|
||||
credential_revision,
|
||||
adapter_spec_revision,
|
||||
result,
|
||||
json.dumps(
|
||||
verified_capabilities,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
),
|
||||
now,
|
||||
error_code,
|
||||
),
|
||||
self._upsert_model_verification(
|
||||
connection,
|
||||
provider_id=provider_id,
|
||||
model_key=model_key,
|
||||
configuration_hash=configuration_hash,
|
||||
credential_revision=credential_revision,
|
||||
adapter_spec_revision=adapter_spec_revision,
|
||||
result=result,
|
||||
verified_capabilities=verified_capabilities,
|
||||
error_code=error_code,
|
||||
now=now,
|
||||
)
|
||||
connection.commit()
|
||||
except BaseException:
|
||||
connection.rollback()
|
||||
raise
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
def record_model_verification_if_current(
|
||||
self,
|
||||
*,
|
||||
expected_registry_revision: int,
|
||||
provider_id: str,
|
||||
model_key: str,
|
||||
configuration_hash: str,
|
||||
credential_revision: int,
|
||||
adapter_spec_revision: int,
|
||||
result: str,
|
||||
verified_capabilities: dict[str, bool],
|
||||
error_code: str | None = None,
|
||||
) -> None:
|
||||
"""The section 9.4 guarded upsert for provider test results.
|
||||
|
||||
The whole guard runs inside one ``BEGIN IMMEDIATE`` transaction: the
|
||||
stored registry revision must still equal
|
||||
``expected_registry_revision`` and the model's current configuration
|
||||
hash must still equal the tested hash. A concurrent change aborts
|
||||
with ``MODEL_CONFIGURATION_CHANGED`` and writes nothing, so a stale
|
||||
test result can never attach to a new configuration.
|
||||
"""
|
||||
if result not in _VERIFICATION_RESULTS:
|
||||
raise ValueError(f"result must be one of {_VERIFICATION_RESULTS}.")
|
||||
now = int(time.time())
|
||||
with self._lock:
|
||||
connection = self._connect()
|
||||
try:
|
||||
connection.execute("BEGIN IMMEDIATE")
|
||||
row = connection.execute(
|
||||
"SELECT revision, registry_json FROM registry_state WHERE id = 1"
|
||||
).fetchone()
|
||||
current_revision = int(row[0]) if row is not None else 1
|
||||
changed = current_revision != expected_registry_revision
|
||||
if not changed:
|
||||
registry = RegistryV4.model_validate(json.loads(str(row[1])))
|
||||
provider = registry.find_provider(provider_id)
|
||||
model = (
|
||||
provider.find_model(model_key) if provider is not None else None
|
||||
)
|
||||
changed = model is None or (
|
||||
_configuration_hash(provider, model) != configuration_hash
|
||||
)
|
||||
if changed:
|
||||
raise ModelRegistryError(
|
||||
MODEL_CONFIGURATION_CHANGED,
|
||||
"The model configuration changed while the provider "
|
||||
"test was running; the result was not recorded.",
|
||||
details=[
|
||||
{
|
||||
"path": "model_ref",
|
||||
"code": MODEL_CONFIGURATION_CHANGED,
|
||||
}
|
||||
],
|
||||
)
|
||||
self._upsert_model_verification(
|
||||
connection,
|
||||
provider_id=provider_id,
|
||||
model_key=model_key,
|
||||
configuration_hash=configuration_hash,
|
||||
credential_revision=credential_revision,
|
||||
adapter_spec_revision=adapter_spec_revision,
|
||||
result=result,
|
||||
verified_capabilities=verified_capabilities,
|
||||
error_code=error_code,
|
||||
now=now,
|
||||
)
|
||||
connection.commit()
|
||||
except BaseException:
|
||||
|
||||
@@ -0,0 +1,600 @@
|
||||
"""Tests for POST /api/model-registry/test (design doc 9.4, 4.3, 5.1, 6.2).
|
||||
|
||||
A Task 2 style fake DNS plus a local threaded HTTP server (registered as a
|
||||
``development_endpoints`` entry) simulates an OpenAI-compatible provider:
|
||||
normal chat completions, 401, 404, connection refusal, and tool-call
|
||||
requests. Covers the full verification flow — frozen resolution, credential
|
||||
resolution, the minimal real call, capability probes, the guarded
|
||||
``model_verifications`` upsert, response redaction, and the
|
||||
``model_config:test`` scope — plus every section 9.4 error path.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import socket
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
|
||||
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.hashing import configuration_hash
|
||||
from EvoScientist.model_registry.http_api import (
|
||||
ApiServices,
|
||||
model_registry_routes,
|
||||
)
|
||||
from EvoScientist.model_registry.platform import DelegationPublicKey
|
||||
from EvoScientist.model_registry.provider_test import ProviderTester
|
||||
from EvoScientist.model_registry.resolver import ModelRegistryResolver
|
||||
from EvoScientist.model_registry.safe_transport import (
|
||||
build_safe_async_http_client,
|
||||
build_safe_http_client,
|
||||
)
|
||||
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"
|
||||
MODEL_REF = {"provider_id": "test-provider", "model_key": "test-model"}
|
||||
|
||||
_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,
|
||||
)
|
||||
|
||||
|
||||
# --- fake OpenAI-compatible provider -----------------------------------------
|
||||
|
||||
|
||||
def _completion(content):
|
||||
return {
|
||||
"id": "chatcmpl-test",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "test-model",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": content},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
}
|
||||
|
||||
|
||||
class _OpenAIHandler(BaseHTTPRequestHandler):
|
||||
"""OpenAI-compatible chat completions with a switchable failure mode."""
|
||||
|
||||
def do_POST(self):
|
||||
length = int(self.headers.get("Content-Length", 0))
|
||||
body = json.loads(self.rfile.read(length))
|
||||
self.server.requests.append(
|
||||
{
|
||||
"path": self.path,
|
||||
"authorization": self.headers.get("Authorization"),
|
||||
"body": body,
|
||||
}
|
||||
)
|
||||
mode = self.server.mode
|
||||
if mode == "reject" or (mode == "fail_tools" and body.get("tools")):
|
||||
status = 401 if mode == "reject" else 400
|
||||
self._respond(
|
||||
status, {"error": {"message": "rejected", "type": "auth_error"}}
|
||||
)
|
||||
return
|
||||
if mode == "missing_model":
|
||||
self._respond(
|
||||
404,
|
||||
{"error": {"message": "model not found", "type": "not_found"}},
|
||||
)
|
||||
return
|
||||
# The structured-output probe sends response_format; answer with JSON.
|
||||
content = json.dumps({"answer": "ok"}) if body.get("response_format") else "ok"
|
||||
self._respond(200, _completion(content))
|
||||
|
||||
def _respond(self, status, payload):
|
||||
body = json.dumps(payload).encode()
|
||||
self.send_response(status)
|
||||
self.send_header("Content-Type", "application/json")
|
||||
self.send_header("Content-Length", str(len(body)))
|
||||
self.end_headers()
|
||||
self.wfile.write(body)
|
||||
|
||||
def log_message(self, *args):
|
||||
pass
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def openai_server():
|
||||
server = ThreadingHTTPServer(("127.0.0.1", 0), _OpenAIHandler)
|
||||
server.requests = []
|
||||
server.mode = "ok"
|
||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
thread.start()
|
||||
yield server
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
thread.join()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def closed_port():
|
||||
"""A loopback port with no listener (connection refused)."""
|
||||
probe = socket.socket()
|
||||
probe.bind(("127.0.0.1", 0))
|
||||
port = probe.getsockname()[1]
|
||||
probe.close()
|
||||
return port
|
||||
|
||||
|
||||
def _fake_getaddrinfo(counter):
|
||||
def fake(host, port, *args, **kwargs):
|
||||
counter["calls"] += 1
|
||||
assert host == "localhost"
|
||||
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("127.0.0.1", port))]
|
||||
|
||||
return fake
|
||||
|
||||
|
||||
# --- registry fixtures ---------------------------------------------------------
|
||||
|
||||
|
||||
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": True,
|
||||
"structured_output": True,
|
||||
},
|
||||
}
|
||||
payload.update(overrides)
|
||||
return payload
|
||||
|
||||
|
||||
def _provider_payload(port, **overrides):
|
||||
provider = {
|
||||
"id": "test-provider",
|
||||
"name": "Test Provider",
|
||||
"adapter": "openai-compatible",
|
||||
"base_url": f"http://localhost:{port}/v1",
|
||||
"auth": {"mode": "api_key", "credential_id": "test-key"},
|
||||
"enabled": True,
|
||||
"runtime": {
|
||||
"timeout_seconds": 120,
|
||||
"max_retries": 2,
|
||||
"default_temperature": 0.7,
|
||||
"default_top_p": 0.95,
|
||||
"default_reasoning_effort": "auto",
|
||||
},
|
||||
"models": [
|
||||
{
|
||||
"key": "test-model",
|
||||
"name": "Test Model",
|
||||
"upstream_model_id": "test-model",
|
||||
"enabled": False,
|
||||
"runtime": _model_runtime(),
|
||||
}
|
||||
],
|
||||
}
|
||||
provider.update(overrides)
|
||||
return provider
|
||||
|
||||
|
||||
def _registry_payload(port, **provider_overrides):
|
||||
return {
|
||||
"version": 4,
|
||||
"revision": 1,
|
||||
"state": "bootstrap",
|
||||
"defaults": {"primary": None, "auxiliary": None},
|
||||
"providers": [_provider_payload(port, **provider_overrides)],
|
||||
}
|
||||
|
||||
|
||||
def _save(store, port, *, credential=True, **provider_overrides):
|
||||
writes = (
|
||||
[CredentialWrite(credential_id="test-key", secret_value=SECRET)]
|
||||
if credential
|
||||
else []
|
||||
)
|
||||
return store.save_registry(
|
||||
expected_revision=1,
|
||||
registry=RegistryV4.model_validate(
|
||||
_registry_payload(port, **provider_overrides)
|
||||
),
|
||||
credential_writes=writes,
|
||||
)
|
||||
|
||||
|
||||
# --- app fixtures --------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def store(tmp_path):
|
||||
return ModelRuntimeStore(config_dir=tmp_path)
|
||||
|
||||
|
||||
def _policy_for(port):
|
||||
return EndpointPolicy(
|
||||
[
|
||||
DevelopmentEndpoint(
|
||||
id="test-provider",
|
||||
url=f"http://localhost:{port}/v1",
|
||||
label="Test provider",
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def _services_for(store, port):
|
||||
resolver = ModelRegistryResolver(store)
|
||||
policy = _policy_for(port)
|
||||
fake_dns = _fake_getaddrinfo({"calls": 0})
|
||||
tester = ProviderTester(
|
||||
store,
|
||||
resolver,
|
||||
policy,
|
||||
sync_client_builder=lambda timeout, retries: build_safe_http_client(
|
||||
policy, timeout=timeout, retries=retries, getaddrinfo=fake_dns
|
||||
),
|
||||
async_client_builder=lambda timeout, retries: build_safe_async_http_client(
|
||||
policy, timeout=timeout, retries=retries, getaddrinfo=fake_dns
|
||||
),
|
||||
)
|
||||
return ApiServices(
|
||||
store=store,
|
||||
resolver=resolver,
|
||||
snapshot_service=SnapshotService(store, resolver),
|
||||
endpoint_policy=policy,
|
||||
authenticator=BffAuthenticator(
|
||||
service_token=SERVICE_TOKEN,
|
||||
service_token_hash=None,
|
||||
delegation_keys=(
|
||||
DelegationPublicKey(deployment_id="webui-1", public_key=_PUBLIC_PEM),
|
||||
),
|
||||
jti_store=store,
|
||||
),
|
||||
provider_tester=tester,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def services(store, openai_server):
|
||||
return _services_for(store, openai_server.server_address[1])
|
||||
|
||||
|
||||
def _client_for(services):
|
||||
app = Starlette(routes=model_registry_routes(lambda: services))
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client(services):
|
||||
return _client_for(services)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def saved_store(store, openai_server):
|
||||
_save(store, openai_server.server_address[1])
|
||||
return store
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def test_client(client, saved_store):
|
||||
return client
|
||||
|
||||
|
||||
# --- auth helpers ---------------------------------------------------------------
|
||||
|
||||
|
||||
def _headers(scopes):
|
||||
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,
|
||||
}
|
||||
return {
|
||||
"Authorization": f"Bearer {SERVICE_TOKEN}",
|
||||
"X-Evo-Actor": jwt.encode(claims, _PRIVATE_PEM, algorithm="ES256"),
|
||||
}
|
||||
|
||||
|
||||
def _test_headers():
|
||||
return _headers(["model_config:test"])
|
||||
|
||||
|
||||
def _post(client, *, model_ref=None, expected_revision=2, headers=None):
|
||||
return client.post(
|
||||
"/api/model-registry/test",
|
||||
headers=headers or _test_headers(),
|
||||
json={
|
||||
"expected_registry_revision": expected_revision,
|
||||
"model_ref": model_ref or MODEL_REF,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
# --- success paths ---------------------------------------------------------------
|
||||
|
||||
|
||||
def test_success_full_flow(test_client, saved_store, openai_server):
|
||||
response = _post(test_client)
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert set(body) == {
|
||||
"ok",
|
||||
"registry_revision",
|
||||
"model_ref",
|
||||
"adapter_spec_revision",
|
||||
"effective_request_options",
|
||||
"model_status",
|
||||
"latency_ms",
|
||||
}
|
||||
assert body["ok"] is True
|
||||
assert body["registry_revision"] == 2
|
||||
assert body["model_ref"] == MODEL_REF
|
||||
assert body["adapter_spec_revision"] == 1
|
||||
# The redacted build_request output matches the section 9.4 example shape.
|
||||
assert body["effective_request_options"] == {
|
||||
"timeout": 120,
|
||||
"max_retries": 2,
|
||||
"max_tokens": 32768,
|
||||
"temperature": 0.7,
|
||||
"top_p": 0.95,
|
||||
}
|
||||
assert isinstance(body["latency_ms"], int)
|
||||
assert body["latency_ms"] >= 0
|
||||
# The model was saved disabled: verified, not yet selectable.
|
||||
assert body["model_status"]["state"] == "verified"
|
||||
assert body["model_status"]["selectable"] is False
|
||||
assert body["model_status"]["verification"]["status"] == "passed"
|
||||
assert body["model_status"]["effective_capabilities"] == {
|
||||
"tools": True,
|
||||
"vision": True,
|
||||
"structured_output": True,
|
||||
}
|
||||
# The response never carries the secret (section 5.1).
|
||||
assert SECRET not in response.text
|
||||
|
||||
# The five-tuple record was upserted (section 4.3).
|
||||
registry = saved_store.load_registry()
|
||||
provider = registry.find_provider("test-provider")
|
||||
model = provider.find_model("test-model")
|
||||
records = saved_store.list_model_verifications()
|
||||
assert len(records) == 1
|
||||
record = records[0]
|
||||
assert record["provider_id"] == "test-provider"
|
||||
assert record["model_key"] == "test-model"
|
||||
assert record["configuration_hash"] == configuration_hash(provider, model)
|
||||
assert record["credential_revision"] == 1
|
||||
assert record["adapter_spec_revision"] == 1
|
||||
assert record["result"] == "passed"
|
||||
assert record["verified_capabilities"] == {
|
||||
"tools": True,
|
||||
"vision": True,
|
||||
"structured_output": True,
|
||||
}
|
||||
assert record["error_code"] is None
|
||||
|
||||
# The provider saw the minimal chat call plus one request per declared
|
||||
# capability probe; every request carried the resolved credential.
|
||||
requests = openai_server.requests
|
||||
assert len(requests) == 4
|
||||
assert all(request["path"] == "/v1/chat/completions" for request in requests)
|
||||
assert all(request["authorization"] == f"Bearer {SECRET}" for request in requests)
|
||||
assert any("tools" in request["body"] for request in requests)
|
||||
assert any("response_format" in request["body"] for request in requests)
|
||||
assert any("image_url" in json.dumps(request["body"]) for request in requests)
|
||||
|
||||
|
||||
def test_enabled_model_reports_enabled_status(client, store, openai_server):
|
||||
port = openai_server.server_address[1]
|
||||
provider = _provider_payload(port)
|
||||
provider["models"][0]["enabled"] = True
|
||||
_save(store, port, models=provider["models"])
|
||||
response = _post(client)
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["ok"] is True
|
||||
assert body["model_status"]["state"] == "enabled"
|
||||
assert body["model_status"]["selectable"] is True
|
||||
|
||||
|
||||
def test_tools_probe_failure_marks_capability_false(
|
||||
test_client, saved_store, openai_server
|
||||
):
|
||||
openai_server.mode = "fail_tools"
|
||||
response = _post(test_client)
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["ok"] is True
|
||||
assert body["model_status"]["effective_capabilities"]["tools"] is False
|
||||
assert body["model_status"]["effective_capabilities"]["vision"] is True
|
||||
|
||||
record = saved_store.list_model_verifications()[0]
|
||||
assert record["result"] == "passed"
|
||||
assert record["verified_capabilities"] == {
|
||||
"tools": False,
|
||||
"vision": True,
|
||||
"structured_output": True,
|
||||
}
|
||||
|
||||
|
||||
# --- precondition and provider errors -------------------------------------------
|
||||
|
||||
|
||||
def test_unsaved_model_ref_is_404(test_client, saved_store):
|
||||
response = _post(
|
||||
test_client,
|
||||
model_ref={"provider_id": "test-provider", "model_key": "free-draft"},
|
||||
)
|
||||
assert response.status_code == 404
|
||||
assert response.json()["code"] == "MODEL_NOT_FOUND"
|
||||
|
||||
response = _post(
|
||||
test_client, model_ref={"provider_id": "draft", "model_key": "test-model"}
|
||||
)
|
||||
assert response.status_code == 404
|
||||
assert response.json()["code"] == "MODEL_NOT_FOUND"
|
||||
assert saved_store.list_model_verifications() == []
|
||||
|
||||
|
||||
def test_missing_credential_is_422(client, store, openai_server):
|
||||
_save(store, openai_server.server_address[1], credential=False)
|
||||
response = _post(client)
|
||||
assert response.status_code == 422
|
||||
assert response.json()["code"] == "CREDENTIAL_NOT_CONFIGURED"
|
||||
assert store.list_model_verifications() == []
|
||||
|
||||
|
||||
def test_credential_rejected_records_failed_result(
|
||||
test_client, saved_store, openai_server
|
||||
):
|
||||
openai_server.mode = "reject"
|
||||
response = _post(test_client)
|
||||
assert response.status_code == 422
|
||||
body = response.json()
|
||||
assert body["code"] == "CREDENTIAL_REJECTED"
|
||||
assert SECRET not in response.text
|
||||
|
||||
record = saved_store.list_model_verifications()[0]
|
||||
assert record["result"] == "failed"
|
||||
assert record["error_code"] == "CREDENTIAL_REJECTED"
|
||||
assert record["verified_capabilities"] == {
|
||||
"tools": False,
|
||||
"vision": False,
|
||||
"structured_output": False,
|
||||
}
|
||||
|
||||
|
||||
def test_provider_404_is_model_not_found(test_client, saved_store, openai_server):
|
||||
openai_server.mode = "missing_model"
|
||||
response = _post(test_client)
|
||||
# The section 9.5 table pins MODEL_NOT_FOUND to 404 for every source.
|
||||
assert response.status_code == 404
|
||||
assert response.json()["code"] == "MODEL_NOT_FOUND"
|
||||
record = saved_store.list_model_verifications()[0]
|
||||
assert record["result"] == "failed"
|
||||
assert record["error_code"] == "MODEL_NOT_FOUND"
|
||||
|
||||
|
||||
def test_connection_refused_is_provider_unreachable(store, closed_port):
|
||||
# The policy must register the unreachable port; the failure must come
|
||||
# from the connection attempt, not the endpoint allowlist.
|
||||
_save(store, closed_port)
|
||||
client = _client_for(_services_for(store, closed_port))
|
||||
response = _post(client)
|
||||
assert response.status_code == 422
|
||||
assert response.json()["code"] == "PROVIDER_UNREACHABLE"
|
||||
record = store.list_model_verifications()[0]
|
||||
assert record["result"] == "failed"
|
||||
assert record["error_code"] == "PROVIDER_UNREACHABLE"
|
||||
|
||||
|
||||
def test_unregistered_local_endpoint_is_rejected(client, store):
|
||||
# Saved directly through the store (no save-time policy hook), so the
|
||||
# test endpoint must re-validate the base URL before any I/O.
|
||||
port = 9 # unregistered loopback endpoint
|
||||
store.save_registry(
|
||||
expected_revision=1,
|
||||
registry=RegistryV4.model_validate(_registry_payload(port)),
|
||||
credential_writes=[
|
||||
CredentialWrite(credential_id="test-key", secret_value=SECRET)
|
||||
],
|
||||
)
|
||||
response = _post(client)
|
||||
assert response.status_code == 422
|
||||
assert response.json()["code"] == "ENDPOINT_NOT_ALLOWED"
|
||||
assert store.list_model_verifications() == []
|
||||
|
||||
|
||||
# --- verification record transaction --------------------------------------------
|
||||
|
||||
|
||||
def test_concurrent_registry_change_is_409_and_writes_nothing(
|
||||
test_client, saved_store, services, monkeypatch
|
||||
):
|
||||
tester = services.provider_tester
|
||||
original_execute = tester._execute
|
||||
|
||||
def mutate_then_execute(*args, **kwargs):
|
||||
# A concurrent admin save between resolution and the record write.
|
||||
registry = saved_store.load_registry()
|
||||
payload = registry.model_dump(mode="json")
|
||||
payload["providers"][0]["models"][0]["name"] = "Renamed Mid-Test"
|
||||
saved_store.save_registry(
|
||||
expected_revision=registry.revision,
|
||||
registry=RegistryV4.model_validate(payload),
|
||||
)
|
||||
return original_execute(*args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(tester, "_execute", mutate_then_execute)
|
||||
|
||||
response = _post(test_client)
|
||||
assert response.status_code == 409
|
||||
assert response.json()["code"] == "MODEL_CONFIGURATION_CHANGED"
|
||||
# The stale test result must not be recorded for the new configuration.
|
||||
assert saved_store.list_model_verifications() == []
|
||||
|
||||
|
||||
def test_stale_expected_revision_is_409(test_client, saved_store):
|
||||
response = _post(test_client, expected_revision=99)
|
||||
assert response.status_code == 409
|
||||
assert response.json()["code"] == "MODEL_CONFIGURATION_CHANGED"
|
||||
assert saved_store.list_model_verifications() == []
|
||||
|
||||
|
||||
# --- authentication ---------------------------------------------------------------
|
||||
|
||||
|
||||
def test_requires_test_scope(test_client, saved_store):
|
||||
response = _post(test_client, headers=_headers(["model_config:write"]))
|
||||
assert response.status_code == 403
|
||||
assert response.json()["code"] == "FORBIDDEN"
|
||||
|
||||
|
||||
def test_requires_authentication(test_client, saved_store):
|
||||
response = test_client.post(
|
||||
"/api/model-registry/test",
|
||||
json={"expected_registry_revision": 2, "model_ref": MODEL_REF},
|
||||
)
|
||||
assert response.status_code == 401
|
||||
Reference in New Issue
Block a user