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:
@@ -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")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user