From 4a0a6944a2d2fc648146f43ac1694e7a4ab99a60 Mon Sep 17 00:00:00 2001 From: kevin9327 <5299031+kevin9327@users.noreply.github.com> Date: Sun, 13 Sep 2026 06:55:35 +0900 Subject: [PATCH] fix(generate): don't kill the model worker when cancelling a job that already ended cancel_job and cancel_run kill the active generator's subprocess so inference stops at once, but they did so whatever the job's status. For a job that had already finished, failed or been cancelled, that subprocess belongs to whatever is generating now, or holds the warm model: cancelling a finished job (the CLI's `legacy cancel` / `workflow-run cancel`, or pressing Cancel as a generation completes) failed the other generation with "Subprocess died during generation" or forced the next one to reload the model from scratch. Only kill the subprocess when the job being cancelled was still pending or running; cancelling an ended job keeps returning {"cancelled": true}. Co-Authored-By: Claude Opus 5 --- api/routers/generation.py | 4 + api/routers/workflow_runs.py | 4 + api/tests/test_workflow_runs_lifecycle.py | 95 +++++++++++++++++++++++ 3 files changed, 103 insertions(+) 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()