* feat: implement bulk cancellation of non-terminal runs before thread deletion * test: enhance thread cancellation tests and add fake restore for orphaned runs sweep * feat: enhance run cancellation logic to support status filtering during thread deletion * feat: add langgraph-sdk dependency for enhanced functionality
This commit is contained in:
@@ -19,7 +19,7 @@ import asyncio
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable, Mapping
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Protocol, TypedDict
|
||||
|
||||
@@ -75,6 +75,12 @@ class _SyncRunsClient(Protocol):
|
||||
|
||||
def get(self, thread_id: str, run_id: str) -> Run: ...
|
||||
|
||||
def list(
|
||||
self, thread_id: str, *, limit: int, offset: int, status: str
|
||||
) -> list[Run]: ...
|
||||
|
||||
def cancel_many(self, *, thread_id: str, run_ids: Sequence[str]) -> object: ...
|
||||
|
||||
|
||||
class SyncLangGraphClient(Protocol):
|
||||
"""Sync subset of the LangGraph SDK used by background runs."""
|
||||
@@ -107,6 +113,14 @@ class _AsyncRunsClient(Protocol):
|
||||
|
||||
async def get(self, thread_id: str, run_id: str) -> Run: ...
|
||||
|
||||
async def list(
|
||||
self, thread_id: str, *, limit: int, offset: int, status: str
|
||||
) -> list[Run]: ...
|
||||
|
||||
async def cancel_many(
|
||||
self, *, thread_id: str, run_ids: Sequence[str]
|
||||
) -> object: ...
|
||||
|
||||
|
||||
class AsyncLangGraphClient(Protocol):
|
||||
"""Async subset of the LangGraph SDK used by background runs."""
|
||||
@@ -246,12 +260,97 @@ async def _aget_run_status(
|
||||
return run["status"]
|
||||
|
||||
|
||||
# Page size for enumerating a thread's runs before deletion. The SDK's
|
||||
# ``runs.list`` defaults to limit=10, which would silently skip runs on
|
||||
# threads with a longer history.
|
||||
_RUN_CANCEL_PAGE_SIZE = 100
|
||||
|
||||
# Statuses worth cancelling; listed server-side so terminal history is
|
||||
# never paged through.
|
||||
_CANCELABLE_RUN_STATUSES = ("pending", "running")
|
||||
|
||||
|
||||
def _cancel_thread_runs(
|
||||
client: SyncLangGraphClient,
|
||||
thread_id: str,
|
||||
*,
|
||||
name: str,
|
||||
) -> None:
|
||||
"""Best-effort interrupt of the thread's pending/running runs.
|
||||
|
||||
The server's ``threads.delete`` cascade-removes queued runs from the
|
||||
registry, but it does not interrupt a run that is already executing —
|
||||
cancelling first sends the interrupt control message so in-flight work
|
||||
actually stops (issue #358). It also protects cleanup paths that
|
||||
mutate the registry without going through ``threads.delete``. The bulk
|
||||
cancel is skipped when nothing is cancellable (the server 404s on an
|
||||
empty cancel set), which keeps the common terminal-only path to two
|
||||
cheap filtered GETs.
|
||||
"""
|
||||
try:
|
||||
run_ids: list[str] = []
|
||||
for status in _CANCELABLE_RUN_STATUSES:
|
||||
offset = 0
|
||||
while True:
|
||||
page = client.runs.list(
|
||||
thread_id,
|
||||
limit=_RUN_CANCEL_PAGE_SIZE,
|
||||
offset=offset,
|
||||
status=status,
|
||||
)
|
||||
run_ids.extend(run["run_id"] for run in page)
|
||||
if len(page) < _RUN_CANCEL_PAGE_SIZE:
|
||||
break
|
||||
offset += _RUN_CANCEL_PAGE_SIZE
|
||||
if run_ids:
|
||||
client.runs.cancel_many(
|
||||
thread_id=thread_id, run_ids=list(dict.fromkeys(run_ids))
|
||||
)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to cancel %s runs on thread %s", name, thread_id, exc_info=True
|
||||
)
|
||||
|
||||
|
||||
async def _acancel_thread_runs(
|
||||
client: AsyncLangGraphClient,
|
||||
thread_id: str,
|
||||
*,
|
||||
name: str,
|
||||
) -> None:
|
||||
"""Async variant of :func:`_cancel_thread_runs`."""
|
||||
try:
|
||||
run_ids: list[str] = []
|
||||
for status in _CANCELABLE_RUN_STATUSES:
|
||||
offset = 0
|
||||
while True:
|
||||
page = await client.runs.list(
|
||||
thread_id,
|
||||
limit=_RUN_CANCEL_PAGE_SIZE,
|
||||
offset=offset,
|
||||
status=status,
|
||||
)
|
||||
run_ids.extend(run["run_id"] for run in page)
|
||||
if len(page) < _RUN_CANCEL_PAGE_SIZE:
|
||||
break
|
||||
offset += _RUN_CANCEL_PAGE_SIZE
|
||||
if run_ids:
|
||||
await client.runs.cancel_many(
|
||||
thread_id=thread_id, run_ids=list(dict.fromkeys(run_ids))
|
||||
)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to cancel %s runs on thread %s", name, thread_id, exc_info=True
|
||||
)
|
||||
|
||||
|
||||
def _delete_thread(
|
||||
client: SyncLangGraphClient,
|
||||
thread_id: str,
|
||||
*,
|
||||
name: str,
|
||||
) -> None:
|
||||
_cancel_thread_runs(client, thread_id, name=name)
|
||||
try:
|
||||
client.threads.delete(thread_id)
|
||||
except Exception:
|
||||
@@ -264,6 +363,7 @@ async def _adelete_thread(
|
||||
*,
|
||||
name: str,
|
||||
) -> None:
|
||||
await _acancel_thread_runs(client, thread_id, name=name)
|
||||
try:
|
||||
await client.threads.delete(thread_id)
|
||||
except Exception:
|
||||
|
||||
@@ -28,6 +28,7 @@ from ..stream.events import (
|
||||
)
|
||||
from ..stream.summarization import _find_summarization_event_payload
|
||||
from ..stream.v3_payloads import _as_raw_map, _event_namespace
|
||||
from .background_runs import _acancel_thread_runs
|
||||
from .types import (
|
||||
DEFAULT_GRAPH_ID,
|
||||
GraphEvent,
|
||||
@@ -324,6 +325,10 @@ class LangGraphServerThreadStore(ThreadStore):
|
||||
return True
|
||||
|
||||
async def delete_thread(self, thread_id: str) -> bool:
|
||||
# Interrupt live runs first: the server's cascade delete clears
|
||||
# queued runs from the registry but does not stop a run that is
|
||||
# already executing (issue #358).
|
||||
await _acancel_thread_runs(self.client, thread_id, name="thread delete")
|
||||
try:
|
||||
await self.client.threads.delete(thread_id)
|
||||
except NotFoundError:
|
||||
|
||||
@@ -32,7 +32,7 @@ import atexit
|
||||
import logging
|
||||
import math
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable, MutableMapping
|
||||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime
|
||||
@@ -1532,7 +1532,7 @@ class _RestoredThreadInfo:
|
||||
model: str | None
|
||||
|
||||
|
||||
async def _restore_webui_threads_to_global_store() -> None:
|
||||
async def _restore_webui_threads_to_global_store() -> bool:
|
||||
"""Re-populate ``GlobalStore["threads"]`` from SQLite on server startup.
|
||||
|
||||
The inmem runtime's thread registry lives in memory (pickled to
|
||||
@@ -1554,7 +1554,10 @@ async def _restore_webui_threads_to_global_store() -> None:
|
||||
``workspace_dir`` are excluded.
|
||||
|
||||
Best-effort: any exception is logged and swallowed so a broken restore
|
||||
never prevents the ``langgraph dev`` server from starting.
|
||||
never prevents the ``langgraph dev`` server from starting. Returns True
|
||||
when the restore completed, False when it failed or the runtime is
|
||||
unavailable — a partially-restored registry must not be swept for
|
||||
orphans, or runs of not-yet-appended valid threads would be dropped.
|
||||
"""
|
||||
try:
|
||||
from langgraph_runtime_inmem.database import (
|
||||
@@ -1562,7 +1565,7 @@ async def _restore_webui_threads_to_global_store() -> None:
|
||||
)
|
||||
except ImportError:
|
||||
# langgraph_runtime_inmem not available (unit tests, plain CLI mode).
|
||||
return
|
||||
return False
|
||||
|
||||
def _to_uuid_safe(v: Any) -> uuid.UUID | None:
|
||||
try:
|
||||
@@ -1769,6 +1772,64 @@ async def _restore_webui_threads_to_global_store() -> None:
|
||||
"empty until new threads are created.",
|
||||
exc_info=True,
|
||||
)
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _sweep_orphaned_global_store_entries(store: MutableMapping[str, Any]) -> int:
|
||||
"""Drop registry runs and crons whose thread no longer exists (issue #358).
|
||||
|
||||
The inmem runtime never discards a run whose thread was deleted:
|
||||
``Runs.next`` reschedules it forever, and ``.langgraph_ops.pckl``
|
||||
persistence reloads it on every start, so zombies accumulate until the
|
||||
queue starves. Thread-bound crons (``crons.create_for_thread``) keep
|
||||
firing against their deleted thread the same way; stateless crons
|
||||
(``thread_id is None``) are never touched. Must run after the
|
||||
thread-registry restore above so the surviving thread set is final.
|
||||
Mutates the store lists in place — the runtime holds references to the
|
||||
same lists.
|
||||
"""
|
||||
valid_thread_ids = {
|
||||
str(entry.get("thread_id")) for entry in store.get("threads") or []
|
||||
}
|
||||
removed = 0
|
||||
runs: list[dict[str, Any]] | None = store.get("runs")
|
||||
if runs:
|
||||
before = len(runs)
|
||||
runs[:] = [run for run in runs if str(run.get("thread_id")) in valid_thread_ids]
|
||||
removed += before - len(runs)
|
||||
crons: list[dict[str, Any]] | None = store.get("crons")
|
||||
if crons:
|
||||
before = len(crons)
|
||||
crons[:] = [
|
||||
cron
|
||||
for cron in crons
|
||||
if cron.get("thread_id") is None
|
||||
or str(cron.get("thread_id")) in valid_thread_ids
|
||||
]
|
||||
removed += before - len(crons)
|
||||
return removed
|
||||
|
||||
|
||||
async def _sweep_orphaned_runs_in_global_store() -> None:
|
||||
"""Best-effort orphaned run/cron sweep against the live LangGraph registry."""
|
||||
try:
|
||||
from langgraph_runtime_inmem.database import (
|
||||
GLOBAL_STORE,
|
||||
)
|
||||
except ImportError:
|
||||
# langgraph_runtime_inmem not available (unit tests, plain CLI mode).
|
||||
return
|
||||
try:
|
||||
removed = _sweep_orphaned_global_store_entries(GLOBAL_STORE)
|
||||
if removed:
|
||||
_logger.info(
|
||||
"Dropped %d orphaned run(s)/cron(s) whose thread no longer "
|
||||
"exists from the LangGraph registry.",
|
||||
removed,
|
||||
)
|
||||
except Exception:
|
||||
_logger.warning("Orphaned-run sweep failed (non-fatal).", exc_info=True)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
@@ -1802,5 +1863,6 @@ async def create_checkpointer_for_langgraph_api() -> AsyncIterator[PruningCheckp
|
||||
) as saver:
|
||||
await saver.setup()
|
||||
await _purge_internal_worker_threads()
|
||||
await _restore_webui_threads_to_global_store()
|
||||
if await _restore_webui_threads_to_global_store():
|
||||
await _sweep_orphaned_runs_in_global_store()
|
||||
yield saver
|
||||
|
||||
@@ -33,6 +33,7 @@ dependencies = [
|
||||
"typer>=0.24",
|
||||
"python-dotenv>=1.0",
|
||||
"langgraph-cli[inmem]>=0.4",
|
||||
"langgraph-sdk>=0.4",
|
||||
"langgraph-checkpoint-sqlite>=3.0",
|
||||
"httpx>=0.28",
|
||||
"psutil>=6.0",
|
||||
|
||||
@@ -306,6 +306,196 @@ async def test_async_status_watcher_aborts_and_deletes_thread_on_error_status():
|
||||
assert deleted == ["thread-1"]
|
||||
|
||||
|
||||
def test_delete_thread_bulk_cancels_nonterminal_runs_before_delete():
|
||||
events: list[tuple] = []
|
||||
|
||||
class _Runs:
|
||||
def list(self, thread_id: str, *, limit: int, offset: int, status: str):
|
||||
events.append(("list", thread_id, status, offset))
|
||||
if status == "pending":
|
||||
return [{"run_id": "run-pending", "status": "pending"}]
|
||||
return []
|
||||
|
||||
def cancel_many(self, *, thread_id: str, run_ids):
|
||||
events.append(("cancel_many", thread_id, list(run_ids)))
|
||||
|
||||
class _Threads:
|
||||
def delete(self, thread_id: str):
|
||||
events.append(("delete", thread_id))
|
||||
|
||||
background_runs._delete_thread(
|
||||
SimpleNamespace(runs=_Runs(), threads=_Threads()),
|
||||
"thread-1",
|
||||
name="test worker",
|
||||
)
|
||||
|
||||
assert events == [
|
||||
("list", "thread-1", "pending", 0),
|
||||
("list", "thread-1", "running", 0),
|
||||
("cancel_many", "thread-1", ["run-pending"]),
|
||||
("delete", "thread-1"),
|
||||
]
|
||||
|
||||
|
||||
def test_delete_thread_skips_cancel_when_all_runs_terminal():
|
||||
events: list[tuple] = []
|
||||
|
||||
class _Runs:
|
||||
def list(self, thread_id: str, *, limit: int, offset: int, status: str):
|
||||
events.append(("list", thread_id, status, offset))
|
||||
return []
|
||||
|
||||
def cancel_many(self, **_kwargs): # pragma: no cover
|
||||
raise AssertionError("cancel_many must not be called")
|
||||
|
||||
class _Threads:
|
||||
def delete(self, thread_id: str):
|
||||
events.append(("delete", thread_id))
|
||||
|
||||
background_runs._delete_thread(
|
||||
SimpleNamespace(runs=_Runs(), threads=_Threads()),
|
||||
"thread-1",
|
||||
name="test worker",
|
||||
)
|
||||
|
||||
assert events == [
|
||||
("list", "thread-1", "pending", 0),
|
||||
("list", "thread-1", "running", 0),
|
||||
("delete", "thread-1"),
|
||||
]
|
||||
|
||||
|
||||
def test_cancel_thread_runs_paginates_past_first_page(monkeypatch):
|
||||
monkeypatch.setattr(background_runs, "_RUN_CANCEL_PAGE_SIZE", 2)
|
||||
cancelled: list[list[str]] = []
|
||||
pages = {
|
||||
("pending", 0): [
|
||||
{"run_id": "p1", "status": "pending"},
|
||||
{"run_id": "p2", "status": "pending"},
|
||||
],
|
||||
("pending", 2): [{"run_id": "p3", "status": "pending"}],
|
||||
("running", 0): [{"run_id": "r1", "status": "running"}],
|
||||
}
|
||||
|
||||
class _Runs:
|
||||
def list(self, thread_id: str, *, limit: int, offset: int, status: str):
|
||||
assert limit == 2
|
||||
return pages.get((status, offset), [])
|
||||
|
||||
def cancel_many(self, *, thread_id: str, run_ids):
|
||||
cancelled.append(list(run_ids))
|
||||
|
||||
background_runs._cancel_thread_runs(
|
||||
SimpleNamespace(runs=_Runs()),
|
||||
"thread-1",
|
||||
name="test worker",
|
||||
)
|
||||
|
||||
assert cancelled == [["p1", "p2", "p3", "r1"]]
|
||||
|
||||
|
||||
def test_delete_thread_still_deletes_when_run_listing_fails():
|
||||
deleted: list[str] = []
|
||||
|
||||
class _Runs:
|
||||
def list(self, thread_id: str, *, limit: int, offset: int, status: str):
|
||||
raise RuntimeError("listing failed")
|
||||
|
||||
def cancel_many(self, **_kwargs): # pragma: no cover
|
||||
raise AssertionError("cancel_many should not be reached")
|
||||
|
||||
class _Threads:
|
||||
def delete(self, thread_id: str):
|
||||
deleted.append(thread_id)
|
||||
|
||||
background_runs._delete_thread(
|
||||
SimpleNamespace(runs=_Runs(), threads=_Threads()),
|
||||
"thread-1",
|
||||
name="test worker",
|
||||
)
|
||||
|
||||
assert deleted == ["thread-1"]
|
||||
|
||||
|
||||
async def test_adelete_thread_bulk_cancels_nonterminal_runs_before_delete():
|
||||
events: list[tuple] = []
|
||||
|
||||
class _Runs:
|
||||
async def list(self, thread_id: str, *, limit: int, offset: int, status: str):
|
||||
events.append(("list", thread_id, status, offset))
|
||||
if status == "running":
|
||||
return [{"run_id": "run-running", "status": "running"}]
|
||||
return []
|
||||
|
||||
async def cancel_many(self, *, thread_id: str, run_ids):
|
||||
events.append(("cancel_many", thread_id, list(run_ids)))
|
||||
|
||||
class _Threads:
|
||||
async def delete(self, thread_id: str):
|
||||
events.append(("delete", thread_id))
|
||||
|
||||
await background_runs._adelete_thread(
|
||||
SimpleNamespace(runs=_Runs(), threads=_Threads()),
|
||||
"thread-1",
|
||||
name="test worker",
|
||||
)
|
||||
|
||||
assert events == [
|
||||
("list", "thread-1", "pending", 0),
|
||||
("list", "thread-1", "running", 0),
|
||||
("cancel_many", "thread-1", ["run-running"]),
|
||||
("delete", "thread-1"),
|
||||
]
|
||||
|
||||
|
||||
async def test_adelete_thread_still_deletes_when_run_listing_fails():
|
||||
deleted: list[str] = []
|
||||
|
||||
class _Runs:
|
||||
async def list(self, thread_id: str, *, limit: int, offset: int, status: str):
|
||||
raise RuntimeError("listing failed")
|
||||
|
||||
async def cancel_many(self, **_kwargs): # pragma: no cover
|
||||
raise AssertionError("cancel_many should not be reached")
|
||||
|
||||
class _Threads:
|
||||
async def delete(self, thread_id: str):
|
||||
deleted.append(thread_id)
|
||||
|
||||
await background_runs._adelete_thread(
|
||||
SimpleNamespace(runs=_Runs(), threads=_Threads()),
|
||||
"thread-1",
|
||||
name="test worker",
|
||||
)
|
||||
|
||||
assert deleted == ["thread-1"]
|
||||
|
||||
|
||||
def test_launch_cancels_stray_run_before_thread_delete_when_run_creation_fails(
|
||||
monkeypatch,
|
||||
):
|
||||
fake_client = _install_sync_launcher(
|
||||
monkeypatch,
|
||||
run_create_error=RuntimeError("run creation failed"),
|
||||
)
|
||||
fake_client.runs.list.side_effect = lambda _thread_id, **kwargs: (
|
||||
[{"run_id": "stray-run", "status": "pending"}]
|
||||
if kwargs.get("status") == "pending"
|
||||
else []
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="run creation failed"):
|
||||
background_runs.launch_background_run(_request())
|
||||
|
||||
fake_client.runs.cancel_many.assert_called_once_with(
|
||||
thread_id="thread-1",
|
||||
run_ids=["stray-run"],
|
||||
)
|
||||
fake_client.threads.delete.assert_called_once_with("thread-1")
|
||||
call_names = [name for name, _args, _kwargs in fake_client.mock_calls]
|
||||
assert call_names.index("runs.cancel_many") < call_names.index("threads.delete")
|
||||
|
||||
|
||||
async def test_async_status_watcher_preserves_run_url():
|
||||
finished: list[background_runs.BackgroundRun] = []
|
||||
|
||||
|
||||
@@ -1087,3 +1087,42 @@ async def test_langgraph_server_gateway_resumes_interrupt_with_thread_stream():
|
||||
}
|
||||
]
|
||||
assert events == [{"type": "done", "content": "", "response": ""}]
|
||||
|
||||
|
||||
async def test_langgraph_server_thread_store_cancels_runs_before_delete():
|
||||
events: list[tuple[str, object]] = []
|
||||
|
||||
class _RecordingThreadsClient(FakeLangGraphThreadsClient):
|
||||
async def delete(self, thread_id: str) -> None:
|
||||
events.append(("delete", thread_id))
|
||||
await super().delete(thread_id)
|
||||
|
||||
threads = _RecordingThreadsClient(threads=[{"thread_id": "abc12345"}])
|
||||
client = FakeLangGraphClient(threads)
|
||||
|
||||
class _FakeRunsClient:
|
||||
async def list(self, thread_id: str, *, limit: int, offset: int, status: str):
|
||||
if status == "pending":
|
||||
return [{"run_id": "run-pending", "status": "pending"}]
|
||||
return []
|
||||
|
||||
async def cancel_many(self, *, thread_id: str, run_ids):
|
||||
events.append(("cancel_many", list(run_ids)))
|
||||
|
||||
client.runs = _FakeRunsClient()
|
||||
store = LangGraphServerThreadStore(client=client)
|
||||
|
||||
assert await store.delete_thread("abc12345") is True
|
||||
assert events == [
|
||||
("cancel_many", ["run-pending"]),
|
||||
("delete", "abc12345"),
|
||||
]
|
||||
assert threads.deleted == ["abc12345"]
|
||||
|
||||
|
||||
async def test_langgraph_server_thread_store_delete_survives_missing_runs_client():
|
||||
threads = FakeLangGraphThreadsClient(threads=[{"thread_id": "abc12345"}])
|
||||
store = LangGraphServerThreadStore(client=FakeLangGraphClient(threads))
|
||||
|
||||
assert await store.delete_thread("abc12345") is True
|
||||
assert threads.deleted == ["abc12345"]
|
||||
|
||||
@@ -2865,5 +2865,177 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.IsolatedAsyncioTestCase):
|
||||
assert restore_called, "_restore_webui_threads_to_global_store must be called"
|
||||
|
||||
|
||||
class TestOrphanedRunSweep:
|
||||
"""Startup sweep for runs whose thread no longer exists (issue #358)."""
|
||||
|
||||
def test_removes_runs_whose_thread_is_missing(self):
|
||||
from EvoScientist.sessions import _sweep_orphaned_global_store_entries
|
||||
|
||||
alive = uuid.uuid4()
|
||||
store = {
|
||||
"threads": [{"thread_id": alive}],
|
||||
"runs": [
|
||||
{"run_id": "keep", "thread_id": alive, "status": "pending"},
|
||||
{
|
||||
"run_id": "zombie-pending",
|
||||
"thread_id": uuid.uuid4(),
|
||||
"status": "pending",
|
||||
},
|
||||
{
|
||||
"run_id": "zombie-error",
|
||||
"thread_id": uuid.uuid4(),
|
||||
"status": "error",
|
||||
},
|
||||
],
|
||||
"crons": [{"cron_id": "c1"}],
|
||||
}
|
||||
|
||||
removed = _sweep_orphaned_global_store_entries(store)
|
||||
|
||||
assert removed == 2
|
||||
assert [r["run_id"] for r in store["runs"]] == ["keep"]
|
||||
assert store["crons"] == [{"cron_id": "c1"}]
|
||||
|
||||
def test_matches_uuid_and_str_thread_ids(self):
|
||||
from EvoScientist.sessions import _sweep_orphaned_global_store_entries
|
||||
|
||||
alive = uuid.uuid4()
|
||||
store = {
|
||||
"threads": [{"thread_id": str(alive)}],
|
||||
"runs": [{"run_id": "keep", "thread_id": alive, "status": "pending"}],
|
||||
}
|
||||
|
||||
assert _sweep_orphaned_global_store_entries(store) == 0
|
||||
assert [r["run_id"] for r in store["runs"]] == ["keep"]
|
||||
|
||||
def test_empty_store_is_noop(self):
|
||||
from EvoScientist.sessions import _sweep_orphaned_global_store_entries
|
||||
|
||||
assert _sweep_orphaned_global_store_entries({}) == 0
|
||||
|
||||
def test_mutates_runs_list_in_place(self):
|
||||
from EvoScientist.sessions import _sweep_orphaned_global_store_entries
|
||||
|
||||
runs = [{"run_id": "zombie", "thread_id": uuid.uuid4(), "status": "pending"}]
|
||||
store = {"threads": [], "runs": runs}
|
||||
|
||||
_sweep_orphaned_global_store_entries(store)
|
||||
|
||||
assert store["runs"] is runs
|
||||
assert runs == []
|
||||
|
||||
def test_removes_thread_bound_crons_whose_thread_is_missing(self):
|
||||
from EvoScientist.sessions import _sweep_orphaned_global_store_entries
|
||||
|
||||
alive = uuid.uuid4()
|
||||
store = {
|
||||
"threads": [{"thread_id": alive}],
|
||||
"runs": [],
|
||||
"crons": [
|
||||
{"cron_id": "keep-stateless", "thread_id": None},
|
||||
{"cron_id": "keep-bound", "thread_id": alive},
|
||||
{"cron_id": "zombie-bound", "thread_id": uuid.uuid4()},
|
||||
],
|
||||
}
|
||||
|
||||
removed = _sweep_orphaned_global_store_entries(store)
|
||||
|
||||
assert removed == 1
|
||||
assert [c["cron_id"] for c in store["crons"]] == [
|
||||
"keep-stateless",
|
||||
"keep-bound",
|
||||
]
|
||||
|
||||
def test_stateless_crons_survive_empty_thread_registry(self):
|
||||
from EvoScientist.sessions import _sweep_orphaned_global_store_entries
|
||||
|
||||
store = {
|
||||
"threads": [],
|
||||
"runs": [],
|
||||
"crons": [{"cron_id": "keep-stateless", "thread_id": None}],
|
||||
}
|
||||
|
||||
assert _sweep_orphaned_global_store_entries(store) == 0
|
||||
assert [c["cron_id"] for c in store["crons"]] == ["keep-stateless"]
|
||||
|
||||
async def test_create_checkpointer_calls_sweep(self):
|
||||
"""create_checkpointer_for_langgraph_api runs the orphan sweep."""
|
||||
from unittest.mock import patch
|
||||
|
||||
from EvoScientist.sessions import create_checkpointer_for_langgraph_api
|
||||
|
||||
sweep_called = []
|
||||
|
||||
async def fake_restore():
|
||||
return True
|
||||
|
||||
async def fake_sweep():
|
||||
sweep_called.append(True)
|
||||
|
||||
with tempfile.TemporaryDirectory() as td:
|
||||
db = os.path.join(td, "sessions.db")
|
||||
with (
|
||||
patch(
|
||||
"EvoScientist.sessions.get_db_path",
|
||||
return_value=_mock_path(db),
|
||||
),
|
||||
patch(
|
||||
"EvoScientist.sessions._restore_webui_threads_to_global_store",
|
||||
side_effect=fake_restore,
|
||||
),
|
||||
patch(
|
||||
"EvoScientist.sessions._sweep_orphaned_runs_in_global_store",
|
||||
side_effect=fake_sweep,
|
||||
),
|
||||
):
|
||||
|
||||
async def _run_inner():
|
||||
async with create_checkpointer_for_langgraph_api():
|
||||
pass
|
||||
|
||||
await _run_inner()
|
||||
|
||||
assert sweep_called, "_sweep_orphaned_runs_in_global_store must be called"
|
||||
|
||||
async def test_sweep_skipped_when_restore_fails(self):
|
||||
"""A failed thread restore must not be followed by a destructive sweep."""
|
||||
from unittest.mock import patch
|
||||
|
||||
from EvoScientist.sessions import create_checkpointer_for_langgraph_api
|
||||
|
||||
sweep_called = []
|
||||
|
||||
async def fake_restore():
|
||||
return False
|
||||
|
||||
async def fake_sweep(): # pragma: no cover - must not run
|
||||
sweep_called.append(True)
|
||||
|
||||
with tempfile.TemporaryDirectory() as td:
|
||||
db = os.path.join(td, "sessions.db")
|
||||
with (
|
||||
patch(
|
||||
"EvoScientist.sessions.get_db_path",
|
||||
return_value=_mock_path(db),
|
||||
),
|
||||
patch(
|
||||
"EvoScientist.sessions._restore_webui_threads_to_global_store",
|
||||
side_effect=fake_restore,
|
||||
),
|
||||
patch(
|
||||
"EvoScientist.sessions._sweep_orphaned_runs_in_global_store",
|
||||
side_effect=fake_sweep,
|
||||
),
|
||||
):
|
||||
|
||||
async def _run_inner():
|
||||
async with create_checkpointer_for_langgraph_api():
|
||||
pass
|
||||
|
||||
await _run_inner()
|
||||
|
||||
assert sweep_called == []
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user