Merge pull request #405 from mcdonc/fix/run-batch-hang-on-dead-workers
fix: run_batch hangs forever when all workers die
This commit is contained in:
commit
64f2b7b7d2
2 changed files with 42 additions and 0 deletions
|
|
@ -205,6 +205,15 @@ class IngesterApp:
|
||||||
counts = await self._jobs.counts_by_status()
|
counts = await self._jobs.counts_by_status()
|
||||||
if not counts.get("queued") and not counts.get("claimed"):
|
if not counts.get("queued") and not counts.get("claimed"):
|
||||||
break
|
break
|
||||||
|
if self._pool.live_workers == 0:
|
||||||
|
outstanding = counts.get("queued", 0) + counts.get("claimed", 0)
|
||||||
|
logger.error(
|
||||||
|
"All workers have died with %d outstanding job(s) "
|
||||||
|
"— aborting batch; stranded jobs will be reaped "
|
||||||
|
"on next start",
|
||||||
|
outstanding,
|
||||||
|
)
|
||||||
|
break
|
||||||
await asyncio.sleep(0.1)
|
await asyncio.sleep(0.1)
|
||||||
completed = await self._jobs.counts_by_status_since(started_at)
|
completed = await self._jobs.counts_by_status_since(started_at)
|
||||||
return BatchReport(
|
return BatchReport(
|
||||||
|
|
|
||||||
|
|
@ -235,6 +235,39 @@ async def test_run_batch_empty_source_returns_immediately(tmp_path, use_client):
|
||||||
client.create_document_from_source.assert_not_awaited()
|
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, caplog
|
||||||
|
):
|
||||||
|
"""If all workers crash with outstanding jobs, run_batch should break
|
||||||
|
out of the drain loop instead of spinning forever."""
|
||||||
|
(tmp_path / "a.md").write_text("hello")
|
||||||
|
|
||||||
|
client = _mock_client()
|
||||||
|
use_client(client)
|
||||||
|
config = _config(tmp_path)
|
||||||
|
|
||||||
|
config.ingester.workers.worker_count = 1
|
||||||
|
|
||||||
|
# Patch _process to raise an unhandled exception, simulating a hard
|
||||||
|
# worker crash. _process only catches CancelledError, PermanentError,
|
||||||
|
# and TransientError — anything else propagates and kills the task.
|
||||||
|
monkeypatch.setattr(
|
||||||
|
WorkerPool,
|
||||||
|
"_process",
|
||||||
|
AsyncMock(side_effect=Exception("worker crash")),
|
||||||
|
)
|
||||||
|
|
||||||
|
with caplog.at_level("ERROR", logger="haiku.rag.ingester.app"):
|
||||||
|
report = await asyncio.wait_for(
|
||||||
|
IngesterApp(config=config, db_path=tmp_path / "db.lancedb").run_batch(),
|
||||||
|
timeout=10.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert report is not None
|
||||||
|
assert "All workers have died" in caplog.text
|
||||||
|
|
||||||
|
|
||||||
async def _wait_until(predicate, *, timeout: float = 5.0):
|
async def _wait_until(predicate, *, timeout: float = 5.0):
|
||||||
deadline = asyncio.get_running_loop().time() + timeout
|
deadline = asyncio.get_running_loop().time() + timeout
|
||||||
while not predicate():
|
while not predicate():
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue