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")
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
"""Tests for scheduler concurrency."""
|
||||
|
||||
import threading
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
@@ -259,6 +260,104 @@ def _make_real_scheduler(device):
|
||||
return scheduler, tokenizer, model
|
||||
|
||||
|
||||
def test_cancel_waiting_task_storm_returns_to_baseline(device):
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
task_ids = [
|
||||
scheduler.add_task(f"waiting-{index}", max_tokens=32) for index in range(32)
|
||||
]
|
||||
|
||||
assert all(scheduler.cancel_task(task_id) for task_id in task_ids)
|
||||
stats = scheduler.get_stats()
|
||||
assert stats["active_tasks"] == 0
|
||||
assert stats["waiting_tasks"] == 0
|
||||
assert stats["in_flight_tasks"] == 0
|
||||
assert stats["kv_cache_tasks"] == 0
|
||||
assert stats["cancelled_total"] == len(task_ids)
|
||||
finally:
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_cancel_active_task_releases_metrics_and_kv(device):
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
task_id = scheduler.add_task("active", max_tokens=32)
|
||||
task = scheduler._task_mgr.pull_candidates(1)[0]
|
||||
assert scheduler._task_cache.task_alloc(task.task_id, task.prompt_ids)
|
||||
assert scheduler._task_mgr.activate(task)
|
||||
|
||||
before = scheduler.get_stats()
|
||||
assert before["active_tasks"] == 1
|
||||
assert before["in_flight_tasks"] == 1
|
||||
assert before["kv_cache_tasks"] == 1
|
||||
|
||||
assert scheduler.cancel_task(task_id)
|
||||
scheduler.start()
|
||||
deadline = time.monotonic() + 5
|
||||
while time.monotonic() < deadline:
|
||||
after = scheduler.get_stats()
|
||||
if (
|
||||
after["active_tasks"] == 0
|
||||
and after["in_flight_tasks"] == 0
|
||||
and after["kv_cache_tasks"] == 0
|
||||
):
|
||||
break
|
||||
time.sleep(0.01)
|
||||
|
||||
assert after["active_tasks"] == 0
|
||||
assert after["waiting_tasks"] == 0
|
||||
assert after["in_flight_tasks"] == 0
|
||||
assert after["kv_cache_tasks"] == 0
|
||||
assert after["cancelled_total"] == 1
|
||||
finally:
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_cancel_during_kv_allocation_releases_metrics_and_kv(device):
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
allocation_started = threading.Event()
|
||||
continue_allocation = threading.Event()
|
||||
original_alloc = scheduler._task_cache.task_alloc
|
||||
|
||||
def blocking_alloc(*args, **kwargs):
|
||||
allocation_started.set()
|
||||
assert continue_allocation.wait(timeout=5)
|
||||
return original_alloc(*args, **kwargs)
|
||||
|
||||
try:
|
||||
with patch.object(
|
||||
scheduler._task_cache,
|
||||
"task_alloc",
|
||||
side_effect=blocking_alloc,
|
||||
):
|
||||
scheduler.start()
|
||||
task_id = scheduler.add_task("allocation-race", max_tokens=32)
|
||||
assert allocation_started.wait(timeout=5)
|
||||
assert scheduler.cancel_task(task_id)
|
||||
continue_allocation.set()
|
||||
|
||||
deadline = time.monotonic() + 5
|
||||
while time.monotonic() < deadline:
|
||||
stats = scheduler.get_stats()
|
||||
if (
|
||||
stats["active_tasks"] == 0
|
||||
and stats["waiting_tasks"] == 0
|
||||
and stats["in_flight_tasks"] == 0
|
||||
and stats["kv_cache_tasks"] == 0
|
||||
):
|
||||
break
|
||||
time.sleep(0.01)
|
||||
|
||||
assert stats["active_tasks"] == 0
|
||||
assert stats["waiting_tasks"] == 0
|
||||
assert stats["in_flight_tasks"] == 0
|
||||
assert stats["kv_cache_tasks"] == 0
|
||||
assert stats["cancelled_total"] == 1
|
||||
finally:
|
||||
continue_allocation.set()
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_returns_token_sequences(device):
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
|
||||
@@ -166,11 +166,15 @@ def test_messages_with_system(client, loaded_model):
|
||||
|
||||
def test_chat_completions_stop_sequence(client, loaded_model):
|
||||
"""POST /v1/chat/completions with stop parameter truncates at stop sequence."""
|
||||
closed = []
|
||||
|
||||
async def async_gen():
|
||||
yield "Hello"
|
||||
yield "X"
|
||||
yield "world"
|
||||
try:
|
||||
yield "Hello"
|
||||
yield "X"
|
||||
yield "world"
|
||||
finally:
|
||||
closed.append(True)
|
||||
|
||||
get_app().state.engine = loaded_model
|
||||
loaded_model.generate_async.return_value = async_gen()
|
||||
@@ -188,15 +192,20 @@ def test_chat_completions_stop_sequence(client, loaded_model):
|
||||
content = data["choices"][0]["message"]["content"]
|
||||
assert "X" in content
|
||||
assert "world" not in content
|
||||
assert closed == [True]
|
||||
|
||||
|
||||
def test_chat_completions_stop_sequence_stream(client, loaded_model):
|
||||
"""POST /v1/chat/completions with stop parameter truncates SSE stream."""
|
||||
closed = []
|
||||
|
||||
async def async_gen():
|
||||
yield "Hello"
|
||||
yield "X"
|
||||
yield "world"
|
||||
try:
|
||||
yield "Hello"
|
||||
yield "X"
|
||||
yield "world"
|
||||
finally:
|
||||
closed.append(True)
|
||||
|
||||
get_app().state.engine = loaded_model
|
||||
loaded_model.generate_async.return_value = async_gen()
|
||||
@@ -217,6 +226,7 @@ def test_chat_completions_stop_sequence_stream(client, loaded_model):
|
||||
assert any(
|
||||
"finish_reason" in line for line in content.split("\n") if "stop" in line
|
||||
)
|
||||
assert closed == [True]
|
||||
|
||||
|
||||
def test_chat_completions_real_engine(tmp_path, client):
|
||||
|
||||
@@ -73,15 +73,21 @@ def test_task_manager_remove_task():
|
||||
assert len(tm.waiting_queue) == 0
|
||||
|
||||
|
||||
def test_task_manager_remove_active_task():
|
||||
def test_task_manager_cancel_active_task_defers_removal():
|
||||
tm = TaskManager(tokenizer=_make_mock_tokenizer())
|
||||
tid = tm.add_task("test")
|
||||
tasks = tm.pull_candidates(1)
|
||||
tm.activate(tasks[0])
|
||||
assert len(tm.active_tasks) == 1
|
||||
removed = tm.remove_task(tid)
|
||||
assert len(removed) == 1
|
||||
immediate, cancelled = tm.cancel_task(tid)
|
||||
assert cancelled is True
|
||||
assert immediate == []
|
||||
assert tm.active_tasks[0].status == TaskStatus.ABORTED
|
||||
|
||||
removed = tm.remove_finished_tasks([])
|
||||
assert removed == tasks
|
||||
assert len(tm.active_tasks) == 0
|
||||
assert tm.get_stats()["cancelled_total"] == 1
|
||||
|
||||
|
||||
def test_task_manager_pull_candidates_fifo():
|
||||
|
||||
Reference in New Issue
Block a user