fix: cancel abandoned generation tasks

- Propagate stream closure and stop-sequence termination into scheduler cancellation
- Defer active KV release to the scheduler owner and close metrics safely
- Expose lifecycle counters and cover waiting, active, and allocation-race cleanup
This commit is contained in:
0z5a
2026-09-02 12:55:05 +08:00
parent 90de5bc1bd
commit 9c3ef0c2a1
10 changed files with 410 additions and 130 deletions
+47
View File
@@ -157,6 +157,28 @@ def test_engine_generate_streaming_yields_tokens():
assert tokens == ["t1", "t2"]
def test_engine_stream_close_cancels_unfinished_task():
mock_model, mock_tokenizer = _make_engine_mocks(decode="tok")
callbacks_saved = []
def capture_cb(prompt, **kwargs):
callbacks_saved.append(kwargs["stream_callback"])
return "task-1"
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
instance = MockSched.return_value
instance.add_task.side_effect = capture_cb
engine = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=1)
stream = engine.generate("hello", stream=True)
callbacks_saved[0]("t1")
assert next(stream) == "t1"
stream.close()
instance.cancel_task.assert_called_once_with("task-1")
def test_engine_generate_async_yields_tokens_until_stop():
mock_model, mock_tokenizer = _make_engine_mocks(decode="tok")
callbacks_saved = []
@@ -186,6 +208,31 @@ def test_engine_generate_async_yields_tokens_until_stop():
assert asyncio.run(collect()) == ["t1", "t2"]
def test_engine_async_close_cancels_unfinished_task():
mock_model, mock_tokenizer = _make_engine_mocks(decode="tok")
callbacks_saved = []
def capture_cb(prompt, **kwargs):
callbacks_saved.append(kwargs["stream_callback"])
return "task-1"
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
instance = MockSched.return_value
instance.add_task.side_effect = capture_cb
engine = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=1)
stream = engine.generate_async("hello")
callbacks_saved[0]("t1")
async def consume_then_close():
assert await anext(stream) == "t1"
await stream.aclose()
asyncio.run(consume_then_close())
instance.cancel_task.assert_called_once_with("task-1")
def test_engine_generate_non_streaming_batch():
mock_model, mock_tokenizer = _make_engine_mocks(decode="r")