Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
106 changes: 47 additions & 59 deletions flexkv/kvtask.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,6 @@
from typing import Callable
import multiprocessing as mp
import copy
from expiring_dict import ExpiringDict
import nvtx
import numpy as np

Expand Down Expand Up @@ -99,6 +98,8 @@ class KVTask:
prefetch_has_swa_remote: bool = False
prefetch_namespace: Optional[List[str]] = None
prefetch_swa_aware: bool = False
cancel_requested: bool = False
prefetch_key: Optional[int] = None

def is_completed(self) -> bool:
return self.status in [TaskStatus.COMPLETED, TaskStatus.CANCELLED, TaskStatus.FAILED]
Expand All @@ -111,6 +112,8 @@ def shed_heavy_resources(self) -> None:
self.slot_mapping = None
self.token_mask = None
self.callback = None
self.op_callback_dict = {}
self.swa_slot_mapping = None

TASK_STATUS_TO_RESPONSE_STATUS = {
TaskStatus.COMPLETED: KVResponseStatus.SUCCESS,
Expand Down Expand Up @@ -214,10 +217,12 @@ def __init__(self,
self.transfer_handles.append(remote_handle)
remote_handle._handle.send_config_to_remotes()

self.tasks: ExpiringDict[int, KVTask] = ExpiringDict(ttl=1800) # 30 minutes
# Active transfers own blocks and callbacks until graph completion.
# A TTL must never evict that ownership while DMA is still in flight.
self.tasks: Dict[int, KVTask] = {}

# hash(token_ids) -> task_id
self.prefetch_tasks: ExpiringDict[int, int] = ExpiringDict(ttl=1800) # 30 minutes
self.prefetch_tasks: Dict[int, int] = {} # removed with the owning task
self._gen_prefetch_key = lambda token_ids, namespace: hash_token(token_ids, namespace)

self.graph_to_task: Dict[int, int] = {}
Expand Down Expand Up @@ -424,7 +429,9 @@ def create_prefetch_task(self,
prefetch_namespace=namespace,
prefetch_swa_aware=swa_aware)

self.prefetch_tasks[self._gen_prefetch_key(token_ids, namespace)] = task_id
key = self._gen_prefetch_key(token_ids, namespace)
self.tasks[task_id].prefetch_key = key
self.prefetch_tasks[key] = task_id

self.graph_to_task[graph.graph_id] = task_id
self._log_task_created(self.tasks[task_id])
Expand Down Expand Up @@ -536,14 +543,6 @@ def _update_tasks(self, timeout: float = 0.001) -> None:
completed_op.xfer_ms,
completed_op.e2e_ms,
)
if task.status == TaskStatus.CANCELLED and task.callback is None:
# Cache was reset while this task was in flight: reset_cache()
# cleared its callbacks and freed the radix nodes / mempool blocks
flexkv_logger.warning(
f"task {task_id}: transfer op {completed_op.op_id} completed "
"after reset_cache(), callback no longer exists and will be skipped."
)
continue
has_callback = completed_op.op_id in task.op_callback_dict
if has_callback:
try:
Expand Down Expand Up @@ -737,30 +736,28 @@ def _fail_task(self, task_id: int) -> None:
flexkv_logger.error(f"[KVTaskEngine] task {task_id} FAILED: a transfer "
f"op of graph {task.graph.graph_id} failed")
self._abort_task_plans(task)
task.status = TaskStatus.FAILED
task.status = TaskStatus.CANCELLED if task.cancel_requested else TaskStatus.FAILED
self._log_task_terminal(task, task.status)
task.task_end_op_finished = True
self.graph_to_task.pop(task.graph.graph_id, None)
task.shed_heavy_resources()
if task.request_returned:
self._release_task(task_id)

def _cancel_task(self, task_id: int) -> None:
if task_id not in self.tasks:
task = self.tasks.get(task_id)
if task is None:
return
task.request_returned = True
if task.status == TaskStatus.RUNNING:
# Cancellation cannot stop submitted DMA. Keep graph mappings,
# plan pins and allocations until every transfer handle drains.
task.cancel_requested = True
return
task = self.tasks[task_id]
if not task.is_completed():
# A task whose graph never launched still holds everything its
# plan acquired at create time: locked radix nodes and staging
# blocks that only a completion callback could mount on the tree.
# Dropping the task without aborting leaks all of it -- the staging
# blocks become unreachable (mempool exhaustion) and the pinned
# nodes stay unevictable. Abort rolls those back; RUNNING tasks
# keep the old behavior (their graph is in flight and completion
# callbacks will still fire).
if task.status in (TaskStatus.UNREADY, TaskStatus.READY):
self._abort_task_plans(task)
self._abort_task_plans(task)
task.status = TaskStatus.CANCELLED
self._log_task_terminal(task, TaskStatus.CANCELLED)
self._log_task_terminal(task, task.status)
self._release_task(task_id)

def check_completed(self, task_id: int, completely: bool = False) -> bool:
Expand All @@ -777,7 +774,7 @@ def check_completed(self, task_id: int, completely: bool = False) -> bool:
# A partial-capable backend may finish the data-path sink after already
# reporting failed blocks. Wait for graph completion so cleanup runs and
# the caller observes FAILED instead of an early RUNNING-as-success result.
if task.transfer_failed:
if task.transfer_failed or task.cancel_requested:
return task.is_completed()
# For tasks with callback (e.g., PUT tasks that need to call insert_and_publish),
# we must wait until _mark_completed is called (i.e., is_completed() returns True)
Expand Down Expand Up @@ -846,13 +843,17 @@ def _release_task(self, task_id: int) -> None:
task = self.tasks[task_id]
if task.graph is not None:
self.graph_to_task.pop(task.graph.graph_id, None)
if task.prefetch_key is not None and self.prefetch_tasks.get(task.prefetch_key) == task_id:
self.prefetch_tasks.pop(task.prefetch_key, None)
self.tasks.pop(task_id, None)

def _mark_completed(self, task_id: int) -> None:
task = self.tasks[task_id]
if task.is_completed():
return
if task.callback:
if task.cancel_requested:
self._abort_task_plans(task)
elif task.callback:
callbacks = (
task.callback if isinstance(task.callback, list)
else [task.callback]
Expand All @@ -872,9 +873,10 @@ def _mark_completed(self, task_id: int) -> None:
if task.task_type == TaskType.PREFETCH:
self._finalize_prefetch_return_mask(task)
task.status = (
TaskStatus.CANCELLED if task.cancel_requested else
TaskStatus.FAILED if task.transfer_failed else TaskStatus.COMPLETED)
task.task_end_op_finished = True
self._log_task_terminal(task, TaskStatus.COMPLETED)
self._log_task_terminal(task, task.status)
self.graph_to_task.pop(task.graph.graph_id, None)
task.shed_heavy_resources()
if task.request_returned:
Expand Down Expand Up @@ -1317,7 +1319,9 @@ def prefetch_async(self,
if task_id == -1:
task_id = self._gen_task_id()
nvtx.push_range(f"prefetch match: task_id={task_id}", color=get_nvtx_default_color())
self.create_prefetch_task(task_id, token_ids, dp_client_id=dp_client_id, namespace=namespace, swa_aware=swa_aware)
self.create_prefetch_task(
task_id, token_ids, dp_client_id=dp_client_id, namespace=namespace, swa_aware=swa_aware
)
self._process_empty_graph(task_id)
nvtx.pop_range()
# trace prefetch async request
Expand Down Expand Up @@ -1467,37 +1471,21 @@ def reset_cache(self) -> None:
Used after a weight update (e.g. verl RL rollout) so that KV computed
against stale weights is never reused.

We do NOT drain or cancel in-flight transfers here. verl issues the
reset at a rollout/weight-update boundary where no new generation
requests are being served, so in practice no task is ongoing. If any
task IS still in flight we only warn: resetting the radix tree +
mempool is cheap and the stale-weight invalidation must not be blocked
on transfer completion. This mirrors vLLM's own reset_encoder_cache /
reset_mm_cache, which likewise only warn on has_unfinished_requests().
The caller must drain transfers first. Resetting a busy cache can
recycle destinations while DMA still writes to them; clearing Python
callbacks cannot make that safe.
"""
ongoing = sum(1 for t in list(self.tasks.values()) if not t.is_completed())
ongoing = sum(1 for t in self.tasks.values() if not t.is_completed())
if ongoing:
flexkv_logger.warning(
f"reset_cache called while {ongoing} task(s) are still in flight; "
f"resetting anyway. In-flight transfers may target blocks that are "
f"being freed — ensure reset is issued at a quiesced boundary."
)

# Invalidate in-flight tasks' callbacks BEFORE dropping the cache. The
# callback / op_callback_dict partials close over the exact radix nodes
# and mempool blocks we are about to free; if a late transfer fired them
# after reset they would unlock or mount onto a deleted node, or recycle a
# block into the freshly-emptied mempool (which raises "already free").
# Clear them here so the dispatch in _update_tasks has nothing to fire.
# Note: reset_cache() runs on the same thread as the callback dispatch
# (_update_tasks), so no lock is needed. We keep the graph_to_task
# mapping so a late-completing op still resolves to its task and warns.
for task_id, task in list(self.tasks.items()):
if task.is_completed():
continue # already-fired callbacks are harmless
task.callback = None
task.op_callback_dict = {}
task.status = TaskStatus.CANCELLED
raise RuntimeError(f"reset_cache requires drained transfers; {ongoing} task(s) active")

# Only terminal response metadata remains after the drain boundary.
self.tasks.clear()
self.prefetch_tasks.clear()
self.graph_to_task.clear()
self.uncompleted_ops.clear()
self.uncompleted_op_results.clear()
self.uncompleted_graphs.clear()

# Drop index (radix tree) + mempool on every tier. CRadixTreeIndex (C++)
# and the pure-Python RadixTreeIndex both expose reset(); GlobalCacheEngine
Expand Down
144 changes: 144 additions & 0 deletions tests/test_task_ownership.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,144 @@
"""CPU task ownership regressions; no DMA or worker process is launched."""

from types import SimpleNamespace as NS
from unittest.mock import Mock, patch

import numpy as np
import pytest

from flexkv.kvtask import KVTask, KVTaskEngine, KVTaskManager, TaskStatus, TaskType
from flexkv.common.transfer import CompletedOp


def engine_with_task(status=TaskStatus.RUNNING, kind=TaskType.GET, request_returned=False):
engine = KVTaskEngine.__new__(KVTaskEngine)
plan = Mock()
task = KVTask(
7,
kind,
11,
False,
status,
np.arange(4),
None,
None,
NS(graph_id=13, num_ops=2, _op_map={}),
np.ones(4, dtype=bool),
plan,
{},
)
task.request_returned = request_returned
task.prefetch_key = 19
engine.tasks = {7: task}
engine.prefetch_tasks = {19: 7}
engine.graph_to_task = {13: 7}
engine.uncompleted_ops = {}
engine.uncompleted_op_results = {}
engine.uncompleted_graphs = {}
engine._log_task_terminal = Mock()
engine._finalize_prefetch_return_mask = Mock()
engine.tracer = Mock()
engine.cache_engine = Mock()
return engine, task, plan


@pytest.mark.parametrize("kind", list(TaskType))
@pytest.mark.parametrize("failed", [False, True])
def test_running_cancel_keeps_ownership_until_graph_drains(kind, failed):
engine, task, plan = engine_with_task(kind=kind)
engine.cancel_tasks([7])
engine.cancel_tasks([7]) # idempotent while draining
assert task.cancel_requested and task.status == TaskStatus.RUNNING
assert engine.tasks[7] is task and engine.graph_to_task == {13: 7}
assert engine.prefetch_tasks == {19: 7}
assert not plan.mock_calls
task.task_end_op_finished = True
assert not engine.check_completed(7) # early data-path completion is insufficient
if failed:
engine._fail_task(7)
else:
engine._mark_completed(7)
plan.abort.assert_called_once_with()
plan.assert_not_called()
assert task.status == TaskStatus.CANCELLED
assert engine.tasks == engine.graph_to_task == engine.prefetch_tasks == {}
assert task.op_callback_dict == {} and task.graph is None


@pytest.mark.parametrize("status", [TaskStatus.UNREADY, TaskStatus.READY])
def test_unlaunched_cancel_releases_immediately(status):
engine, task, plan = engine_with_task(status=status)
engine._cancel_task(7)
plan.abort.assert_called_once_with()
assert engine.tasks == engine.prefetch_tasks == engine.graph_to_task == {}


def test_terminal_response_wait_recycles_all_indexes():
engine, task, plan = engine_with_task()
engine._mark_completed(7)
assert engine.tasks[7] is task
engine._update_tasks = Mock()
result = engine.try_wait([7])
assert result[7].status.value == "success"
plan.assert_called_once_with()
assert engine.tasks == engine.prefetch_tasks == engine.graph_to_task == {}


def test_releasing_old_prefetch_does_not_remove_new_owner():
engine, task, plan = engine_with_task(status=TaskStatus.READY)
engine.prefetch_tasks[19] = 8
engine._cancel_task(7)
assert engine.prefetch_tasks == {19: 8}


@pytest.mark.parametrize("status", [TaskStatus.RUNNING, TaskStatus.READY, TaskStatus.UNREADY])
def test_busy_reset_has_no_side_effects(status):
engine, task, plan = engine_with_task(status=status)
with pytest.raises(RuntimeError, match="requires drained transfers"):
engine.reset_cache()
assert engine.tasks[7] is task
assert engine.graph_to_task == {13: 7}
assert not engine.cache_engine.mock_calls and not plan.mock_calls


def test_drained_reset_clears_terminal_metadata():
engine, task, plan = engine_with_task()
engine._mark_completed(7)
engine.reset_cache()
engine.cache_engine.reset.assert_called_once_with()
assert engine.tasks == engine.prefetch_tasks == engine.graph_to_task == {}


def test_constructor_uses_explicit_task_ownership():
import flexkv.kvtask as module

config = NS(enable_cpu=True, enable_remote=True, enable_gds=False, enable_kv_sharing=False, enable_nixl=False)
model = NS(use_trtllm_subprocess=False, nnodes=1)
with patch.object(module, "GlobalCacheEngine"), patch.object(module, "TransferManagerHandle"):
manager = KVTaskManager(model, config)
assert type(manager.tasks) is dict and type(manager.prefetch_tasks) is dict


@pytest.mark.parametrize("failed", [False, True])
def test_cancel_waits_for_every_transfer_handle_before_abort(failed):
engine, task, plan = engine_with_task()
writer_commit = Mock()
task.op_callback_dict[11] = writer_commit
engine.required_completed_count = 2
first = [CompletedOp(13, 11), CompletedOp(13, -1)]
second = [CompletedOp(13, 11), CompletedOp(13, -1, failed=failed)]
engine.transfer_handles = [
NS(wait=Mock(side_effect=[first, []])),
NS(wait=Mock(side_effect=[[], second])),
]
engine.cancel_tasks([7])
engine._update_tasks(timeout=0)
assert engine.tasks[7] is task
plan.abort.assert_not_called()
writer_commit.assert_not_called()
engine._update_tasks(timeout=0)
writer_commit.assert_called_once() # completed writers can safely publish
plan.abort.assert_called_once_with()
plan.assert_not_called()
assert engine.tasks == engine.graph_to_task == engine.prefetch_tasks == {}
assert engine.uncompleted_graphs == engine.uncompleted_ops == {}
Loading