Files
EvoScientist-Multi/tests/test_checkpoint_writer_race.py
T
m4 d4b53bfb08
Docker / build (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.12) (push) Has been cancelled
Test / pytest (windows-latest, 3.11) (push) Has been cancelled
Test / pytest (windows-latest, 3.12) (push) Has been cancelled
Lint / ruff (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.11) (push) Has been cancelled
Build / build (push) Has been cancelled
test: cover stop contract, execution adapters, checkpointer race and runtime identity
2026-09-13 15:12:17 +08:00

158 lines
9.2 KiB
Python

"""Distinct cross-Turn writer exclusion scenario, real Graph/SQLite, offline."""
import asyncio
import hashlib
import socket
import uuid
from dataclasses import replace
from typing import TypedDict
import pytest
def test_cross_turn_checkpoint_writer_barrier(tmp_path, monkeypatch):
from EvoScientist.llm.contracts import AgentInputV3, WebHostContext, EvoRuntimeError, canonical_json_v1
from EvoScientist.llm.host_execution_registry import SQLiteHostRegistry
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
from langgraph.graph import StateGraph, START, END
from langgraph.types import Command, interrupt
from tests.test_web_model_runtime import _runtime, _preparation, _admission, _Sink
def forbidden(*args, **kwargs):
raise AssertionError("network forbidden")
monkeypatch.setattr(socket.socket, "connect", forbidden)
monkeypatch.setattr(socket, "getaddrinfo", forbidden)
# Bare StateGraph has no DeepAgents subagent event lane. Adapt only the
# event projection; execution, interrupts and persistence remain real.
async def graph_events(agent, message, thread_id, *, configurable, **kwargs):
from EvoScientist.stream.events import build_agent_stream_input
value = await build_agent_stream_input(message, media=None)
async for update in agent.astream(value, {"configurable": {
**configurable, "thread_id": thread_id}}, stream_mode="updates"):
yield {"type": "graph_update"}
monkeypatch.setattr("EvoScientist.stream.events.stream_agent_events", graph_events)
class State(TypedDict):
messages: list
_verified_review_mode: dict
result: str
consumed = []
async def review(state):
if state["messages"][-1]["content"] == "parallel":
return {"result": "parallel"}
decision = interrupt({"action_requests": [{"name": "probe", "args": {}}]})
consumed.append(decision)
return {"result": "approved"}
async def scenario():
runtime, authority = _runtime(tmp_path, monkeypatch)
peer_root = tmp_path / "peer"
peer_root.mkdir()
peer, _ = _runtime(peer_root, monkeypatch)
registry_path = tmp_path / "registry.sqlite"
runtime.host_registry = SQLiteHostRegistry(registry_path, host_id="host", boot_id="a")
peer.host_registry = SQLiteHostRegistry(registry_path, host_id="host", boot_id="b")
async with AsyncSqliteSaver.from_conn_string(str(tmp_path / "graph.sqlite")) as saver:
async with AsyncSqliteSaver.from_conn_string(str(tmp_path / "graph.sqlite")) as peer_saver:
builder = StateGraph(State)
builder.add_node("mode", lambda state: {"_verified_review_mode": {"mode": "manual"}})
builder.add_node("review", review)
builder.add_edge(START, "mode")
builder.add_edge("mode", "review")
builder.add_edge("review", END)
runtime.agent_factory = lambda snapshot, host, models: builder.compile(checkpointer=host.checkpointer)
peer.agent_factory = runtime.agent_factory
host = WebHostContext(str(tmp_path), str(tmp_path), object(), saver, runtime_event_sink=_Sink())
peer_host = replace(host, checkpointer=peer_saver, runtime_event_sink=_Sink())
def grant(value, turn, **changes):
template = _preparation(authority, value, title_policy="disabled")
fields = dict(template.unsigned_payload())
for key in ("issued_at", "expires_at", "key_id", "schema_version", "contract_type"):
fields.pop(key, None)
fields.update(request_id=str(uuid.uuid4()), turn_id=turn, **changes)
return authority.sign_preparation(**fields, ttl_ms=60000)
async def prepare(rt, h, value, turn, **changes):
return await rt.prepare_model_run(grant(value, turn, **changes), value, h)
value = AgentInputV3("approval", "user:checkpoint")
turn = str(uuid.uuid4())
parent_q = await prepare(runtime, host, value, turn)
parent = await runtime.start_web_run(_admission(authority, parent_q))
assert await parent.wait_stopped(timeout=5) == "awaiting_input"
payload = parent._terminal_event.payload
identity = {k: payload[k] for k in ("checkpoint_thread_id", "checkpoint_id", "checkpoint_ns", "pending_interrupts")}
approved = AgentInputV3(Command(resume={"decisions": [{"type": "approve"}]}), value.checkpoint_thread_id)
child_q = await prepare(runtime, host, approved, turn,
checkpoint_snapshot_id=payload["checkpoint_id"], predecessor_execution_id=parent.run_id,
predecessor_checkpoint_id=payload["checkpoint_id"], predecessor_owner_epoch=1,
continuation_pending_hash=hashlib.sha256(canonical_json_v1(identity)).hexdigest(),
continuation_decision_hash=hashlib.sha256(canonical_json_v1(approved.message.resume)).hexdigest())
other_q = await prepare(runtime, host, value, str(uuid.uuid4()))
peer_q = await prepare(peer, peer_host, value, str(uuid.uuid4()))
parallel_value = AgentInputV3("parallel", "user:other-checkpoint")
parallel_q = await prepare(peer, peer_host, parallel_value, str(uuid.uuid4()))
checked, release = asyncio.Event(), asyncio.Event()
closing, closed = asyncio.Event(), asyncio.Event()
class Client:
async def aclose(self):
closing.set()
await closed.wait()
factory = runtime.agent_factory
def owned_factory(*args):
from EvoScientist.llm.runtime import _construction_owner
owner = _construction_owner.get()
client = Client()
owner._owned_clients[id(client)] = client
return factory(*args)
runtime.agent_factory = owned_factory
validate = runtime._validate_continuation
async def barrier(g, v, h):
await validate(g, v, h)
if g.predecessor_execution_id:
checked.set()
await release.wait()
monkeypatch.setattr(runtime, "_validate_continuation", barrier)
starting = asyncio.create_task(runtime.start_web_run(_admission(authority, child_q)))
late = []
try:
await asyncio.wait_for(checked.wait(), 5)
parallel = await peer.start_web_run(_admission(authority, parallel_q))
assert await parallel.wait_stopped(timeout=5) == "completed"
for rt, quote in ((runtime, other_q), (peer, peer_q)):
try:
late.append(await rt.start_web_run(_admission(authority, quote)))
except EvoRuntimeError as exc:
assert "CHECKPOINT_WRITER_BUSY" in str(exc)
assert not late, "cross-Turn writers admitted after continuation latest validation"
finally:
for run in late:
await run.wait_stopped(timeout=5)
release.set()
child = await starting
await asyncio.wait_for(closing.wait(), 5)
try:
assert await child.wait_stopped(timeout=0.02) == "unknown"
with pytest.raises(EvoRuntimeError, match="CHECKPOINT_WRITER_BUSY"):
await peer.start_web_run(_admission(authority, peer_q))
finally:
closed.set()
await child.wait_stopped(timeout=5)
assert child._terminal_event.payload["outcome"] == "completed"
assert consumed == [{"decisions": [{"type": "approve"}]}]
import sqlite3
with sqlite3.connect(registry_path) as db:
assert db.execute("SELECT count(*) FROM checkpoint_writers").fetchone()[0] == 0
# A stale continuation never consumes a later Turn's pending.
newer = await peer.start_web_run(_admission(authority, peer_q))
assert await newer.wait_stopped(timeout=5) == "awaiting_input"
with pytest.raises(EvoRuntimeError, match="CONTINUATION_CHECKPOINT_STALE"):
await runtime.prepare_model_run(grant(approved, turn,
checkpoint_snapshot_id=payload["checkpoint_id"], predecessor_execution_id=parent.run_id,
predecessor_checkpoint_id=payload["checkpoint_id"], predecessor_owner_epoch=1,
continuation_pending_hash=hashlib.sha256(canonical_json_v1(identity)).hexdigest(),
continuation_decision_hash=hashlib.sha256(canonical_json_v1(approved.message.resume)).hexdigest()), approved, host)
print("cross-Turn same/peer runtime blocked; distinct checkpoint completes; pinned resume consumed once")
asyncio.run(scenario())