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:
parent
d5e5733f67
commit
da8e7dc568
2 changed files with 50 additions and 0 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
Loading…
Reference in a new issue