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:
@@ -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`(涉及文件)→ 全净
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user