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.
This commit is contained in:
m4
2026-07-21 08:45:23 +08:00
parent 0cc995eb80
commit c8c46eab16
4 changed files with 123 additions and 32 deletions
+14
View File
@@ -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`(涉及文件)→ 全净
+22 -14
View File
@@ -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,
+24
View File
@@ -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).
+63 -18
View File
@@ -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):