Files
EvoScientist/tests/test_snapshots.py
T
m4 c8c46eab16 fix(model-registry): make snapshot abort atomic against concurrent bind
The abort path was read-then-write with an unconditional UPDATE, so a bind
committing between the two calls was clobbered back to aborted, losing its
langgraph_run_id. Add a conditional store-level abort_run_snapshot
(prepared-only UPDATE, rowcount-checked) and re-read on a lost race, matching
the bind loop. Also pin the inherit selection_hash test to a hardcoded
SHA-256 literal instead of reimplementing the serialization in the test.
2026-07-21 08:45:33 +08:00

613 lines
24 KiB
Python

"""Tests for the run runtime snapshot service (design doc 5.2, 8.1, 8.2).
Covers creation with inherit semantics, the selection-hash idempotency
rules, bind/abort state transitions, binding validation on reads, TTLs,
adapter-spec revalidation, and credential resolution against frozen
revisions — with no in-process secret caching.
"""
from __future__ import annotations
import json
import time
import pytest
from EvoScientist.model_registry.adapters import adapter_specs, find_adapter_spec
from EvoScientist.model_registry.errors import (
ADAPTER_NOT_SUPPORTED,
MODEL_REGISTRY_NOT_READY,
RUN_CREDENTIAL_REVISION_UNAVAILABLE,
RUN_REQUEST_CONFLICT,
SNAPSHOT_ALREADY_BOUND,
SNAPSHOT_EXPIRED,
SNAPSHOT_NOT_FOUND,
ModelRegistryError,
)
from EvoScientist.model_registry.hashing import configuration_hash
from EvoScientist.model_registry.resolver import ModelRegistryResolver
from EvoScientist.model_registry.schemas import (
CredentialWrite,
ModelRef,
RegistryV4,
)
from EvoScientist.model_registry.snapshots import (
BOUND_RETENTION_SECONDS,
PREPARED_TTL_SECONDS,
SnapshotCreateRequest,
SnapshotService,
compute_selection_hash,
config_for_role,
public_snapshot_view,
)
from EvoScientist.model_registry.store import ModelRuntimeStore
ZHIPU_REF = ModelRef(provider_id="zhipu-glm", model_key="glm-5.2")
OLLAMA_REF = ModelRef(provider_id="local-ollama", model_key="qwen3")
SECRET = "sk-live-9876abcd"
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():
return {
"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(),
}
],
}
def _ollama_provider():
return {
"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(),
}
],
}
def _registry(*, auxiliary_default=True) -> RegistryV4:
return RegistryV4.model_validate(
{
"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"}
if auxiliary_default
else None
),
},
"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,
},
)
@pytest.fixture
def store(tmp_path):
return ModelRuntimeStore(config_dir=tmp_path)
@pytest.fixture
def active_store(store):
registry = store.save_registry(
expected_revision=1,
registry=_registry(),
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 service(active_store):
return SnapshotService(active_store, ModelRegistryResolver(active_store))
def _request(**overrides):
payload = {
"run_request_id": "req-1",
"thread_id": "thread-1",
"deployment_id": "local",
"model_selection_revision": 0,
"primary": None,
"auxiliary": None,
}
payload.update(overrides)
return SnapshotCreateRequest.model_validate(payload)
# SHA-256 of the canonical JSON '{"auxiliary":null,"primary":null}' — the
# inherit/inherit selection. Pinned as a literal to lock the wire contract.
INHERIT_SELECTION_HASH = (
"697c046214ccc3ddee1018af7eb6c21dbd5bfd01fbca9cb594ffc273deb55abc"
)
class TestCreate:
def test_inherit_resolves_registry_defaults(self, service, active_store):
creation = service.create(_request(model_selection_revision=4))
assert creation.created is True
snapshot = creation.snapshot
assert snapshot.status == "prepared"
assert snapshot.langgraph_run_id is None
payload = snapshot.payload
assert payload.registry_revision == 2
assert payload.model_selection_revision == 4
assert payload.primary.model_ref == ZHIPU_REF
assert payload.primary.role == "primary"
assert payload.auxiliary is not None
assert payload.auxiliary.model_ref == OLLAMA_REF
assert payload.auxiliary.role == "auxiliary"
def test_inherit_without_auxiliary_default(self, active_store):
store = active_store
registry = store.load_registry()
store.save_registry(
expected_revision=registry.revision,
registry=_registry(auxiliary_default=False),
)
service = SnapshotService(store, ModelRegistryResolver(store))
creation = service.create(_request())
assert creation.snapshot.payload.auxiliary is None
def test_explicit_selection_freezes_both_roles(self, service):
creation = service.create(_request(primary=ZHIPU_REF, auxiliary=OLLAMA_REF))
payload = creation.snapshot.payload
assert payload.primary.model_ref == ZHIPU_REF
assert payload.auxiliary is not None
assert payload.auxiliary.model_ref == OLLAMA_REF
def test_bootstrap_registry_is_not_ready(self, store):
service = SnapshotService(store, ModelRegistryResolver(store))
with pytest.raises(ModelRegistryError) as excinfo:
service.create(_request())
assert excinfo.value.code == MODEL_REGISTRY_NOT_READY
assert excinfo.value.http_status == 422
def test_selection_hash_uses_pre_resolution_semantics(self, service):
creation = service.create(_request())
snapshot = creation.snapshot
# The inherit/inherit hash is pinned to a known literal (computed
# offline from the canonical JSON) so the contract can't drift
# together with a reimplemented expectation.
assert snapshot.selection_hash == INHERIT_SELECTION_HASH
assert compute_selection_hash(None, None) == INHERIT_SELECTION_HASH
# An explicit selection equal to the defaults still hashes differently.
assert compute_selection_hash(ZHIPU_REF, OLLAMA_REF) != INHERIT_SELECTION_HASH
def test_selection_hash_ignores_selection_revision(self, service):
first = service.create(_request(model_selection_revision=4))
second = service.create(_request(model_selection_revision=9))
assert second.created is False
assert second.snapshot.snapshot_id == first.snapshot.snapshot_id
def test_idempotent_same_selection_returns_original(self, service):
first = service.create(_request())
second = service.create(_request())
assert first.created is True
assert second.created is False
assert second.snapshot.snapshot_id == first.snapshot.snapshot_id
def test_conflicting_selection_same_triplet(self, service):
service.create(_request())
with pytest.raises(ModelRegistryError) as excinfo:
service.create(_request(primary=ZHIPU_REF))
assert excinfo.value.code == RUN_REQUEST_CONFLICT
assert excinfo.value.http_status == 409
def test_recreate_after_expired(self, service, active_store):
first = service.create(_request())
active_store.set_run_snapshot_status(first.snapshot.snapshot_id, "expired")
second = service.create(_request())
assert second.created is True
assert second.snapshot.snapshot_id != first.snapshot.snapshot_id
def test_recreate_after_aborted(self, service):
first = service.create(_request())
service.abort(first.snapshot.snapshot_id)
second = service.create(_request())
assert second.created is True
assert second.snapshot.snapshot_id != first.snapshot.snapshot_id
def test_payload_freezes_credential_revision_without_secret(self, service):
creation = service.create(_request())
payload = creation.snapshot.payload
assert payload.primary.auth_ref.credential_id == "zhipu-primary"
assert payload.primary.auth_ref.credential_revision == 1
encoded = json.dumps(payload.model_dump(mode="json"))
assert SECRET not in encoded
def test_prepared_ttl_is_fifteen_minutes(self, service):
creation = service.create(_request())
snapshot = creation.snapshot
assert PREPARED_TTL_SECONDS == 15 * 60
delta = snapshot.expires_at - snapshot.created_at
assert PREPARED_TTL_SECONDS - 1 <= delta <= PREPARED_TTL_SECONDS + 1
class TestBind:
def test_bind_once(self, service):
creation = service.create(_request())
before = int(time.time())
bound = service.bind(creation.snapshot.snapshot_id, "lg-run-1")
assert bound.status == "bound"
assert bound.langgraph_run_id == "lg-run-1"
assert BOUND_RETENTION_SECONDS == 24 * 60 * 60
delta = bound.expires_at - before
assert BOUND_RETENTION_SECONDS - 1 <= delta <= BOUND_RETENTION_SECONDS + 1
def test_bind_same_run_id_is_idempotent(self, service):
creation = service.create(_request())
service.bind(creation.snapshot.snapshot_id, "lg-run-1")
again = service.bind(creation.snapshot.snapshot_id, "lg-run-1")
assert again.status == "bound"
assert again.langgraph_run_id == "lg-run-1"
def test_bind_different_run_id_conflicts(self, service):
creation = service.create(_request())
service.bind(creation.snapshot.snapshot_id, "lg-run-1")
with pytest.raises(ModelRegistryError) as excinfo:
service.bind(creation.snapshot.snapshot_id, "lg-run-2")
assert excinfo.value.code == SNAPSHOT_ALREADY_BOUND
def test_bind_expired_snapshot(self, service, active_store):
creation = service.create(_request())
active_store.set_run_snapshot_status(creation.snapshot.snapshot_id, "expired")
with pytest.raises(ModelRegistryError) as excinfo:
service.bind(creation.snapshot.snapshot_id, "lg-run-1")
assert excinfo.value.code == SNAPSHOT_EXPIRED
def test_bind_aborted_snapshot(self, service):
creation = service.create(_request())
service.abort(creation.snapshot.snapshot_id)
with pytest.raises(ModelRegistryError) as excinfo:
service.bind(creation.snapshot.snapshot_id, "lg-run-1")
assert excinfo.value.code == SNAPSHOT_EXPIRED
def test_bind_unknown_snapshot(self, service):
with pytest.raises(ModelRegistryError) as excinfo:
service.bind("snap-missing", "lg-run-1")
assert excinfo.value.code == SNAPSHOT_NOT_FOUND
assert excinfo.value.http_status == 404
class TestAbort:
def test_abort_prepared(self, service, active_store):
creation = service.create(_request())
service.abort(creation.snapshot.snapshot_id)
row = active_store.get_run_snapshot(creation.snapshot.snapshot_id)
assert row["status"] == "aborted"
def test_abort_is_idempotent(self, service):
creation = service.create(_request())
service.abort(creation.snapshot.snapshot_id)
service.abort(creation.snapshot.snapshot_id)
def test_abort_bound_conflicts(self, service):
creation = service.create(_request())
service.bind(creation.snapshot.snapshot_id, "lg-run-1")
with pytest.raises(ModelRegistryError) as excinfo:
service.abort(creation.snapshot.snapshot_id)
assert excinfo.value.code == SNAPSHOT_ALREADY_BOUND
def test_abort_expired(self, service, active_store):
creation = service.create(_request())
active_store.set_run_snapshot_status(creation.snapshot.snapshot_id, "expired")
with pytest.raises(ModelRegistryError) as excinfo:
service.abort(creation.snapshot.snapshot_id)
assert excinfo.value.code == SNAPSHOT_EXPIRED
def test_abort_unknown(self, service):
with pytest.raises(ModelRegistryError) as excinfo:
service.abort("snap-missing")
assert excinfo.value.code == SNAPSHOT_NOT_FOUND
def test_abort_run_snapshot_store_update_is_conditional(
self, service, active_store
):
creation = service.create(_request())
snapshot_id = creation.snapshot.snapshot_id
# The store-level abort only fires on prepared rows.
assert active_store.abort_run_snapshot(snapshot_id) is True
assert active_store.abort_run_snapshot(snapshot_id) is False
bound = service.create(_request(run_request_id="req-2"))
service.bind(bound.snapshot.snapshot_id, "lg-run-7")
assert active_store.abort_run_snapshot(bound.snapshot.snapshot_id) is False
row = active_store.get_run_snapshot(bound.snapshot.snapshot_id)
assert row["status"] == "bound"
assert row["langgraph_run_id"] == "lg-run-7"
def test_abort_losing_bind_race_keeps_bound_state(
self, service, active_store, monkeypatch
):
"""A bind committing between abort's read and write must win.
The abort read observes ``prepared``; a concurrent bind then commits;
the abort write must fail its conditional update, re-read, and raise
``SNAPSHOT_ALREADY_BOUND`` instead of clobbering the bound row (and
its ``langgraph_run_id``) back to ``aborted``.
"""
creation = service.create(_request())
snapshot_id = creation.snapshot.snapshot_id
original_get = active_store.get_run_snapshot
raced = False
def get_with_concurrent_bind(sid):
nonlocal raced
row = original_get(sid)
if not raced and row is not None and row["status"] == "prepared":
raced = True
assert active_store.bind_run_snapshot(
sid,
langgraph_run_id="lg-run-9",
expires_at=row["expires_at"] + 100,
)
return row
monkeypatch.setattr(active_store, "get_run_snapshot", get_with_concurrent_bind)
with pytest.raises(ModelRegistryError) as excinfo:
service.abort(snapshot_id)
assert raced is True
assert excinfo.value.code == SNAPSHOT_ALREADY_BOUND
row = active_store.get_run_snapshot(snapshot_id)
assert row["status"] == "bound"
assert row["langgraph_run_id"] == "lg-run-9"
class TestGet:
def test_get_returns_snapshot(self, service):
creation = service.create(_request())
snapshot = service.get(
creation.snapshot.snapshot_id,
deployment_id="local",
thread_id="thread-1",
)
assert snapshot.snapshot_id == creation.snapshot.snapshot_id
assert snapshot.payload.primary.model_ref == ZHIPU_REF
def test_get_rejects_foreign_thread(self, service):
creation = service.create(_request())
with pytest.raises(ModelRegistryError) as excinfo:
service.get(
creation.snapshot.snapshot_id,
deployment_id="local",
thread_id="thread-2",
)
assert excinfo.value.code == SNAPSHOT_NOT_FOUND
def test_get_rejects_foreign_deployment(self, service):
creation = service.create(_request())
with pytest.raises(ModelRegistryError) as excinfo:
service.get(
creation.snapshot.snapshot_id,
deployment_id="other",
thread_id="thread-1",
)
assert excinfo.value.code == SNAPSHOT_NOT_FOUND
def test_get_expired_fails(self, service, active_store):
creation = service.create(_request())
active_store.set_run_snapshot_status(creation.snapshot.snapshot_id, "expired")
with pytest.raises(ModelRegistryError) as excinfo:
service.get(
creation.snapshot.snapshot_id,
deployment_id="local",
thread_id="thread-1",
)
assert excinfo.value.code == SNAPSHOT_EXPIRED
def test_get_fails_when_frozen_spec_revision_removed(self, active_store):
service = SnapshotService(active_store, ModelRegistryResolver(active_store))
creation = service.create(_request())
# Simulate a contract upgrade that removed spec_revision 1 entirely.
bumped = tuple(
spec.model_copy(update={"spec_revision": 2}) for spec in adapter_specs()
)
stale_service = SnapshotService(
active_store, ModelRegistryResolver(active_store, specs=bumped)
)
with pytest.raises(ModelRegistryError) as excinfo:
stale_service.get(
creation.snapshot.snapshot_id,
deployment_id="local",
thread_id="thread-1",
)
assert excinfo.value.code == ADAPTER_NOT_SUPPORTED
class TestCleanupExpired:
def _insert(self, store, snapshot_id, *, status, expires_at):
store.insert_run_snapshot(
snapshot_id=snapshot_id,
deployment_id="local",
thread_id="thread-1",
run_request_id=f"req-{snapshot_id}",
selection_hash="b" * 64,
payload={},
expires_at=expires_at,
status=status,
langgraph_run_id="lg-run" if status == "bound" else None,
)
def test_cleanup_marks_due_snapshots_expired(self, store):
service = SnapshotService(store, ModelRegistryResolver(store))
now = 1_000_000
self._insert(store, "snap-p-old", status="prepared", expires_at=now - 1)
self._insert(store, "snap-b-old", status="bound", expires_at=now - 1)
self._insert(store, "snap-p-new", status="prepared", expires_at=now + 900)
expired = service.cleanup_expired(now)
assert sorted(expired) == ["snap-b-old", "snap-p-old"]
assert store.get_run_snapshot("snap-p-old")["status"] == "expired"
assert store.get_run_snapshot("snap-b-old")["status"] == "expired"
assert store.get_run_snapshot("snap-p-new")["status"] == "prepared"
class TestRoleMapping:
def test_auxiliary_roles_fall_back_to_primary(self, active_store):
store = active_store
registry = store.load_registry()
store.save_registry(
expected_revision=registry.revision,
registry=_registry(auxiliary_default=False),
)
service = SnapshotService(store, ModelRegistryResolver(store))
snapshot = service.create(_request()).snapshot
for role in ("primary", "auxiliary", "summary", "tool_selector"):
assert config_for_role(snapshot, role).model_ref == ZHIPU_REF
def test_auxiliary_roles_use_frozen_auxiliary(self, service):
snapshot = service.create(_request()).snapshot
assert config_for_role(snapshot, "primary").model_ref == ZHIPU_REF
for role in ("auxiliary", "summary", "tool_selector"):
config = config_for_role(snapshot, role)
assert config.model_ref == OLLAMA_REF
assert config.role == "auxiliary"
def test_unknown_role_rejected(self, service):
snapshot = service.create(_request()).snapshot
with pytest.raises(ValueError, match="role"):
config_for_role(snapshot, "planner")
class TestSnapshotCredentials:
def test_resolves_frozen_revision(self, service):
snapshot = service.create(_request()).snapshot
assert service.resolve_snapshot_credential(snapshot, "primary") == SECRET
def test_rotation_keeps_frozen_revision(self, service, active_store):
snapshot = service.create(_request()).snapshot
active_store.write_credential_version("zhipu-primary", "sk-rotated-2222")
assert service.resolve_snapshot_credential(snapshot, "primary") == SECRET
def test_destroyed_revision_fails_loudly(self, service, active_store):
snapshot = service.create(_request()).snapshot
active_store.retire_credential_version("zhipu-primary", 1)
with pytest.raises(ModelRegistryError) as excinfo:
service.resolve_snapshot_credential(snapshot, "primary")
assert excinfo.value.code == RUN_CREDENTIAL_REVISION_UNAVAILABLE
def test_no_in_process_secret_cache(self, service, active_store):
snapshot = service.create(_request()).snapshot
# A brand-new store/service pair on the same database still resolves.
fresh_store = ModelRuntimeStore(config_dir=active_store.config_dir)
fresh_service = SnapshotService(fresh_store, ModelRegistryResolver(fresh_store))
assert fresh_service.resolve_snapshot_credential(snapshot, "primary") == SECRET
def test_auxiliary_role_resolves_auxiliary_credential(self, service):
snapshot = service.create(_request()).snapshot
# The frozen auxiliary is the mode=none ollama model: no secret.
assert service.resolve_snapshot_credential(snapshot, "summary") == ""
assert service.resolve_snapshot_credential(snapshot, "tool_selector") == ""
def test_mode_none_returns_empty_string(self, service):
snapshot = service.create(_request(primary=OLLAMA_REF)).snapshot
assert service.resolve_snapshot_credential(snapshot, "primary") == ""
class TestPublicView:
def test_matches_section_8_2_shape(self, service):
snapshot = service.create(_request()).snapshot
view = public_snapshot_view(snapshot)
assert view["snapshot_id"] == snapshot.snapshot_id
assert view["registry_revision"] == 2
primary = view["primary"]
assert primary["provider_id"] == "zhipu-glm"
assert primary["model_key"] == "glm-5.2"
assert primary["adapter_spec_revision"] == 1
assert primary["runtime"] == {
"max_output_tokens": 32768,
"temperature": 0.7,
"top_p": 0.95,
"timeout_seconds": 120,
"max_retries": 2,
}
assert view["auxiliary"]["provider_id"] == "local-ollama"
def test_never_contains_secret_or_base_url(self, service):
snapshot = service.create(_request()).snapshot
encoded = json.dumps(public_snapshot_view(snapshot))
assert SECRET not in encoded
assert "open.bigmodel.cn" not in encoded
assert "localhost" not in encoded