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
334 lines
17 KiB
Python
334 lines
17 KiB
Python
"""B03 isolated history contract, real Graph and SQLite."""
|
|
import asyncio
|
|
import copy
|
|
import importlib.util
|
|
import socket
|
|
from dataclasses import replace
|
|
import pytest
|
|
|
|
|
|
@pytest.mark.parametrize("boundary", ["START", "END"])
|
|
def test_initialization_recovers_lost_write_response(tmp_path, boundary):
|
|
from EvoScientist.llm import history_rebuild as h
|
|
assert hasattr(h, "SqliteInitializationStore"), "durable initialization gate missing"
|
|
from langchain_core.messages import HumanMessage, AIMessage
|
|
from langgraph.graph import StateGraph, MessagesState, START, END
|
|
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
|
|
|
|
async def scenario():
|
|
scope = h.HistoryScope("a", "t", "recover", "w", "g", "tools", 1)
|
|
records = [h.HistoryRecord("m", 1, HumanMessage(content="history"))]
|
|
history = h.normalize_history(records, source=scope, expected=scope)
|
|
original = copy.deepcopy(records)
|
|
ran = []
|
|
def build(saver):
|
|
builder = StateGraph(MessagesState)
|
|
def model(state):
|
|
ran.append(state["messages"][-1].content)
|
|
return {"messages": [AIMessage(content="answer")]}
|
|
builder.add_node("model", model)
|
|
builder.add_edge(START, "model")
|
|
builder.add_edge("model", END)
|
|
return builder.compile(checkpointer=saver)
|
|
class LostResponse:
|
|
def __init__(self, graph):
|
|
self.graph = graph
|
|
def __getattr__(self, name):
|
|
return getattr(self.graph, name)
|
|
async def aupdate_state(self, *args, **kwargs):
|
|
result = await self.graph.aupdate_state(*args, **kwargs)
|
|
if kwargs["as_node"] == (START if boundary == "START" else END):
|
|
raise ConnectionError("write committed, response lost")
|
|
return result
|
|
db = str(tmp_path / "graph.sqlite")
|
|
store_path = str(tmp_path / "initialization.sqlite")
|
|
store = h.SqliteInitializationStore(store_path)
|
|
kwargs = dict(expected=scope, store=store, attempt="attempt-1", owner="host-1", fence=1)
|
|
async with AsyncSqliteSaver.from_conn_string(db) as saver:
|
|
graph = build(saver)
|
|
with pytest.raises(ConnectionError):
|
|
await h.create_history_checkpoint(LostResponse(graph), history, **kwargs)
|
|
assert store.get(h.history_key(scope))["status"] == "INITIALIZING"
|
|
with pytest.raises(ValueError, match="READY"):
|
|
await h.invoke_history_checkpoint(graph, scope=scope, store=store,
|
|
attempt="attempt-1", owner="host-1", fence=1, input={"messages": []})
|
|
assert ran == []
|
|
store = h.SqliteInitializationStore(store_path)
|
|
kwargs.update(store=store, owner="host-2", fence=2)
|
|
async with AsyncSqliteSaver.from_conn_string(db) as saver:
|
|
graph = build(saver)
|
|
config = await h.create_history_checkpoint(graph, history, **kwargs)
|
|
state = await graph.aget_state(config)
|
|
assert not state.tasks and not state.next
|
|
assert await h.create_history_checkpoint(graph, history, **kwargs) == config
|
|
with pytest.raises(ValueError):
|
|
await h.create_history_checkpoint(graph, history, **dict(kwargs, owner="host-1", fence=1))
|
|
changed = h.normalize_history([h.HistoryRecord("m", 1, HumanMessage(content="changed"))], source=scope, expected=scope)
|
|
with pytest.raises(ValueError):
|
|
await h.create_history_checkpoint(graph, changed, **kwargs)
|
|
with pytest.raises(ValueError):
|
|
await h.create_history_checkpoint(graph, history, **dict(kwargs, attempt="other"))
|
|
result = await h.invoke_history_checkpoint(graph, scope=scope, store=store,
|
|
attempt="attempt-1", owner="host-2", fence=2,
|
|
input={"messages": [HumanMessage(content="new")]})
|
|
assert result["messages"][-1].content == "answer" and ran == ["new"]
|
|
with pytest.raises(ValueError):
|
|
await h.create_history_checkpoint(graph, history, **kwargs)
|
|
assert records == original
|
|
asyncio.run(scenario())
|
|
|
|
|
|
@pytest.mark.parametrize("kind", ["wrong-name", "missing-id", "duplicate-id", "orphan"])
|
|
def test_tool_association_conflicts_are_rejected(kind):
|
|
from EvoScientist.llm.history_rebuild import HistoryScope, HistoryRecord, normalize_history
|
|
from langchain_core.messages import AIMessage, ToolMessage
|
|
scope = HistoryScope("a", "t", "turn", "w", "g", "tools", 1)
|
|
calls = [{"id": "stable", "name": "A", "args": {}}]
|
|
if kind == "duplicate-id":
|
|
calls.append({"id": "stable", "name": "B", "args": {}})
|
|
messages = [AIMessage(content="", tool_calls=calls), ToolMessage(content="result B",
|
|
tool_call_id="" if kind == "missing-id" else "unknown" if kind == "orphan" else "stable",
|
|
name="B" if kind in ("wrong-name", "missing-id") else "A")]
|
|
records = [HistoryRecord(str(i), 1, m) for i, m in enumerate(messages)]
|
|
original = copy.deepcopy(records)
|
|
with pytest.raises(ValueError, match="tool"):
|
|
normalize_history(records, source=scope, expected=scope)
|
|
assert records == original
|
|
|
|
|
|
def test_new_turn_uses_fresh_checkpoint(tmp_path, monkeypatch):
|
|
def no_network(*args, **kwargs):
|
|
raise AssertionError("network forbidden")
|
|
monkeypatch.setattr(socket.socket, "connect", no_network)
|
|
assert importlib.util.find_spec("EvoScientist.llm.history_rebuild"), "missing history builder"
|
|
from EvoScientist.llm.history_rebuild import (HistoryScope, HistoryRecord, normalize_history,
|
|
create_history_checkpoint, invoke_history_checkpoint, SqliteInitializationStore)
|
|
from EvoScientist.llm.patches import _validate_openai_tool_history
|
|
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
|
from langgraph.graph import StateGraph, MessagesState, START, END
|
|
from langgraph.types import interrupt
|
|
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
|
|
|
|
scope = HistoryScope("alice", "thread", "new-turn", "workspace", "g1", "t1", 7)
|
|
calls = [{"name": "probe", "args": {}, "id": k} for k in ("complete", "missing")]
|
|
records = [
|
|
HistoryRecord("human", 1, HumanMessage(content="old question")),
|
|
HistoryRecord("partial", 2, AIMessage(content="partial answer", tool_calls=calls), partial=True),
|
|
HistoryRecord("result", 1, ToolMessage(content="real result", tool_call_id="complete")),
|
|
HistoryRecord("media", 1, HumanMessage(content=[
|
|
{"type": "text", "text": "attachment"},
|
|
{"type": "image_url", "image_url": {"url": "data:image/png;base64,SECRET"}},
|
|
]), file_refs=("file:owned-document",)),
|
|
]
|
|
original = copy.deepcopy(records)
|
|
history = normalize_history(records, source=scope, expected=scope)
|
|
assert records == original
|
|
_validate_openai_tool_history(list(history.messages))
|
|
tools = [m for m in history.messages if isinstance(m, ToolMessage)]
|
|
assert len(tools) == 1 and tools[0].content == "real result"
|
|
ai = [m for m in history.messages if isinstance(m, AIMessage) and m.tool_calls]
|
|
assert [c["id"] for c in ai[0].tool_calls] == ["complete"]
|
|
text = str(history.messages)
|
|
assert "SECRET" not in text and "base64" not in text
|
|
assert "file:owned-document" in text
|
|
assert "partial" in text and "missing" in text and "incomplete" in text
|
|
for field, value in (("tenant_id", "bob"), ("workspace_id", "other"),
|
|
("graph_version", "g2"), ("tool_version", "t2"), ("history_revision", 8)):
|
|
with pytest.raises(ValueError):
|
|
normalize_history(records, source=scope, expected=replace(scope, **{field: value}))
|
|
for bad in ([records[0], records[0]], [replace(records[0], revision=-1)]):
|
|
with pytest.raises(ValueError):
|
|
normalize_history(bad, source=scope, expected=scope)
|
|
executed = []
|
|
def build(saver):
|
|
graph = StateGraph(MessagesState)
|
|
def model(state):
|
|
last = state["messages"][-1].content
|
|
if last == "old pending":
|
|
interrupt({"tool": "old-danger", "pending": True})
|
|
executed.append("OLD TOOL")
|
|
executed.append(last)
|
|
return {"messages": [AIMessage(content="new answer")]}
|
|
graph.add_node("model", model)
|
|
graph.add_edge(START, "model")
|
|
graph.add_edge("model", END)
|
|
return graph.compile(checkpointer=saver)
|
|
|
|
async def scenario():
|
|
db = str(tmp_path / "history.sqlite")
|
|
store = SqliteInitializationStore(str(tmp_path / "initialization.sqlite"))
|
|
gate = dict(store=store, attempt="original", owner="host", fence=1)
|
|
old = {"configurable": {"thread_id": "legacy-pending"}}
|
|
async with AsyncSqliteSaver.from_conn_string(db) as saver:
|
|
graph = build(saver)
|
|
await graph.ainvoke({"messages": [HumanMessage(content="old pending")]}, old)
|
|
before = await graph.aget_state(old)
|
|
assert before.next == ("model",) and before.tasks[0].interrupts
|
|
old_tuple = await saver.aget_tuple(old)
|
|
async with AsyncSqliteSaver.from_conn_string(db) as saver:
|
|
graph = build(saver)
|
|
config = await create_history_checkpoint(graph, history, expected=scope, **gate)
|
|
fresh = await graph.aget_state(config)
|
|
assert not fresh.next and not fresh.tasks and executed == []
|
|
with pytest.raises(ValueError):
|
|
await create_history_checkpoint(graph, history, expected=scope, **dict(gate, attempt="other"))
|
|
result = await invoke_history_checkpoint(graph, scope=scope, **gate,
|
|
input={"messages": [HumanMessage(content="new input")]})
|
|
assert result["messages"][-1].content == "new answer"
|
|
bob = replace(scope, tenant_id="bob")
|
|
bob_history = normalize_history([], source=bob, expected=bob)
|
|
bob_config = await create_history_checkpoint(graph, bob_history, expected=bob, **gate)
|
|
assert bob_config["configurable"]["thread_id"] != config["configurable"]["thread_id"]
|
|
assert not (await graph.aget_state(bob_config)).values.get("messages")
|
|
async with AsyncSqliteSaver.from_conn_string(db) as saver:
|
|
graph = build(saver)
|
|
assert await graph.aget_state(old) == before
|
|
assert await saver.aget_tuple(old) == old_tuple
|
|
assert (await graph.aget_state(config)).values["messages"][-1].content == "new answer"
|
|
assert executed == ["new input"]
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_committed_v2_history_replaces_checkpoint_without_replaying_tools():
|
|
from typing import Any, cast
|
|
from EvoScientist.llm.history_rebuild import committed_history_input
|
|
from langchain_core.messages import AIMessage, convert_to_messages
|
|
from langgraph.graph.message import add_messages
|
|
|
|
history = {
|
|
"schema": "ai4sci.committed-history.v1", "thread_id": "thread",
|
|
"user_uid": "user", "conversation_revision": 2, "excluded_message_id": "new",
|
|
"records": [
|
|
{"message_id": "u", "message_index": 1, "revision": 1,
|
|
"role": "user", "payload": {"content": "old question"}},
|
|
{"message_id": "a", "message_index": 2, "revision": 1,
|
|
"role": "assistant", "payload": {"items": [
|
|
{"item_id": "answer", "item_sequence": 1, "type": "message",
|
|
"content": [{"type": "output_text", "text": "old answer"}]},
|
|
{"item_id": "tool", "item_sequence": 2, "type": "tool_call",
|
|
"name": "search", "input": {"q": "old"}, "status": "failed"},
|
|
]}},
|
|
],
|
|
}
|
|
result = committed_history_input(history, {"messages": [{"role": "user", "content": "next"}]},
|
|
thread_id="thread", run_id="run")
|
|
reduce_messages = cast(Any, add_messages)
|
|
messages = reduce_messages([AIMessage(content="stale checkpoint", id="stale")],
|
|
convert_to_messages(result["messages"]))
|
|
assert [m.id for m in messages] == ["u", "a", "current:run:0"]
|
|
assert "historical tool_call" in messages[1].content
|
|
assert not messages[1].tool_calls
|
|
assert messages[-1].content == "next"
|
|
assert result["_summarization_event"] is None
|
|
assert len(reduce_messages(messages, convert_to_messages(result["messages"]))) == 3
|
|
|
|
# A compatible checkpoint is authoritative after the first initialization.
|
|
# Its tool protocol and summarization indexes must survive the next turn.
|
|
current = {"messages": [{"role": "user", "content": "next"}]}
|
|
original = copy.deepcopy(current)
|
|
appended = committed_history_input(history, current, thread_id="thread", run_id="run",
|
|
checkpoint_exists=True)
|
|
checkpoint = [AIMessage(content="checkpoint answer", id="checkpoint",
|
|
tool_calls=[{"id": "tool-1", "name": "search", "args": {}}])]
|
|
continued = reduce_messages(checkpoint, convert_to_messages(appended["messages"]))
|
|
assert [m.id for m in continued] == ["checkpoint", "current:run:0"]
|
|
assert continued[0].tool_calls[0]["id"] == "tool-1"
|
|
assert "_summarization_event" not in appended
|
|
assert current == original
|
|
|
|
|
|
@pytest.mark.parametrize("existing,pending,operation,history,legacy,expected", [
|
|
(False, False, "start", True, False, "initialize"),
|
|
(True, False, "start", True, False, "append"),
|
|
(True, True, "start", True, False, "THREAD_AWAITING_INPUT"),
|
|
(True, True, "resume", False, False, "resume"),
|
|
(False, False, "resume", False, True, "CHECKPOINT_RESUME_UNAVAILABLE"),
|
|
(False, False, "start", False, True, "HISTORY_REQUIRED"),
|
|
(False, False, "start", False, False, "new"),
|
|
])
|
|
def test_runtime_history_admission(existing, pending, operation, history, legacy, expected,
|
|
monkeypatch):
|
|
from EvoScientist.langgraph_dev import http
|
|
|
|
async def checkpoint(*args):
|
|
return existing, pending
|
|
async def old_history(*args):
|
|
return legacy
|
|
monkeypatch.setattr(http, "_compatible_checkpoint", checkpoint, raising=False)
|
|
monkeypatch.setattr(http, "_has_legacy_history", old_history, raising=False)
|
|
result = asyncio.run(http._history_admission(None, "thread", "EvoScientist", {},
|
|
operation, {} if history else None))
|
|
assert result == expected
|
|
|
|
|
|
def test_runtime_reads_checkpoint_through_real_async_factory(monkeypatch):
|
|
from contextlib import asynccontextmanager
|
|
from langgraph.checkpoint.base import empty_checkpoint
|
|
from langgraph.checkpoint.memory import InMemorySaver
|
|
from langgraph.graph import StateGraph, MessagesState, START, END
|
|
# Installed API factory dispatch, without startup or network services.
|
|
monkeypatch.setenv("REDIS_URI", "redis://unused")
|
|
monkeypatch.setenv("LANGGRAPH_RUNTIME_VARIANT", "inmem")
|
|
from langgraph_api import graph as api_graph, _checkpointer
|
|
from langgraph_api import store as api_store
|
|
from EvoScientist.langgraph_dev import http
|
|
|
|
saver = InMemorySaver()
|
|
entered = []
|
|
@asynccontextmanager
|
|
async def factory(config):
|
|
entered.append("enter")
|
|
builder = StateGraph(MessagesState)
|
|
builder.add_node("model", lambda state: {})
|
|
builder.add_edge(START, "model")
|
|
builder.add_edge("model", END)
|
|
try:
|
|
yield builder.compile()
|
|
finally:
|
|
entered.append("exit")
|
|
monkeypatch.setitem(api_graph.GRAPHS, "history-contract", factory)
|
|
api_graph.classify_factory(factory, "history-contract")
|
|
async def get_saver(**kwargs):
|
|
return saver
|
|
async def get_store():
|
|
return None
|
|
monkeypatch.setattr(_checkpointer, "get_checkpointer", get_saver)
|
|
monkeypatch.setattr(api_store, "get_store", get_store)
|
|
|
|
async def scenario():
|
|
config = {"configurable": {"thread_id": "thread", "checkpoint_ns": ""}}
|
|
assert await http._compatible_checkpoint(None, "thread", "history-contract", {}) == (False, False)
|
|
cp = empty_checkpoint()
|
|
cp["channel_values"] = {"messages": []}
|
|
cp["channel_versions"] = {"messages": 1}
|
|
await saver.aput(config, cp, {"source": "update", "step": 0,
|
|
"parents": {}, "graph_id": "history-contract"}, {"messages": 1})
|
|
assert await http._compatible_checkpoint(None, "thread", "history-contract", {}) == (True, False)
|
|
assert entered == ["enter", "exit"]
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_legacy_sqlite_history_is_detected_read_only(tmp_path, monkeypatch):
|
|
import sqlite3
|
|
import sys
|
|
import types
|
|
from EvoScientist.langgraph_dev import http
|
|
path = tmp_path / "sessions.db"
|
|
with sqlite3.connect(path) as conn:
|
|
conn.execute("CREATE TABLE checkpoints(thread_id TEXT)")
|
|
conn.execute("INSERT INTO checkpoints VALUES ('old')")
|
|
before = path.read_bytes()
|
|
monkeypatch.setenv("EVOSCIENTIST_LEGACY_CHECKPOINT_PATHS", __import__('json').dumps([str(path)]))
|
|
class Threads:
|
|
@staticmethod
|
|
async def get(*args):
|
|
from starlette.exceptions import HTTPException
|
|
raise HTTPException(404)
|
|
monkeypatch.setitem(sys.modules, "langgraph_runtime.ops", types.SimpleNamespace(Threads=Threads))
|
|
async def scenario():
|
|
assert await http._has_legacy_history(None, "old")
|
|
assert not await http._has_legacy_history(None, "new")
|
|
asyncio.run(scenario())
|
|
assert before == path.read_bytes()
|
|
|