fix: cascade-cancel runs on thread deletion + startup orphan sweep (#358) (#359)

* 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:
Xi Zhang
2026-07-16 12:46:09 +01:00
committed by GitHub
parent 05dfffbc73
commit 0f709cff8b
7 changed files with 575 additions and 6 deletions
+190
View File
@@ -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] = []
+39
View File
@@ -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"]
+172
View File
@@ -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()