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
158 lines
9.2 KiB
Python
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()) |