From c8c46eab162b47b60416796cbd47fa62a5ea545b Mon Sep 17 00:00:00 2001 From: m4 Date: Tue, 21 Jul 2026 08:45:23 +0800 Subject: [PATCH] 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. --- .superpowers/sdd/briefs/task-4-report.md | 14 ++++ EvoScientist/model_registry/snapshots.py | 36 +++++++---- EvoScientist/model_registry/store.py | 24 +++++++ tests/test_snapshots.py | 81 ++++++++++++++++++------ 4 files changed, 123 insertions(+), 32 deletions(-) diff --git a/.superpowers/sdd/briefs/task-4-report.md b/.superpowers/sdd/briefs/task-4-report.md index ae09295..095f0c9 100644 --- a/.superpowers/sdd/briefs/task-4-report.md +++ b/.superpowers/sdd/briefs/task-4-report.md @@ -73,3 +73,17 @@ 2. 终态(expired/aborted)快照的 bind/get 统一抛 `SNAPSHOT_EXPIRED`;文档只明文规定 expired 的情形,aborted 按同一终态语义处理。 3. 冻结的 `budget.message_budget` 取 base(仅扣系统预留);逐次调用的 has_tools/has_attachments 重算属 Task 6 的 MessageBudgetMiddleware。 4. `resolve` 支持 `registry=` 参数供 `create` 传入同一份 Registry,保证快照 `registry_revision` 与解析所用文档一致。 + +## 评审修复(2026-07-21,commit 见下) + +1. **Important:`abort` read-then-write 竞态**。原实现先读后写且 `set_run_snapshot_status` 为无条件 UPDATE,并发 bind 在两次调用间提交时会把 bound 改写为 aborted 并丢失 `langgraph_run_id`。修复:store 层新增 `abort_run_snapshot`(`UPDATE ... SET status='aborted' WHERE snapshot_id=? AND status='prepared'`,按 rowcount 判定),`SnapshotService.abort` 改为与 `bind` 相同的读-条件写-失败重读循环;已 aborted 重复调用保持幂等成功。 +2. **Minor:selection_hash 期望值自证**。`test_selection_hash_uses_pre_resolution_semantics` 原先用测试内重复实现的同一序列化逻辑计算期望值(两侧同变不红)。改为钉死离线算出的 SHA-256 字面值(`{"auxiliary":null,"primary":null}` → `697c0462...55abc`,常量 `INHERIT_SELECTION_HASH`),锁住对外契约;删除测试内的重复实现。 + +新增测试(先红后绿): +- `test_abort_run_snapshot_store_update_is_conditional`:store 层条件 UPDATE 的 rowcount 语义(prepared→True,重复/bound→False 且行不被改写)。 +- `test_abort_losing_bind_race_keeps_bound_state`:monkeypatch `get_run_snapshot` 在 abort 读与写之间插入并发 bind,断言 abort 抛 `SNAPSHOT_ALREADY_BOUND` 且行保持 bound、`langgraph_run_id` 完好(旧实现此测试必红)。 + +验证: +- `.venv/bin/python -m pytest tests/test_snapshots.py -x -q` → **42 passed** +- `.venv/bin/python -m pytest tests/ -x -q` → **3133 passed, 10 skipped**(无回归) +- `ruff check` / `ruff format --check`(涉及文件)→ 全净 diff --git a/EvoScientist/model_registry/snapshots.py b/EvoScientist/model_registry/snapshots.py index 6e1f003..a0b67d2 100644 --- a/EvoScientist/model_registry/snapshots.py +++ b/EvoScientist/model_registry/snapshots.py @@ -312,20 +312,28 @@ class SnapshotService: # Lost a state-transition race; re-read and apply the rules. def abort(self, snapshot_id: str) -> None: - """Mark a ``prepared`` snapshot ``aborted`` (run creation failed).""" - row = self._store.get_run_snapshot(snapshot_id) - if row is None: - raise _snapshot_not_found() - if row["status"] == "aborted": - return - if row["status"] == "bound": - raise ModelRegistryError( - SNAPSHOT_ALREADY_BOUND, - "A bound snapshot cannot be aborted.", - ) - if row["status"] == "expired": - raise _snapshot_expired() - self._store.set_run_snapshot_status(snapshot_id, "aborted") + """Mark a ``prepared`` snapshot ``aborted`` (run creation failed). + + The write is a conditional update: a bind that commits between the + read and the write wins the race, and the abort re-reads and raises + instead of overwriting the bound row. + """ + while True: + row = self._store.get_run_snapshot(snapshot_id) + if row is None: + raise _snapshot_not_found() + if row["status"] == "aborted": + return + if row["status"] == "bound": + raise ModelRegistryError( + SNAPSHOT_ALREADY_BOUND, + "A bound snapshot cannot be aborted.", + ) + if row["status"] == "expired": + raise _snapshot_expired() + if self._store.abort_run_snapshot(snapshot_id): + return + # Lost a state-transition race; re-read and apply the rules. def get( self, diff --git a/EvoScientist/model_registry/store.py b/EvoScientist/model_registry/store.py index 39db28a..5bee146 100644 --- a/EvoScientist/model_registry/store.py +++ b/EvoScientist/model_registry/store.py @@ -826,6 +826,30 @@ class ModelRuntimeStore: finally: connection.close() + def abort_run_snapshot(self, snapshot_id: str) -> bool: + """Abort a ``prepared`` snapshot; return False when it left that state. + + The conditional update keeps the prepared→aborted transition atomic + so a concurrent bind cannot be clobbered back to ``aborted`` (which + would lose its ``langgraph_run_id``). + """ + with self._lock: + connection = self._connect() + try: + connection.execute("BEGIN IMMEDIATE") + cursor = connection.execute( + "UPDATE run_runtime_snapshots SET status = 'aborted' " + "WHERE snapshot_id = ? AND status = 'prepared'", + (snapshot_id,), + ) + connection.commit() + return cursor.rowcount == 1 + except BaseException: + connection.rollback() + raise + finally: + connection.close() + def expire_due_run_snapshots(self, now: int) -> list[str]: """Mark every due prepared/bound snapshot ``expired`` (terminal). diff --git a/tests/test_snapshots.py b/tests/test_snapshots.py index 70a6111..1b6b8ba 100644 --- a/tests/test_snapshots.py +++ b/tests/test_snapshots.py @@ -8,7 +8,6 @@ revisions — with no in-process secret caching. from __future__ import annotations -import hashlib import json import time @@ -197,18 +196,11 @@ def _request(**overrides): return SnapshotCreateRequest.model_validate(payload) -def _expected_selection_hash(primary, auxiliary): - def entry(ref): - if ref is None: - return None - return {"provider_id": ref.provider_id, "model_key": ref.model_key} - - encoded = json.dumps( - {"primary": entry(primary), "auxiliary": entry(auxiliary)}, - sort_keys=True, - separators=(",", ":"), - ) - return hashlib.sha256(encoded.encode("utf-8")).hexdigest() +# 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: @@ -255,12 +247,13 @@ class TestCreate: def test_selection_hash_uses_pre_resolution_semantics(self, service): creation = service.create(_request()) snapshot = creation.snapshot - assert snapshot.selection_hash == _expected_selection_hash(None, None) - assert compute_selection_hash(None, None) == _expected_selection_hash( - None, None - ) + # 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) != snapshot.selection_hash + 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)) @@ -389,6 +382,58 @@ class TestAbort: 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):