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
119 lines
6.6 KiB
Python
119 lines
6.6 KiB
Python
"""Offline real-v3 stop proofs; each fault has its own five-run budget."""
|
|
import asyncio
|
|
import sqlite3
|
|
import socket
|
|
import threading
|
|
|
|
import pytest
|
|
|
|
|
|
@pytest.mark.parametrize("fault", ["cancel", "close_error", "slow_commit"])
|
|
def test_real_v3_stop(tmp_path, monkeypatch, fault):
|
|
from langchain_core.language_models.fake_chat_models import FakeListChatModel
|
|
from langchain_core.messages import AIMessage
|
|
from langgraph.graph import StateGraph, MessagesState, START, END
|
|
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
|
|
from EvoScientist.llm.contracts import AgentInputV3, WebHostContext
|
|
from EvoScientist.llm.host_execution_registry import SQLiteHostRegistry
|
|
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)
|
|
|
|
async def scenario():
|
|
entered, settled = asyncio.Event(), asyncio.Event()
|
|
close_failed, close_release = asyncio.Event(), asyncio.Event()
|
|
release = threading.Event()
|
|
runtime, authority = _runtime(tmp_path, monkeypatch)
|
|
registry_path = tmp_path / "registry.sqlite"
|
|
runtime.host_registry = SQLiteHostRegistry(registry_path, host_id="host", boot_id="stop")
|
|
async with AsyncSqliteSaver.from_conn_string(str(tmp_path / "graph.sqlite")) as saver:
|
|
model = FakeListChatModel(responses=["controlled streaming response"], sleep=0.01)
|
|
async def model_node(state):
|
|
try:
|
|
async for chunk in model.astream(state["messages"]):
|
|
entered.set()
|
|
await asyncio.Event().wait()
|
|
return {"messages": [AIMessage(content="finished")]}
|
|
finally:
|
|
settled.set()
|
|
builder = StateGraph(MessagesState)
|
|
builder.add_node("model", model_node)
|
|
builder.add_edge(START, "model")
|
|
builder.add_edge("model", END)
|
|
from langchain.agents._subagent_transformer import SubagentTransformer
|
|
runtime.agent_factory = lambda snapshot, host, models: builder.compile(
|
|
checkpointer=saver, transformers=[SubagentTransformer])
|
|
sink = _Sink()
|
|
host = WebHostContext(str(tmp_path), str(tmp_path), object(), saver, runtime_event_sink=sink)
|
|
value = AgentInputV3("stop probe", "user:stop")
|
|
quote = await runtime.prepare_model_run(_preparation(authority, value, title_policy="disabled"), value, host)
|
|
run = await runtime.start_web_run(_admission(authority, quote))
|
|
queued = None
|
|
try:
|
|
await asyncio.wait_for(entered.wait(), 3)
|
|
if fault == "close_error":
|
|
# Instance-local dependency fault, not a replacement event pipeline.
|
|
raw = run._stream_stop.graph
|
|
class FailOnce:
|
|
failures = 0
|
|
def __aiter__(self): return self
|
|
async def __anext__(self): return await raw.__anext__()
|
|
async def aclose(self):
|
|
if not self.failures:
|
|
self.failures += 1
|
|
close_failed.set()
|
|
raise RuntimeError("controlled close failure")
|
|
await close_release.wait()
|
|
await raw.aclose()
|
|
run._stream_stop.graph = FailOnce()
|
|
if fault == "slow_commit":
|
|
await saver.conn.execute("CREATE TABLE stop_probe(value TEXT)")
|
|
await saver.conn.commit()
|
|
busy = threading.Event()
|
|
def blocked_write():
|
|
busy.set()
|
|
release.wait(5)
|
|
saver.conn._conn.execute("INSERT INTO stop_probe VALUES ('cancelled future executed')")
|
|
queued = asyncio.create_task(saver.conn._execute(blocked_write))
|
|
await asyncio.wait_for(asyncio.to_thread(busy.wait), 3)
|
|
queued.cancel()
|
|
await asyncio.gather(queued, return_exceptions=True)
|
|
run._request_cancel("probe")
|
|
if fault == "close_error":
|
|
await asyncio.wait_for(close_failed.wait(), 3)
|
|
assert await run.wait_stopped(timeout=0.02) == "unknown"
|
|
assert run._stream_stop.errors
|
|
assert not run._stream_stop.task.done()
|
|
with sqlite3.connect(registry_path) as db:
|
|
assert db.execute("SELECT count(*) FROM checkpoint_writers").fetchone()[0] == 1
|
|
close_release.set()
|
|
if fault == "slow_commit":
|
|
assert await run.wait_stopped(timeout=0.03) == "unknown"
|
|
with sqlite3.connect(registry_path) as db:
|
|
assert db.execute("SELECT count(*) FROM checkpoint_writers").fetchone()[0] == 1
|
|
print("slow_commit: cancelled queued future; unknown retains writer")
|
|
release.set()
|
|
assert await run.wait_stopped(timeout=3) == "cancelled"
|
|
assert settled.is_set()
|
|
assert run._checkpoint_writes_stopped
|
|
assert sum(e.payload.get("kind") == "run_terminal" for e in run._journal) == 1
|
|
with sqlite3.connect(registry_path) as db:
|
|
assert db.execute("SELECT count(*) FROM checkpoint_writers").fetchone()[0] == 0
|
|
if fault == "slow_commit":
|
|
with sqlite3.connect(tmp_path / "graph.sqlite") as db:
|
|
assert db.execute("SELECT value FROM stop_probe").fetchone()[0] == "cancelled future executed"
|
|
if fault == "close_error":
|
|
assert run._stream_stop.errors
|
|
print("close_error: retained failure, retried closure, released")
|
|
print(f"{fault}: real v3 model settled, terminal=1, writers=0")
|
|
print(f"{fault}: exit_tasks_observed={run._stream_stop.exit_tasks_observed}, pulls={len(run._stream_stop.pulls)}")
|
|
finally:
|
|
release.set()
|
|
close_release.set()
|
|
if run._terminal_event is None:
|
|
run._request_cancel("test cleanup")
|
|
await asyncio.gather(run.wait_stopped(timeout=3), return_exceptions=True)
|
|
asyncio.run(scenario()) |