* 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:
@@ -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