Fix run_batch hanging forever when all workers die

The drain loop in run_batch() polls counts_by_status() waiting for
queued and claimed counts to reach zero. If all worker tasks crash
(unhandled exception, OOM), claimed jobs stay claimed forever and
the loop never exits — the CLI command hangs.

Check live_workers during the drain loop. If claimed jobs exist but
no workers are alive to process them, log an error and break out.
The stranded jobs will be reaped on the next start.
This commit is contained in:
Chris McDonough 2026-06-01 07:52:17 -04:00
parent d5e5733f67
commit da8e7dc568
2 changed files with 50 additions and 0 deletions

View file

@ -204,6 +204,14 @@ class IngesterApp:
counts = await self._jobs.counts_by_status()
if not counts.get("queued") and not counts.get("claimed"):
break
if counts.get("claimed") and self._pool.live_workers == 0:
logger.error(
"All workers have died with %d claimed job(s) — "
"aborting batch; stranded jobs will be reaped on "
"next start",
counts["claimed"],
)
break
await asyncio.sleep(0.1)
completed = await self._jobs.counts_by_status_since(started_at)
return BatchReport(

View file

@ -234,6 +234,48 @@ async def test_run_batch_empty_source_returns_immediately(tmp_path, use_client):
client.create_document_from_source.assert_not_awaited()
@pytest.mark.asyncio
async def test_run_batch_aborts_when_all_workers_die(tmp_path, use_client, monkeypatch):
"""If all workers crash with claimed jobs still outstanding, run_batch
should break out of the drain loop instead of spinning forever."""
(tmp_path / "a.md").write_text("hello")
client = _mock_client()
# Block the worker forever so the job stays claimed until we kill it.
stall = asyncio.Event()
client.create_document_from_source.side_effect = lambda *a, **k: stall.wait()
use_client(client)
config = _config(tmp_path)
app = IngesterApp(config=config, db_path=tmp_path / "db.lancedb")
async def _kill_workers_after_claim():
"""Wait until at least one job is claimed, then kill all workers."""
pool = app._pool
assert pool is not None
for _ in range(200):
counts = await app._jobs.counts_by_status()
if counts.get("claimed"):
break
await asyncio.sleep(0.05)
for task in pool._workers:
task.cancel()
await asyncio.gather(*pool._workers, return_exceptions=True)
# Run the killer concurrently with run_batch.
batch_task = asyncio.create_task(app.run_batch())
# Give run_batch a moment to start, then schedule the killer.
await asyncio.sleep(0.1)
killer_task = asyncio.create_task(_kill_workers_after_claim())
report = await asyncio.wait_for(batch_task, timeout=10.0)
# Killer may still be running against a closed DB — suppress errors.
killer_task.cancel()
await asyncio.gather(killer_task, return_exceptions=True)
# The batch should have exited without hanging.
assert report is not None
async def _wait_until(predicate, *, timeout: float = 5.0):
deadline = asyncio.get_running_loop().time() + timeout
while not predicate():