Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions api/routers/generation.py
Original file line number Diff line number Diff line change
Expand Up @@ -156,6 +156,10 @@ async def cancel_job(job_id: str):
if job.status in ("pending", "running"):
job.status = "cancelled"
_completed_at[job_id] = time.monotonic()
else:
# The job already ended, so the active subprocess isn't running it: it
# holds the warm model or another job's generation. Leave it alone.
return {"cancelled": True}
# Kill the active generator subprocess immediately so inference stops now.
# _run_generation will catch the resulting exception, see job_id in _cancelled,
# and return cleanly without setting an error status.
Expand Down
4 changes: 4 additions & 0 deletions api/routers/workflow_runs.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,6 +121,10 @@ async def cancel_run(run_id: str):
if job.status in ("pending", "running"):
job.status = "cancelled"
_completed_at[run_id] = time.monotonic()
else:
# The run already ended, so the active subprocess isn't running it: it
# holds the warm model or another job's generation. Leave it alone.
return {"cancelled": True}

try:
gen = generator_registry._generators.get(generator_registry._active_id)
Expand Down
95 changes: 95 additions & 0 deletions api/tests/test_workflow_runs_lifecycle.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,5 +92,100 @@ def test_cancel_run_records_completion_so_it_can_be_purged(self) -> None:
self.assertIn(run_id, generation._completed_at)


class _FakeProc:
"""Stands in for a loaded extension worker's subprocess."""

def __init__(self) -> None:
self.killed = False

def poll(self):
return -9 if self.killed else None

def kill(self) -> None:
self.killed = True


class _FakeWorker:
def __init__(self) -> None:
self._proc = _FakeProc()
self._loaded = True


class _RegistryWithLoadedWorker:
def __init__(self) -> None:
self.worker = _FakeWorker()
self._generators = {"ext/model": self.worker}
self._active_id = "ext/model"


class CancelEndedJobTests(unittest.TestCase):
"""Cancelling kills the active generator's subprocess so inference stops at
once. That is only right while the job being cancelled is still pending or
running: once it has ended, that subprocess belongs to whatever is generating
now (or holds the warm model), and killing it fails that other generation or
forces the next one to reload the model from scratch."""

def setUp(self) -> None:
self._prev_generation_registry = generation.generator_registry
self._prev_runs_registry = workflow_runs.generator_registry
self.registry = _RegistryWithLoadedWorker()
generation.generator_registry = self.registry
workflow_runs.generator_registry = self.registry
_clear_job_stores()

def tearDown(self) -> None:
generation.generator_registry = self._prev_generation_registry
workflow_runs.generator_registry = self._prev_runs_registry
_clear_job_stores()

def _file_job(self, job_id: str, status: str) -> None:
generation._jobs[job_id] = JobStatus(job_id=job_id, status=status, progress=50)
generation._cancel_events[job_id] = threading.Event()

def test_cancelling_a_finished_job_leaves_the_active_worker_alone(self) -> None:
self._file_job("finished-job", "done")
proc = self.registry.worker._proc

asyncio.run(generation.cancel_job("finished-job"))

self.assertEqual(generation._jobs["finished-job"].status, "done")
self.assertFalse(proc.killed)
self.assertIs(self.registry.worker._proc, proc)
self.assertTrue(self.registry.worker._loaded)

def test_cancelling_a_failed_run_leaves_the_active_worker_alone(self) -> None:
self._file_job("failed-run", "error")
proc = self.registry.worker._proc

asyncio.run(workflow_runs.cancel_run("failed-run"))

self.assertEqual(generation._jobs["failed-run"].status, "error")
self.assertFalse(proc.killed)
self.assertIs(self.registry.worker._proc, proc)
self.assertTrue(self.registry.worker._loaded)

def test_cancelling_a_running_job_still_stops_the_worker(self) -> None:
self._file_job("running-job", "running")
proc = self.registry.worker._proc

asyncio.run(generation.cancel_job("running-job"))

self.assertEqual(generation._jobs["running-job"].status, "cancelled")
self.assertTrue(proc.killed)
self.assertIsNone(self.registry.worker._proc)
self.assertFalse(self.registry.worker._loaded)

def test_cancelling_a_running_run_still_stops_the_worker(self) -> None:
self._file_job("running-run", "running")
proc = self.registry.worker._proc

asyncio.run(workflow_runs.cancel_run("running-run"))

self.assertEqual(generation._jobs["running-run"].status, "cancelled")
self.assertTrue(proc.killed)
self.assertIsNone(self.registry.worker._proc)
self.assertFalse(self.registry.worker._loaded)


if __name__ == "__main__":
unittest.main()