diff --git a/api/routers/generation.py b/api/routers/generation.py index 8481deb4..d6b9d702 100644 --- a/api/routers/generation.py +++ b/api/routers/generation.py @@ -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. diff --git a/api/routers/workflow_runs.py b/api/routers/workflow_runs.py index 9b78a205..b7d0b262 100644 --- a/api/routers/workflow_runs.py +++ b/api/routers/workflow_runs.py @@ -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) diff --git a/api/tests/test_workflow_runs_lifecycle.py b/api/tests/test_workflow_runs_lifecycle.py index f8fe7371..8f294a91 100644 --- a/api/tests/test_workflow_runs_lifecycle.py +++ b/api/tests/test_workflow_runs_lifecycle.py @@ -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()