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
+5
View File
@@ -322,6 +322,11 @@ class TaskCacheManager:
state.length = pos + 1 state.length = pos + 1
return True return True
@property
def task_count(self) -> int:
"""Number of tasks currently holding KV request state."""
return len(self._states)
def task_cached(self, task_id: str) -> int: def task_cached(self, task_id: str) -> int:
state = self._states.get(task_id) state = self._states.get(task_id)
return state.cached if state is not None else 0 return state.cached if state is not None else 0
+43 -32
View File
@@ -43,8 +43,7 @@ class GenerateResult:
with self._cond: with self._cond:
out = self.tokens.copy() out = self.tokens.copy()
self.tokens.clear() self.tokens.clear()
if not out: self._event.clear()
self._event.clear()
return out return out
def wait(self, timeout: Optional[float] = None) -> bool: def wait(self, timeout: Optional[float] = None) -> bool:
@@ -141,25 +140,34 @@ class InferenceEngine:
frequency_penalty: float = 0.0, frequency_penalty: float = 0.0,
rep_window: int = 64, rep_window: int = 64,
) -> AsyncGenerator[str, None]: ) -> AsyncGenerator[str, None]:
sync_gen = self._generate( request_backend = get_backend(use_default=False)
[prompt], result = GenerateResult()
False, task_id = self.scheduler.add_task(
True, prompt=prompt,
max_tokens, max_tokens=max_tokens,
temperature, temperature=temperature,
top_p, top_p=top_p,
top_k, top_k=top_k,
frequency_penalty, frequency_penalty=frequency_penalty,
rep_window, rep_window=rep_window,
backend=request_backend,
stream_callback=result.append,
) )
async def _agen(): async def _agen():
loop = asyncio.get_event_loop() finished = False
while True: try:
token = await loop.run_in_executor(None, next, sync_gen, None) while not finished:
if token is None: for _idx, token in result.pop_all():
break if token is STOP:
yield token finished = True
break
yield token
if not finished:
await asyncio.to_thread(result.wait, 0.05)
finally:
if not finished:
self.scheduler.cancel_task(task_id)
return _agen() return _agen()
@@ -198,10 +206,8 @@ class InferenceEngine:
result.wait_completion() result.wait_completion()
except TimeoutError: except TimeoutError:
for tid in task_ids: for tid in task_ids:
self.scheduler.remove_task(tid) self.scheduler.cancel_task(tid)
raise raise
for tid in task_ids:
self.scheduler.remove_task(tid)
res = result.get_results() res = result.get_results()
return res if is_batch else res[0] return res if is_batch else res[0]
@@ -210,17 +216,22 @@ class InferenceEngine:
def gen(): def gen():
nonlocal remaining nonlocal remaining
while remaining > 0: try:
items = result.pop_all() while remaining > 0:
for idx, token in items: items = result.pop_all()
if token is STOP: for idx, token in items:
if not finished[idx]: if token is STOP:
finished[idx] = True if not finished[idx]:
remaining -= 1 finished[idx] = True
else: remaining -= 1
yield (idx, token) if is_batch else token else:
if remaining > 0: yield (idx, token) if is_batch else token
result.wait(timeout=0.05) if remaining > 0:
result.wait(timeout=0.05)
finally:
for idx, task_id in enumerate(task_ids):
if not finished[idx]:
self.scheduler.cancel_task(task_id)
return gen() return gen()
+41 -33
View File
@@ -1,5 +1,6 @@
"""Unified per-task perf/stats: timing records, context-manager scopes, aggregate reporting.""" """Unified per-task perf/stats: timing records, context-manager scopes, aggregate reporting."""
import threading
import time import time
from collections import deque from collections import deque
from contextlib import contextmanager from contextlib import contextmanager
@@ -125,6 +126,7 @@ class MetricsCollector:
def __init__(self, max_recent: int = 128): def __init__(self, max_recent: int = 128):
self._timings: Dict[str, TaskTiming] = {} self._timings: Dict[str, TaskTiming] = {}
self._completed: Deque[TaskTiming] = deque(maxlen=max_recent) self._completed: Deque[TaskTiming] = deque(maxlen=max_recent)
self._lock = threading.Lock()
self._ttft_ms_sum = 0.0 self._ttft_ms_sum = 0.0
self._ttft_ms_count = 0 self._ttft_ms_count = 0
@@ -135,18 +137,22 @@ class MetricsCollector:
def register(self, task_id: str): def register(self, task_id: str):
"""Create a timing record for a newly-created task.""" """Create a timing record for a newly-created task."""
self._timings[task_id] = TaskTiming(task_id=task_id, arrival_time=time.time()) with self._lock:
self._timings[task_id] = TaskTiming(
task_id=task_id, arrival_time=time.time()
)
def mark_finished(self, task_id: str, input_tokens: int, output_tokens: int): def mark_finished(self, task_id: str, input_tokens: int, output_tokens: int):
"""Close timing for a finished/aborted task and move it to completed.""" """Close timing for a finished/aborted task and move it to completed."""
timing = self._timings.pop(task_id, None) with self._lock:
if timing is None: timing = self._timings.pop(task_id, None)
return if timing is None:
timing.finish_time = time.time() return
timing.input_tokens = input_tokens timing.finish_time = time.time()
timing.output_tokens = output_tokens timing.input_tokens = input_tokens
self._completed.append(timing) timing.output_tokens = output_tokens
self._accumulate(timing) self._completed.append(timing)
self._accumulate(timing)
# timing scopes # timing scopes
@@ -158,34 +164,36 @@ class MetricsCollector:
yield yield
toc = time.time() toc = time.time()
dt = toc - tic dt = toc - tic
for tid in task_ids: with self._lock:
t = self._timings.get(tid) for tid in task_ids:
if t is None: t = self._timings.get(tid)
continue if t is None:
if phase == "prefill": continue
t.prefill_start_time = tic if phase == "prefill":
t.first_token_time = toc t.prefill_start_time = tic
elif phase == "decode": t.first_token_time = toc
t._decode_steps += 1 elif phase == "decode":
t._decode_total_s += dt t._decode_steps += 1
t._decode_total_s += dt
# aggregate stats # aggregate stats
def get_stats(self) -> Dict[str, Any]: def get_stats(self) -> Dict[str, Any]:
stats: Dict[str, Any] = {} with self._lock:
if self._ttft_ms_count > 0: stats: Dict[str, Any] = {"in_flight_tasks": len(self._timings)}
stats["avg_ttft_ms"] = round(self._ttft_ms_sum / self._ttft_ms_count, 2) if self._ttft_ms_count > 0:
if self._decode_tps_count > 0: stats["avg_ttft_ms"] = round(self._ttft_ms_sum / self._ttft_ms_count, 2)
stats["avg_decode_tps"] = round( if self._decode_tps_count > 0:
self._decode_tps_sum / self._decode_tps_count, 2 stats["avg_decode_tps"] = round(
) self._decode_tps_sum / self._decode_tps_count, 2
if self._e2e_ms_count > 0: )
stats["avg_e2e_latency_ms"] = round( if self._e2e_ms_count > 0:
self._e2e_ms_sum / self._e2e_ms_count, 2 stats["avg_e2e_latency_ms"] = round(
) self._e2e_ms_sum / self._e2e_ms_count, 2
if self._completed: )
stats["recent_tasks"] = [t.to_dict() for t in self._completed] if self._completed:
return stats stats["recent_tasks"] = [t.to_dict() for t in self._completed]
return stats
# internal # internal
+28 -22
View File
@@ -146,25 +146,28 @@ class ProtocolHandler:
yielded = "" yielded = ""
matched = None matched = None
token_ids: List[int] = [] token_ids: List[int] = []
async for token in agen: try:
body += token async for token in agen:
body += token
new_ids = self.engine.tokenizer.encode(token) new_ids = self.engine.tokenizer.encode(token)
token_ids.extend(new_ids) token_ids.extend(new_ids)
matched = checker.check(body) matched = checker.check(body)
if matched: if matched:
break break
ctx.completion_tokens += 1 ctx.completion_tokens += 1
for event in self.builder.format_chunk( for event in self.builder.format_chunk(
token, token,
body=body, body=body,
current_token_ids=token_ids, current_token_ids=token_ids,
delta_token_ids=new_ids, delta_token_ids=new_ids,
): ):
yield event yield event
yielded += token yielded += token
finally:
await agen.aclose()
stop = StopInfo(matched=matched, body=body, yielded=yielded) stop = StopInfo(matched=matched, body=body, yielded=yielded)
for event in self.builder.format_stream_end(ctx, stop): for event in self.builder.format_stream_end(ctx, stop):
@@ -184,14 +187,17 @@ class ProtocolHandler:
body = "" body = ""
matched = None matched = None
async for token in agen: try:
body += token async for token in agen:
body += token
matched = checker.check(body) matched = checker.check(body)
if matched: if matched:
break break
ctx.completion_tokens += 1 ctx.completion_tokens += 1
finally:
await agen.aclose()
stop = StopInfo(matched=matched, body=body) stop = StopInfo(matched=matched, body=body)
return self.builder.format_response(ctx, body, stop) return self.builder.format_response(ctx, body, stop)
+41 -8
View File
@@ -107,12 +107,25 @@ class InferenceScheduler:
def add_task(self, prompt: str, **kwargs) -> str: def add_task(self, prompt: str, **kwargs) -> str:
return self._task_mgr.add_task(prompt, **kwargs) return self._task_mgr.add_task(prompt, **kwargs)
def remove_task(self, task_id: str): def cancel_task(self, task_id: str) -> bool:
for task in self._task_mgr.remove_task(task_id): """Cancel a waiting or active task without freeing in-use KV state."""
self._task_cache.task_free(task.task_id) immediate, cancelled = self._task_mgr.cancel_task(task_id)
for task in immediate:
self._metrics.mark_finished(
task.task_id, task.input_tokens, task.output_tokens
)
if cancelled:
self._task_mgr.wake()
return cancelled
def remove_task(self, task_id: str) -> bool:
"""Backward-compatible alias for cancellation."""
return self.cancel_task(task_id)
def get_stats(self) -> Dict[str, Any]: def get_stats(self) -> Dict[str, Any]:
return self._task_mgr.get_stats() stats = self._task_mgr.get_stats()
stats["kv_cache_tasks"] = self._task_cache.task_count
return stats
@property @property
def backend_name(self) -> str: def backend_name(self) -> str:
@@ -254,7 +267,13 @@ class InferenceScheduler:
if self._task_cache.task_alloc( if self._task_cache.task_alloc(
task.task_id, task.prompt_ids task.task_id, task.prompt_ids
): ):
self._task_mgr.activate(task) if not self._task_mgr.activate(task):
self._task_cache.task_free(task.task_id)
self._metrics.mark_finished(
task.task_id,
task.input_tokens,
task.output_tokens,
)
else: else:
failed.append(task) failed.append(task)
if failed: if failed:
@@ -264,7 +283,11 @@ class InferenceScheduler:
self._task_mgr.wait_for_tasks(timeout=1.0) self._task_mgr.wait_for_tasks(timeout=1.0)
continue continue
active = self._task_mgr.get_active_tasks() active = [
task
for task in self._task_mgr.get_active_tasks()
if task.status != TaskStatus.ABORTED
]
decoded, aborted = self._step(active) decoded, aborted = self._step(active)
@@ -272,6 +295,8 @@ class InferenceScheduler:
self._task_mgr.invoke_callback(t.task_id, STOP) self._task_mgr.invoke_callback(t.task_id, STOP)
for t in decoded: for t in decoded:
if t.status == TaskStatus.ABORTED:
continue
new_text = t.decode_new_token(self._task_mgr.tokenizer) new_text = t.decode_new_token(self._task_mgr.tokenizer)
if new_text: if new_text:
self._task_mgr.invoke_callback(t.task_id, new_text) self._task_mgr.invoke_callback(t.task_id, new_text)
@@ -303,13 +328,21 @@ class InferenceScheduler:
def _abort_and_clear(self, free_waiting: bool): def _abort_and_clear(self, free_waiting: bool):
"""Invoke STOP callbacks, release cache slots, and clear task queues.""" """Invoke STOP callbacks, release cache slots, and clear task queues."""
for task in self._task_mgr.get_active_tasks(): active = self._task_mgr.get_active_tasks()
waiting = self._task_mgr.get_waiting_tasks()
for task in active:
self._task_mgr.invoke_callback(task.task_id, STOP) self._task_mgr.invoke_callback(task.task_id, STOP)
self._task_cache.task_free(task.task_id) self._task_cache.task_free(task.task_id)
for task in self._task_mgr.get_waiting_tasks(): self._metrics.mark_finished(
task.task_id, task.input_tokens, task.output_tokens
)
for task in waiting:
self._task_mgr.invoke_callback(task.task_id, STOP) self._task_mgr.invoke_callback(task.task_id, STOP)
if free_waiting: if free_waiting:
self._task_cache.task_free(task.task_id) self._task_cache.task_free(task.task_id)
self._metrics.mark_finished(
task.task_id, task.input_tokens, task.output_tokens
)
self._task_mgr.clear_queues() self._task_mgr.clear_queues()
def run_batch( def run_batch(
+81 -26
View File
@@ -4,7 +4,17 @@ import uuid
from collections import deque from collections import deque
from dataclasses import dataclass from dataclasses import dataclass
from enum import Enum from enum import Enum
from typing import TYPE_CHECKING, Any, Callable, Deque, Dict, List, Literal, Optional from typing import (
TYPE_CHECKING,
Any,
Callable,
Deque,
Dict,
List,
Literal,
Optional,
Tuple,
)
from tokenizers.decoders import DecodeStream from tokenizers.decoders import DecodeStream
@@ -150,12 +160,14 @@ class TaskManager:
self.waiting_queue: Deque[Task] = deque() self.waiting_queue: Deque[Task] = deque()
self.active_tasks: List[Task] = [] self.active_tasks: List[Task] = []
self._callbacks: Dict[str, Callable[[str], None]] = {} self._callbacks: Dict[str, Callable[[str], None]] = {}
self._tasks: Dict[str, Task] = {}
self._task_event = threading.Event() self._task_event = threading.Event()
self._lock = threading.Lock() self._lock = threading.Lock()
self._total_tasks = 0 self._total_tasks = 0
self._total_tokens = 0 self._total_tokens = 0
self._cancelled_total = 0
self._metrics = metrics self._metrics = metrics
@@ -195,6 +207,7 @@ class TaskManager:
with self._lock: with self._lock:
self.waiting_queue.append(task) self.waiting_queue.append(task)
self._tasks[task_id] = task
self._total_tasks += 1 self._total_tasks += 1
if stream_callback: if stream_callback:
self._callbacks[task_id] = stream_callback self._callbacks[task_id] = stream_callback
@@ -205,28 +218,49 @@ class TaskManager:
self._task_event.set() self._task_event.set()
return task_id return task_id
def remove_task(self, task_id: str) -> List[Task]: def cancel_task(self, task_id: str) -> Tuple[List[Task], bool]:
"""Mark a task cancelled and return tasks safe to clean immediately."""
with self._lock: with self._lock:
removed_active = [t for t in self.active_tasks if t.task_id == task_id] task = self._tasks.get(task_id)
self.waiting_queue = deque(
t for t in self.waiting_queue if t.task_id != task_id
)
self.active_tasks = [t for t in self.active_tasks if t.task_id != task_id]
self._callbacks.pop(task_id, None) self._callbacks.pop(task_id, None)
return removed_active if task is None or task.status in (
TaskStatus.FINISHED,
TaskStatus.ABORTED,
):
return [], False
task.status = TaskStatus.ABORTED
self._cancelled_total += 1
if task in self.waiting_queue:
self.waiting_queue = deque(
waiting for waiting in self.waiting_queue if waiting is not task
)
self._tasks.pop(task_id, None)
return [task], True
return [], True
def remove_task(self, task_id: str) -> List[Task]:
"""Backward-compatible alias for cancellation."""
immediate, _ = self.cancel_task(task_id)
return immediate
def invoke_callback(self, task_id: str, token: str): def invoke_callback(self, task_id: str, token: str):
cb = self._callbacks.get(task_id) with self._lock:
cb = self._callbacks.get(task_id)
if cb: if cb:
cb(token) cb(token)
def get_stats(self) -> Dict[str, Any]: def get_stats(self) -> Dict[str, Any]:
stats: Dict[str, Any] = { with self._lock:
"total_tasks": self._total_tasks, waiting = len(self.waiting_queue)
"total_tokens": self._total_tokens, stats: Dict[str, Any] = {
"active_tasks": len(self.active_tasks), "total_tasks": self._total_tasks,
"waiting_queue": len(self.waiting_queue), "total_tokens": self._total_tokens,
} "active_tasks": len(self.active_tasks),
"waiting_tasks": waiting,
"waiting_queue": waiting,
"cancelled_total": self._cancelled_total,
}
if self._metrics is not None: if self._metrics is not None:
stats.update(self._metrics.get_stats()) stats.update(self._metrics.get_stats())
return stats return stats
@@ -242,18 +276,21 @@ class TaskManager:
finished.append(task) finished.append(task)
self._total_tokens += task.output_tokens self._total_tokens += task.output_tokens
if self._metrics is not None:
for task in finished:
self._metrics.mark_finished(
task.task_id, task.input_tokens, task.output_tokens
)
self.active_tasks = [ self.active_tasks = [
t t
for t in self.active_tasks for t in self.active_tasks
if t.status not in (TaskStatus.FINISHED, TaskStatus.ABORTED) if t.status not in (TaskStatus.FINISHED, TaskStatus.ABORTED)
] ]
return finished for task in finished:
self._tasks.pop(task.task_id, None)
self._callbacks.pop(task.task_id, None)
if self._metrics is not None:
for task in finished:
self._metrics.mark_finished(
task.task_id, task.input_tokens, task.output_tokens
)
return finished
def pull_candidates(self, n: int) -> List[Task]: def pull_candidates(self, n: int) -> List[Task]:
to_add: List[Task] = [] to_add: List[Task] = []
@@ -263,18 +300,35 @@ class TaskManager:
to_add.append(self.waiting_queue.popleft()) to_add.append(self.waiting_queue.popleft())
return to_add return to_add
def activate(self, task: Task): def activate(self, task: Task) -> bool:
task.status = TaskStatus.RUNNING
with self._lock: with self._lock:
if task.status == TaskStatus.ABORTED:
self._tasks.pop(task.task_id, None)
self._callbacks.pop(task.task_id, None)
return False
task.status = TaskStatus.RUNNING
self.active_tasks.append(task) self.active_tasks.append(task)
return True
def return_to_waiting(self, tasks: List[Task]): def return_to_waiting(self, tasks: List[Task]):
cancelled = []
with self._lock: with self._lock:
for task in reversed(tasks): for task in reversed(tasks):
self.waiting_queue.appendleft(task) if task.status == TaskStatus.ABORTED:
self._tasks.pop(task.task_id, None)
self._callbacks.pop(task.task_id, None)
cancelled.append(task)
else:
self.waiting_queue.appendleft(task)
if self._metrics is not None:
for task in cancelled:
self._metrics.mark_finished(
task.task_id, task.input_tokens, task.output_tokens
)
def has_work(self) -> bool: def has_work(self) -> bool:
return bool(self.active_tasks or self.waiting_queue) with self._lock:
return bool(self.active_tasks or self.waiting_queue)
def wait_for_tasks(self, timeout: float = 1.0): def wait_for_tasks(self, timeout: float = 1.0):
with self._lock: with self._lock:
@@ -296,6 +350,7 @@ class TaskManager:
self.waiting_queue.clear() self.waiting_queue.clear()
self.active_tasks.clear() self.active_tasks.clear()
self._callbacks.clear() self._callbacks.clear()
self._tasks.clear()
def wake(self): def wake(self):
self._task_event.set() self._task_event.set()
+47
View File
@@ -157,6 +157,28 @@ def test_engine_generate_streaming_yields_tokens():
assert tokens == ["t1", "t2"] 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(): def test_engine_generate_async_yields_tokens_until_stop():
mock_model, mock_tokenizer = _make_engine_mocks(decode="tok") mock_model, mock_tokenizer = _make_engine_mocks(decode="tok")
callbacks_saved = [] callbacks_saved = []
@@ -186,6 +208,31 @@ def test_engine_generate_async_yields_tokens_until_stop():
assert asyncio.run(collect()) == ["t1", "t2"] 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(): def test_engine_generate_non_streaming_batch():
mock_model, mock_tokenizer = _make_engine_mocks(decode="r") mock_model, mock_tokenizer = _make_engine_mocks(decode="r")
+99
View File
@@ -1,6 +1,7 @@
"""Tests for scheduler concurrency.""" """Tests for scheduler concurrency."""
import threading import threading
import time
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
@@ -259,6 +260,104 @@ def _make_real_scheduler(device):
return scheduler, tokenizer, model 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): def test_run_batch_returns_token_sequences(device):
scheduler, _tok, _model = _make_real_scheduler(device) scheduler, _tok, _model = _make_real_scheduler(device)
try: try:
+16 -6
View File
@@ -166,11 +166,15 @@ def test_messages_with_system(client, loaded_model):
def test_chat_completions_stop_sequence(client, loaded_model): def test_chat_completions_stop_sequence(client, loaded_model):
"""POST /v1/chat/completions with stop parameter truncates at stop sequence.""" """POST /v1/chat/completions with stop parameter truncates at stop sequence."""
closed = []
async def async_gen(): async def async_gen():
yield "Hello" try:
yield "X" yield "Hello"
yield "world" yield "X"
yield "world"
finally:
closed.append(True)
get_app().state.engine = loaded_model get_app().state.engine = loaded_model
loaded_model.generate_async.return_value = async_gen() 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"] content = data["choices"][0]["message"]["content"]
assert "X" in content assert "X" in content
assert "world" not in content assert "world" not in content
assert closed == [True]
def test_chat_completions_stop_sequence_stream(client, loaded_model): def test_chat_completions_stop_sequence_stream(client, loaded_model):
"""POST /v1/chat/completions with stop parameter truncates SSE stream.""" """POST /v1/chat/completions with stop parameter truncates SSE stream."""
closed = []
async def async_gen(): async def async_gen():
yield "Hello" try:
yield "X" yield "Hello"
yield "world" yield "X"
yield "world"
finally:
closed.append(True)
get_app().state.engine = loaded_model get_app().state.engine = loaded_model
loaded_model.generate_async.return_value = async_gen() loaded_model.generate_async.return_value = async_gen()
@@ -217,6 +226,7 @@ def test_chat_completions_stop_sequence_stream(client, loaded_model):
assert any( assert any(
"finish_reason" in line for line in content.split("\n") if "stop" in line "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): def test_chat_completions_real_engine(tmp_path, client):
+9 -3
View File
@@ -73,15 +73,21 @@ def test_task_manager_remove_task():
assert len(tm.waiting_queue) == 0 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()) tm = TaskManager(tokenizer=_make_mock_tokenizer())
tid = tm.add_task("test") tid = tm.add_task("test")
tasks = tm.pull_candidates(1) tasks = tm.pull_candidates(1)
tm.activate(tasks[0]) tm.activate(tasks[0])
assert len(tm.active_tasks) == 1 assert len(tm.active_tasks) == 1
removed = tm.remove_task(tid) immediate, cancelled = tm.cancel_task(tid)
assert len(removed) == 1 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 len(tm.active_tasks) == 0
assert tm.get_stats()["cancelled_total"] == 1
def test_task_manager_pull_candidates_fifo(): def test_task_manager_pull_candidates_fifo():