From 35a69a5c252bd74a1d77d0f4fbced44125027faf Mon Sep 17 00:00:00 2001 From: Ketor Date: Thu, 17 Sep 2026 17:24:36 +0800 Subject: [PATCH 1/3] fix(vllm): support legacy and native request failure protocols --- CMakeLists.txt | 8 + docs/CONNECTORS.md | 26 +- integration/vllm/README.md | 42 +- integration/vllm/src/dfkv_vllm/connector.py | 59 +- integration/vllm/src/dfkv_vllm/scheduler.py | 3 +- .../vllm/src/dfkv_vllm/transfer_protocol.py | 299 +++++++++ integration/vllm/src/dfkv_vllm/worker.py | 8 +- .../vllm/tests/test_request_load_failures.py | 85 ++- .../tests/test_scheduler_full_block_ids.py | 15 +- .../vllm/tests/test_transfer_protocol.py | 593 ++++++++++++++++++ test/python/test_dfkv_vllm_connector.py | 105 +++- test/python/test_dfkv_vllm_worker.py | 132 ++-- 12 files changed, 1244 insertions(+), 131 deletions(-) create mode 100644 integration/vllm/src/dfkv_vllm/transfer_protocol.py create mode 100644 integration/vllm/tests/test_transfer_protocol.py diff --git a/CMakeLists.txt b/CMakeLists.txt index f3adfb4..4b0f11b 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -264,6 +264,14 @@ if(DFKV_BUILD_TESTS) -p test_client_ranks.py) set_tests_properties(python_vllm_client_ranks PROPERTIES ENVIRONMENT "PYTHONPATH=${CMAKE_CURRENT_SOURCE_DIR}/integration/common/src:${CMAKE_CURRENT_SOURCE_DIR}/integration/vllm/src") + add_test(NAME python_vllm_connector_lifecycle + COMMAND ${PYTHON3} + ${CMAKE_CURRENT_SOURCE_DIR}/test/python/test_dfkv_vllm_connector.py + WORKING_DIRECTORY ${CMAKE_CURRENT_SOURCE_DIR}) + add_test(NAME python_vllm_worker_lifecycle + COMMAND ${PYTHON3} + ${CMAKE_CURRENT_SOURCE_DIR}/test/python/test_dfkv_vllm_worker.py + WORKING_DIRECTORY ${CMAKE_CURRENT_SOURCE_DIR}) # Shared telemetry layer: pure-python, no build/lib needed (OTel-push test # self-skips when the SDK is absent). Includes the vendored-copy drift guard. add_test(NAME python_telemetry diff --git a/docs/CONNECTORS.md b/docs/CONNECTORS.md index 96a8084..651307a 100644 --- a/docs/CONNECTORS.md +++ b/docs/CONNECTORS.md @@ -882,14 +882,24 @@ lsmod | grep nvidia_peermem #### 混合模型的 KV 加载故障恢复 -- 混合/多 group 和单 group 非 full-attention 布局使用 vLLM 原生请求级失败 - 协议,要求引擎提供 `KVConnectorTransferResults.failed_recving`、 - `KVConnectorOutput.failed_recving` 的完整透传,以及对 - `WAITING_FOR_REMOTE_KVS` 请求的调度器恢复处理。请升级到具备完整 API 的 - 引擎;不再提供或要求外置 invalid-blocks 调度器补丁,也不按未经验证的 - 版本号推定支持。 -- 这些布局按 request ID 报错,`invalid_block_ids` 留空;只有原生 I/O 和 GPU - 写入都经过终止 fence 后,才同时报告 `failed_recving` 与 `finished_recving`。 +- 一个 connector wheel 自动适配两类引擎,不按未经验证的版本号推定支持。 + 混合/多 group 和单 group 非 full-attention 布局始终按 request ID 报错: + **原生路径**通过 `get_transfer_results()` 返回 + `KVConnectorTransferResults.failed_recving`,由引擎透传 + `KVConnectorOutput.failed_recving` 并恢复等待请求,不安装调度器桥接。 + **旧引擎路径**通过 `get_finished()` 收集同一份 worker 结果,将失败 ID + 保存在 `build_connector_worker_meta()` 返回的 worker metadata 中;多次 + poll 不会丢失尚未上报的失败,原生路径不会重复发送这份 metadata。 +- 旧引擎必须具备 `KVConnectorWorkerMetadata.aggregate()` 及其传输、executor + 跨 rank receive-completion 聚合,以及调度器 `update_from_output`、 + `_handle_invalid_blocks`、等待请求恢复和 load-failure policy 接口。 + 仅在旧引擎的 scheduler-role connector 使用请求级布局时,自动安装经过 + 能力检查的进程内桥接;缺失所需能力时启动明确报错。无需修改已安装的 + vLLM 源文件,也无需手动打引擎补丁,不影响普通 block error 或其他连接器。 + 桥接跨 step 保留失败,直到聚合后的 `finished_recving` 确认所有 rank 完成, + 才按请求身份交给既有恢复/终止路径,绝不伪造或扫描混合布局的 block ID。 +- 这些布局的 worker `invalid_block_ids` 留空;只有原生 I/O 和 GPU + 写入都经过终止 fence 后,才发布失败与 receive completion。 `kv_load_failure_policy="recompute"` 时,引擎释放失败分配并重新本地准入; dfkv 丢弃旧 lookup、load spec 和 tracker,按新分配及实际计算前缀重建 SAVE 状态。失败请求在本次生命周期内(包括后续抢占)不再查询远端命中,避免同一 diff --git a/integration/vllm/README.md b/integration/vllm/README.md index c478a2c..4774013 100644 --- a/integration/vllm/README.md +++ b/integration/vllm/README.md @@ -41,18 +41,36 @@ device pointers. Set `DFKV_RDMA=1` in every engine process. Construction rejects and closes any native handle whose reported transport is not `rdma`; there is no TCP or host-bounce fallback for this connector. -Hybrid/multi-group caches and single-group non-full-attention caches require -vLLM's native request-level load-failure protocol: -`KVConnectorTransferResults.failed_recving`, its propagation through -`KVConnectorOutput.failed_recving`, and scheduler recovery of failed -`WAITING_FOR_REMOTE_KVS` requests. Upgrade to an engine exposing that complete -API; no out-of-tree invalid-block scheduler patch is supported or required. -The connector reports failed request IDs together with receive completion only -after native I/O and GPU writes are fenced, leaving block-ID errors empty for -these layouts. With `kv_load_failure_policy="recompute"`, the engine releases the -failed allocation and retries locally; dfkv bypasses the failed remote source -for the remainder of that request, including subsequent preemptions. The engine's -`"fail"` policy terminates the affected request instead. +Hybrid/multi-group caches and single-group non-full-attention caches report +request-level load failures. One connector wheel automatically selects the +engine's protocol: + +- **Native engines:** `get_transfer_results()` returns + `KVConnectorTransferResults.failed_recving` alongside receive completion. + The engine propagates `KVConnectorOutput.failed_recving` and performs its + native scheduler recovery. No scheduler bridge is installed. +- **Legacy engines:** `get_finished()` uses the same worker outcome collection; + `build_connector_worker_meta()` carries failed request IDs through + `KVConnectorWorkerMetadata.aggregate()`. For request-level layouts, the + scheduler-role connector installs a capability-checked, in-process bridge. + It retains failures across steps until executor-aggregated `finished_recving` + confirms all ranks are done, then routes request identities through the + engine's existing recovery/error path without inventing block IDs. + +The legacy bridge requires worker-metadata aggregation/transport, all-rank +receive-completion aggregation, and the scheduler's `update_from_output`, +`_handle_invalid_blocks`, waiting-request recovery and load-failure policy +hooks. Missing bridge capabilities fail explicitly at startup. It does not +modify installed vLLM source files or require a manually applied engine patch; +compatibility is based on these capabilities, not a blanket version cutoff. +Single-group full-attention retains its existing block-level recovery. + +Failed request IDs and receive completion are published only after native I/O +and GPU writes are fenced. These layouts leave block-ID errors empty. With +`kv_load_failure_policy="recompute"`, the engine releases the failed allocation +and retries locally; dfkv bypasses the failed remote source for the remainder +of that request, including subsequent preemptions. The engine's `"fail"` policy +terminates the affected request instead. The request-level completion fence is mandatory even when the legacy `DFKV_GPU_LOAD_FENCE=0` override is present. diff --git a/integration/vllm/src/dfkv_vllm/connector.py b/integration/vllm/src/dfkv_vllm/connector.py index dc8f619..e2f68a3 100644 --- a/integration/vllm/src/dfkv_vllm/connector.py +++ b/integration/vllm/src/dfkv_vllm/connector.py @@ -27,7 +27,6 @@ KVConnectorBase_V1, KVConnectorMetadata, KVConnectorRole, - KVConnectorTransferResults, SupportsHMA, ) from vllm.distributed.kv_transfer.kv_connector.v1.metrics import ( @@ -45,9 +44,19 @@ from vllm.v1.outputs import KVConnectorOutput from vllm.v1.request import Request -from .data import DfkvStoreConnectorMetadata, VLLM_RAW_LAYOUT +from .data import ( + DfkvStoreConnectorMetadata, + VLLM_RAW_LAYOUT, + requires_request_level_loads, +) from .metrics import DfkvStoreConnectorStats, DfkvStorePromMetrics from .scheduler import DfkvStoreScheduler +from .transfer_protocol import ( + HAS_NATIVE_TRANSFER_RESULTS, + LegacyReceiveFailures, + TransferResults, + install_legacy_failure_bridge, +) from .worker import DfkvStoreWorker logger = init_logger(__name__) @@ -144,11 +153,17 @@ def __init__( self._inflight_calls = 0 self._shutdown = False self._shutdown_complete = False + self._legacy_failed_recving: set[str] = set() self.connector_scheduler: DfkvStoreScheduler | None = None self.connector_worker: DfkvStoreWorker | None = None if role == KVConnectorRole.SCHEDULER: + if ( + not HAS_NATIVE_TRANSFER_RESULTS + and requires_request_level_loads(kv_cache_config) + ): + install_legacy_failure_bridge() self.connector_scheduler = DfkvStoreScheduler( vllm_config, kv_cache_config ) @@ -361,15 +376,42 @@ def wait_for_save(self): finally: self._finish_call() - def get_transfer_results( + def _collect_transfer_results(self, finished_req_ids: set[str]) -> TransferResults: + assert self.connector_worker is not None + metadata = self._get_connector_metadata() + assert isinstance(metadata, DfkvStoreConnectorMetadata) + return self.connector_worker.get_transfer_results(finished_req_ids, metadata) + + def get_transfer_results(self, finished_req_ids: set[str]) -> TransferResults: + """Return native outcomes without also publishing legacy metadata.""" + self._begin_call() + try: + return self._collect_transfer_results(finished_req_ids) + finally: + self._finish_call() + + def get_finished( self, finished_req_ids: set[str] - ) -> KVConnectorTransferResults: + ) -> tuple[set[str], set[str]]: + """Adapt one outcome snapshot to the legacy completion/metadata hooks.""" self._begin_call() try: - assert self.connector_worker is not None - metadata = self._get_connector_metadata() - assert isinstance(metadata, DfkvStoreConnectorMetadata) - return self.connector_worker.get_transfer_results(finished_req_ids, metadata) + results = self._collect_transfer_results(finished_req_ids) + with self._shutdown_condition: + self._legacy_failed_recving.update(results.failed_recving) + return results.finished_sending, results.finished_recving + finally: + self._finish_call() + + def build_connector_worker_meta(self) -> LegacyReceiveFailures | None: + self._begin_call() + try: + with self._shutdown_condition: + failed_recving = self._legacy_failed_recving + if not failed_recving: + return None + self._legacy_failed_recving = set() + return LegacyReceiveFailures(failed_recving=failed_recving) finally: self._finish_call() @@ -432,6 +474,7 @@ def shutdown(self) -> None: errors.append(error) finally: with self._shutdown_condition: + self._legacy_failed_recving.clear() self._shutdown_complete = True self._shutdown_condition.notify_all() diff --git a/integration/vllm/src/dfkv_vllm/scheduler.py b/integration/vllm/src/dfkv_vllm/scheduler.py index 1899a0c..1f31e88 100644 --- a/integration/vllm/src/dfkv_vllm/scheduler.py +++ b/integration/vllm/src/dfkv_vllm/scheduler.py @@ -28,6 +28,7 @@ RequestTracker, requires_request_level_loads, ) +from .transfer_protocol import failed_requests_from_output from .worker import ( LookupKeyClient, ) @@ -190,7 +191,7 @@ def update_connector_output(self, output: KVConnectorOutput) -> None: """Prevent a failed remote source from being re-admitted on retry.""" if not self.request_level_loads: return - for req_id in output.failed_recving or (): + for req_id in failed_requests_from_output(output): # A receive can finish after an aborted request's final metadata. # Such an output must not recreate scheduler-side request state. if req_id not in self._unfinished_request_ids: diff --git a/integration/vllm/src/dfkv_vllm/transfer_protocol.py b/integration/vllm/src/dfkv_vllm/transfer_protocol.py new file mode 100644 index 0000000..8d72c35 --- /dev/null +++ b/integration/vllm/src/dfkv_vllm/transfer_protocol.py @@ -0,0 +1,299 @@ +"""Native transfer snapshots and a metadata bridge for older vLLM schedulers.""" + +from dataclasses import dataclass, field, fields +from functools import wraps +from inspect import signature +from weakref import WeakValueDictionary +from threading import Lock + +from vllm.distributed.kv_transfer.kv_connector.v1 import base +from vllm.v1.outputs import KVConnectorOutput + + +def _require_method(owner, name: str, *parameters: str) -> None: + method = getattr(owner, name, None) + if not callable(method): + raise RuntimeError(f"dfkv requires vLLM {owner.__name__}.{name}") + actual = signature(method).parameters + if not all(parameter in actual for parameter in parameters): + raise RuntimeError(f"Unsupported vLLM {owner.__name__}.{name} signature") + + +def _require_fields(cls, names: set[str]) -> None: + if not names.issubset(getattr(cls, "__dataclass_fields__", {})): + raise RuntimeError(f"Unsupported vLLM {cls.__name__} transfer fields") + + +_native_results = getattr(base, "KVConnectorTransferResults", None) +_native_method = getattr(base.KVConnectorBase_V1, "get_transfer_results", None) +_native_output = "failed_recving" in {item.name for item in fields(KVConnectorOutput)} +HAS_NATIVE_TRANSFER_RESULTS = _native_results is not None +if HAS_NATIVE_TRANSFER_RESULTS or _native_method is not None or _native_output: + if not (HAS_NATIVE_TRANSFER_RESULTS and callable(_native_method) and _native_output): + raise RuntimeError("Incomplete native vLLM request-level transfer protocol") + _require_fields( + _native_results, {"finished_sending", "finished_recving", "failed_recving"}, + ) + _require_method(base.KVConnectorBase_V1, "get_transfer_results", "finished_req_ids") + TransferResults = _native_results +else: + _require_method(base.KVConnectorBase_V1, "get_finished", "finished_req_ids") + _require_method(base.KVConnectorBase_V1, "build_connector_worker_meta") + _require_method(base.KVConnectorBase_V1, "update_connector_output", "connector_output") + _require_fields( + KVConnectorOutput, + {"finished_recving", "invalid_block_ids", "kv_connector_worker_meta"}, + ) + + @dataclass + class TransferResults: + """One worker snapshot; failed receives also finish receiving.""" + + finished_sending: set[str] = field(default_factory=set) + finished_recving: set[str] = field(default_factory=set) + failed_recving: set[str] = field(default_factory=set) + + +_require_method(base.KVConnectorWorkerMetadata, "aggregate", "other") + + +@dataclass +class LegacyReceiveFailures(base.KVConnectorWorkerMetadata): + """Worker errors, unioned by vLLM's supported metadata aggregation.""" + + failed_recving: set[str] = field(default_factory=set) + # Only the scheduler bridge marks a local callback packet ready. Worker + # metadata alone does not prove completion on every worker. + _ready: bool = field(default=False, init=False, repr=False, compare=False) + + def aggregate(self, other: base.KVConnectorWorkerMetadata): + if not isinstance(other, LegacyReceiveFailures): + raise TypeError("Cannot aggregate dfkv receive failures with other metadata") + self.failed_recving.update(other.failed_recving) + self._ready = False + return self + + +def failed_requests_from_output(output) -> set[str]: + """Return native failures or failures fenced by the legacy scheduler.""" + native = getattr(output, "failed_recving", None) + if native is not None: + return native + metadata = getattr(output, "kv_connector_worker_meta", None) + if isinstance(metadata, LegacyReceiveFailures) and metadata._ready: + return metadata.failed_recving + return set() + + +class _RequestErrorBatch(set): + """Local dispatch only, never sent through the executor or worker wire. + + Its elements are exclusively the original real block IDs. The private + request set makes the old invalid-block guard dispatch even when there are + no block errors; it is never interpreted as, or converted to, block IDs. + """ + + def __init__(self, block_ids: set[int], failed_recving: set[str]): + super().__init__(block_ids) + self.block_ids = block_ids + self.failed_recving = failed_recving + + def __bool__(self): + return bool(self.block_ids or self.failed_recving) + + +class _PendingFailures: + def __init__(self, connector): + self.connector = connector + # Remember the actual request, not just a reusable string ID, without + # keeping removed requests (and their prompt/cache state) alive. + self.requests = WeakValueDictionary() + # This tree contains only dfkv failure leaves and None foreign leaves. + # Retaining a whole worker packet would replay other connectors' state. + self.metadata = None + + def prune(self, scheduler, waiting_status) -> set[str]: + replaced = set() + for req_id, request in list(self.requests.items()): + if scheduler.requests.get(req_id) is not request: + replaced.add(req_id) + if ( + req_id in replaced + or request.status != waiting_status + ): + del self.requests[req_id] + return replaced + +_install_lock = Lock() + + +def install_legacy_failure_bridge() -> None: + """Install legacy hooks once, including concurrent engine construction.""" + with _install_lock: + _install_legacy_failure_bridge() + + +def _install_legacy_failure_bridge() -> None: + """Install the two legacy dispatch hooks once; native engines stay untouched.""" + if HAS_NATIVE_TRANSFER_RESULTS: + return + + # Lazy import: connector discovery can occur while vLLM imports Scheduler. + from vllm.distributed.kv_transfer.kv_connector.v1.multi_connector import ( + MultiKVConnectorWorkerMetadata, + ) + from vllm.v1.core.sched.scheduler import Scheduler + from vllm.v1.request import RequestStatus + + if getattr(Scheduler.update_from_output, "_dfkv_legacy_failure_bridge", False): + return + for name, parameters in ( + ("update_from_output", ("scheduler_output", "model_runner_output")), + ("_handle_invalid_blocks", ("invalid_block_ids", "num_scheduled_tokens")), + ("_update_from_kv_xfer_finished", ("kv_connector_output",)), + ("_update_waiting_for_remote_kv", ("request",)), + ("finish_requests", ("request_ids", "finished_status")), + ): + _require_method(Scheduler, name, *parameters) + if not hasattr(RequestStatus, "WAITING_FOR_REMOTE_KVS"): + raise RuntimeError("Unsupported vLLM asynchronous receive lifecycle") + + original_update = Scheduler.update_from_output + original_invalid = Scheduler._handle_invalid_blocks + waiting = RequestStatus.WAITING_FOR_REMOTE_KVS + + def failure_leaves(metadata): + if isinstance(metadata, LegacyReceiveFailures): + yield metadata + elif isinstance(metadata, MultiKVConnectorWorkerMetadata): + for child in metadata.metadata: + yield from failure_leaves(child) + + def graft_failures(metadata, failures, request_ids, *, ready=False): + """Copy only allowed failure leaves, preserving current foreign data.""" + if isinstance(metadata, LegacyReceiveFailures) or ( + metadata is None and isinstance(failures, LegacyReceiveFailures) + ): + leaf = LegacyReceiveFailures( + failures.failed_recving.intersection(request_ids) + if isinstance(failures, LegacyReceiveFailures) else set() + ) + leaf._ready = ready + return leaf + if isinstance(metadata, MultiKVConnectorWorkerMetadata) or ( + metadata is None and isinstance(failures, MultiKVConnectorWorkerMetadata) + ): + template = metadata if metadata is not None else failures + children = tuple( + graft_failures( + metadata.metadata[index] if metadata is not None else None, + failures.metadata[index] + if isinstance(failures, MultiKVConnectorWorkerMetadata) else None, + request_ids, + ready=ready, + ) + for index in range(len(template.metadata)) + ) + if metadata is not None and all( + child is original for child, original in zip(children, metadata.metadata) + ): + return metadata + return MultiKVConnectorWorkerMetadata(metadata=children) + return metadata + + @wraps(original_invalid) + def handle_invalid(self, invalid_block_ids, num_scheduled_tokens): + if not isinstance(invalid_block_ids, _RequestErrorBatch): + return original_invalid(self, invalid_block_ids, num_scheduled_tokens) + batch = invalid_block_ids + affected = ( + original_invalid(self, batch.block_ids, num_scheduled_tokens) + if batch.block_ids else set() + ) + for req_id in batch.failed_recving: + request = self.requests.get(req_id) + if request is None or request.status != waiting: + continue + if self.recompute_kv_load_failures: + request.num_computed_tokens = 0 + self.failed_recving_kv_req_ids.add(req_id) + else: + affected.add(req_id) + return affected + + @wraps(original_update) + def update(self, scheduler_output, model_runner_output): + packet = model_runner_output.kv_connector_output + metadata = getattr(packet, "kv_connector_worker_meta", None) + pending = getattr(self, "_dfkv_pending_receive_failures", None) + if pending is not None and pending.connector is not self.connector: + del self._dfkv_pending_receive_failures + pending = None + if metadata is None and pending is None: + return original_update(self, scheduler_output, model_runner_output) + leaves = tuple(failure_leaves(metadata)) + if not leaves and pending is None: + return original_update(self, scheduler_output, model_runner_output) + if pending is None: + pending = _PendingFailures(self.connector) + self._dfkv_pending_receive_failures = pending + + replaced = pending.prune(self, waiting) + pending.metadata = graft_failures(None, pending.metadata, pending.requests) + for leaf in leaves: + for req_id in leaf.failed_recving: + request = self.requests.get(req_id) + if ( + req_id not in replaced and request is not None + and request.status == waiting + ): + pending.requests[req_id] = request + if leaves: + incoming = graft_failures(None, metadata, pending.requests) + pending.metadata = ( + pending.metadata.aggregate(incoming) + if pending.metadata is not None else incoming + ) + + if packet is None: + try: + return original_update(self, scheduler_output, model_runner_output) + finally: + pending.prune(self, waiting) + if not pending.requests: + del self._dfkv_pending_receive_failures + + ready = set(pending.requests).intersection(packet.finished_recving or ()) + original_blocks = packet.invalid_block_ids + original_finished = packet.finished_recving + callback_metadata = graft_failures( + metadata, pending.metadata, ready, ready=True, + ) + try: + # Ignore already-retired receive notifications only in this dfkv + # scope. Keep terminal requests still owned by the engine: its + # existing completion path must free their deferred allocations. + packet.finished_recving = { + req_id for req_id in original_finished or () + if req_id not in replaced + and (request := self.requests.get(req_id)) is not None + and (request.status == waiting or request.is_finished()) + } + packet.kv_connector_worker_meta = callback_metadata + if ready: + packet.invalid_block_ids = _RequestErrorBatch(original_blocks, ready) + result = original_update(self, scheduler_output, model_runner_output) + for req_id in ready: + pending.requests.pop(req_id, None) + return result + finally: + packet.invalid_block_ids = original_blocks + packet.finished_recving = original_finished + packet.kv_connector_worker_meta = metadata + pending.prune(self, waiting) + if not pending.requests: + del self._dfkv_pending_receive_failures + + update._dfkv_legacy_failure_bridge = True + Scheduler._handle_invalid_blocks = handle_invalid + Scheduler.update_from_output = update diff --git a/integration/vllm/src/dfkv_vllm/worker.py b/integration/vllm/src/dfkv_vllm/worker.py index 80544ea..bc2efa2 100644 --- a/integration/vllm/src/dfkv_vllm/worker.py +++ b/integration/vllm/src/dfkv_vllm/worker.py @@ -48,9 +48,6 @@ get_tensor_model_parallel_world_size, ) from vllm.distributed.kv_events import BlockStored -from vllm.distributed.kv_transfer.kv_connector.v1.base import ( - KVConnectorTransferResults, -) from vllm.logger import init_logger from vllm.utils.network_utils import make_zmq_socket from vllm.v1.core import kv_cache_utils @@ -91,6 +88,7 @@ from .dfkv_utils import get_dp_engine_index from .metrics import DfkvStoreConnectorStats from .rail_affinity import physical_affinity_rank +from .transfer_protocol import TransferResults from ._telemetry import config as _tcfg # connector identity + client_register switch from .protocol import ( LOOKUP_MSG, @@ -2196,7 +2194,7 @@ def get_transfer_results( self, finished_req_ids: set[str], meta: DfkvStoreConnectorMetadata, - ) -> KVConnectorTransferResults: + ) -> TransferResults: """Submit post-forward I/O and atomically collect receive outcomes. Mutable and windowed stores finish before the next model step can @@ -2226,7 +2224,7 @@ def get_transfer_results( done_recving, failed_recving = ( self.kv_recv_thread.get_and_clear_receive_results() ) - return KVConnectorTransferResults( + return TransferResults( finished_sending=done_sending, finished_recving=done_recving, failed_recving=failed_recving, diff --git a/integration/vllm/tests/test_request_load_failures.py b/integration/vllm/tests/test_request_load_failures.py index 67e4ea0..c930caa 100644 --- a/integration/vllm/tests/test_request_load_failures.py +++ b/integration/vllm/tests/test_request_load_failures.py @@ -19,6 +19,7 @@ ReqMeta, ) from dfkv_vllm.worker import DfkvStoreWorker, KVCacheStoreRecvingThread +from dfkv_vllm.transfer_protocol import LegacyReceiveFailures BLOCK = 16 @@ -255,13 +256,17 @@ def connector_for(receiver, *, inline): connector._shutdown_condition = threading.Condition() connector._shutdown = False connector._inflight_calls = 0 + connector._legacy_failed_recving = set() metadata = DfkvStoreConnectorMetadata({"load"}, set()) metadata.add_request(request()) connector.bind_connector_metadata(metadata) return connector -def test_inline_parked_load_fences_both_sides_and_reports_once(make_receiver, monkeypatch): +@pytest.mark.parametrize("legacy", [False, True]) +def test_inline_parked_load_fences_both_sides_and_reports_once( + make_receiver, monkeypatch, legacy, +): client = MemoryClient("miss") receiver, _ = make_receiver(client) receiver._cuda_device = 7 @@ -276,34 +281,88 @@ def synchronize(device): connector.start_load_kv(None) assert client.thread_ids == [threading.get_ident()] assert calls_seen_at_fence == [0, 1] - result = connector.get_transfer_results(set()) - assert result.finished_recving == result.failed_recving == {"load"} - assert result.finished_sending == set() + if legacy: + assert connector.get_finished(set()) == (set(), {"load"}) + else: + result = connector.get_transfer_results(set()) + assert result.finished_recving == result.failed_recving == {"load"} + assert result.finished_sending == set() # Repeated polling/hooks with the same metadata must not resubmit writes. connector.start_load_kv(None) - result = connector.get_transfer_results(set()) - assert result.finished_recving == result.failed_recving == set() + if legacy: + assert connector.get_finished(set()) == (set(), set()) + metadata = connector.build_connector_worker_meta() + assert isinstance(metadata, LegacyReceiveFailures) + assert metadata.failed_recving == {"load"} + else: + result = connector.get_transfer_results(set()) + assert result.finished_recving == result.failed_recving == set() + assert connector.build_connector_worker_meta() is None assert client.calls == 1 assert connector._inflight_calls == 0 assert connector.get_block_ids_with_load_errors() == set() -def test_pool_workers_emit_one_completion_and_never_resubmit_polled_metadata(make_receiver): +@pytest.mark.parametrize("legacy", [False, True]) +def test_pool_workers_emit_one_completion_and_never_resubmit_polled_metadata( + make_receiver, legacy, +): client = MemoryClient() receiver, pool = make_receiver(client, workers=3) connector = connector_for(receiver, inline=False) receiver.start() connector.start_load_kv(None) assert client.calls == 0 - first = connector.get_transfer_results(set()) + def poll(): + if legacy: + return connector.get_finished(set())[1] + result = connector.get_transfer_results(set()) + assert result.failed_recving == set() + return result.finished_recving + + first = poll() receiver.request_queue.join() - second = connector.get_transfer_results(set()) - assert first.finished_recving.isdisjoint(second.finished_recving) - assert first.finished_recving | second.finished_recving == {"load"} - assert first.failed_recving == second.failed_recving == set() + second = poll() + assert first.isdisjoint(second) + assert first | second == {"load"} assert ctypes.string_at(ctypes.addressof(pool) + BLOCK, BLOCK) == bytes([17]) * BLOCK assert client.calls == 1 - assert connector.get_transfer_results(set()).finished_recving == set() + assert poll() == set() + assert connector.build_connector_worker_meta() is None + + +def test_legacy_failure_metadata_unions_multiple_polls_and_drains_once(make_receiver): + client = MemoryClient("miss") + receiver, _ = make_receiver(client) + connector = connector_for(receiver, inline=True) + connector.start_load_kv(None) + assert connector.get_finished(set()) == (set(), {"load"}) + receiver.load_request_sync(request("second", 2)) + assert connector.get_finished(set()) == (set(), {"second"}) + assert connector.get_finished(set()) == (set(), set()) + metadata = connector.build_connector_worker_meta() + assert metadata.failed_recving == {"load", "second"} + assert connector.build_connector_worker_meta() is None + # A subsequent drain owns a new batch and cannot mutate the emitted packet. + receiver.load_request_sync(request("third", 3)) + assert connector.get_finished(set()) == (set(), {"third"}) + assert connector.build_connector_worker_meta().failed_recving == {"third"} + assert metadata.failed_recving == {"load", "second"} + assert client.calls == 3 + + +@pytest.mark.parametrize( + "hook", ["get_finished", "get_transfer_results", "build_connector_worker_meta"], +) +def test_result_hooks_reject_calls_after_shutdown(make_receiver, hook): + client = MemoryClient() + receiver, _ = make_receiver(client) + connector = connector_for(receiver, inline=True) + connector._shutdown = True + args = () if hook == "build_connector_worker_meta" else (set(),) + with pytest.raises(RuntimeError): + getattr(connector, hook)(*args) + assert client.calls == 0 def test_connector_failure_update_is_not_skipped_without_kv_events(): diff --git a/integration/vllm/tests/test_scheduler_full_block_ids.py b/integration/vllm/tests/test_scheduler_full_block_ids.py index 81dbe6a..54c8e48 100644 --- a/integration/vllm/tests/test_scheduler_full_block_ids.py +++ b/integration/vllm/tests/test_scheduler_full_block_ids.py @@ -5,6 +5,7 @@ from dfkv_vllm.data import LoadSpec from dfkv_vllm.scheduler import DfkvStoreScheduler +from dfkv_vllm.transfer_protocol import LegacyReceiveFailures def test_new_request_metadata_uses_complete_allocated_block_table(): @@ -334,8 +335,9 @@ def test_load_admission_parks_request_level_layouts( @pytest.mark.parametrize("cached_resume", [False, True]) @pytest.mark.parametrize("load_async", [False, True]) +@pytest.mark.parametrize("legacy", [False, True]) def test_failed_load_bypasses_persistent_hit_and_rebuilds_computed_save( - cached_resume, load_async, + cached_resume, load_async, legacy, ): scheduler = object.__new__(DfkvStoreScheduler) scheduler.kv_role = "kv_both" @@ -379,7 +381,14 @@ def lookup(req_id, *args, **kwargs): assert pending.requests[0].load_spec.can_load assert pending.requests[0].can_save is False - scheduler.update_connector_output(SimpleNamespace(failed_recving={"retry"})) + if legacy: + failures = LegacyReceiveFailures(failed_recving={"retry"}) + # Model the callback packet after the bridge's all-rank receive fence. + failures._ready = True + output = SimpleNamespace(kv_connector_worker_meta=failures) + else: + output = SimpleNamespace(failed_recving={"retry"}) + scheduler.update_connector_output(output) assert cached == {} assert scheduler.build_connector_meta(step).requests == [] assert scheduler.get_num_new_matched_tokens(request, 0) == (0, False) @@ -427,5 +436,5 @@ def lookup(req_id, *args, **kwargs): assert scheduler.build_connector_meta(step).requests == [] # Late receive completion of an aborted/finalized request cannot resurrect # quarantine or stale allocation metadata. - scheduler.update_connector_output(SimpleNamespace(failed_recving={"retry"})) + scheduler.update_connector_output(output) assert scheduler.get_num_new_matched_tokens(request, 0) == (12, True) diff --git a/integration/vllm/tests/test_transfer_protocol.py b/integration/vllm/tests/test_transfer_protocol.py new file mode 100644 index 0000000..8cb52cc --- /dev/null +++ b/integration/vllm/tests/test_transfer_protocol.py @@ -0,0 +1,593 @@ +"""Exercise legacy metadata fencing through vLLM's real aggregator/scheduler.""" + +from dataclasses import dataclass +from types import SimpleNamespace +import weakref + +import pytest + +pytest.importorskip("vllm") + +from vllm import SamplingParams +from vllm.distributed.kv_transfer.kv_connector.utils import KVOutputAggregator +from vllm.distributed.kv_transfer.kv_connector.v1 import base +from vllm.distributed.kv_transfer.kv_connector.v1.multi_connector import ( + MultiConnector, +) +from vllm.v1.core.sched.request_queue import FCFSRequestQueue +from vllm.v1.core.sched.scheduler import Scheduler +from vllm.v1.outputs import KVConnectorOutput, ModelRunnerOutput +from vllm.v1.request import Request, RequestStatus + +from dfkv_vllm import transfer_protocol as protocol + + +def request(req_id, status=RequestStatus.WAITING_FOR_REMOTE_KVS): + req = Request(req_id, list(range(32)), SamplingParams(max_tokens=4), None) + req.status = status + req.num_computed_tokens = 16 + return req + + +def worker_output(finished=(), failed=(), *, blocks=(), metadata=True): + return ModelRunnerOutput( + req_ids=[], req_id_to_index={}, + kv_connector_output=KVConnectorOutput( + finished_recving=set(finished), invalid_block_ids=set(blocks), + kv_connector_worker_meta=( + protocol.LegacyReceiveFailures(set(failed)) if metadata else None + ), + ), + ) + + +@pytest.fixture +def legacy_scheduler(monkeypatch): + if protocol.HAS_NATIVE_TRANSFER_RESULTS: + pytest.skip("Legacy integration runs against the old engine installations") + # Restore process-wide hooks after this test, even when another connector + # has already installed them during test collection. + monkeypatch.setattr(Scheduler, "update_from_output", Scheduler.update_from_output) + monkeypatch.setattr(Scheduler, "_handle_invalid_blocks", Scheduler._handle_invalid_blocks) + protocol.install_legacy_failure_bridge() + protocol.install_legacy_failure_bridge() + + def make(*requests, recompute=True): + scheduler = object.__new__(Scheduler) + scheduler.requests = {req.request_id: req for req in requests} + scheduler.running = [] + scheduler.waiting = FCFSRequestQueue() + scheduler.skipped_waiting = FCFSRequestQueue(requests) + scheduler.recompute_kv_load_failures = recompute + scheduler.failed_recving_kv_req_ids = set() + scheduler.finished_recving_kv_req_ids = set() + scheduler.finished_req_ids = set() + scheduler.finished_req_ids_dict = None + scheduler.grammar_compile_error_reqs = set() + scheduler.defer_block_free = False + scheduler.perf_metrics = None + scheduler.ec_connector = None + scheduler._inflight_prefills = set() + scheduler.needs_kv_cache_zeroing = False + scheduler.block_size = 4 + scheduler.make_stats = lambda *args: None + scheduler._connector_finished = lambda req: (False, None) + scheduler.encoder_cache_manager = SimpleNamespace(free=lambda req: None) + scheduler.freed = [] + scheduler.cached = [] + scheduler.evicted = [] + scheduler.tables = {} + scheduler.kv_cache_manager = SimpleNamespace( + take_events=lambda: [], + free=lambda req: scheduler.freed.append(req.request_id), + cache_blocks=lambda req, count: scheduler.cached.append((req.request_id, count)), + get_block_ids=lambda req_id: scheduler.tables[req_id], + evict_blocks=lambda blocks: scheduler.evicted.append(set(blocks)), + ) + scheduler.callbacks = [] + scheduler.connector = SimpleNamespace( + update_connector_output=lambda output: scheduler.callbacks.append( + set(protocol.failed_requests_from_output(output)) + ), + get_kv_connector_stats=lambda: None, + take_events=lambda: [], + ) + return scheduler + + return make + + +def step(scheduler, output): + return scheduler.update_from_output( + SimpleNamespace(num_scheduled_tokens={}, total_num_scheduled_tokens=0), output, + ) + + +def multi_connector(*children): + connector = object.__new__(MultiConnector) + connector._connectors = list(children) + return connector + + +def multi_worker_output(finished=(), *, metadata): + def worker(child): + if isinstance(child, tuple): + return multi_connector(*(worker(item) for item in child)) + return SimpleNamespace(build_connector_worker_meta=lambda: child) + + output = worker_output(finished, metadata=False) + output.kv_connector_output.kv_connector_worker_meta = ( + worker(metadata).build_connector_worker_meta() + ) + return output + + +def dfkv_child(*requests): + from dfkv_vllm.scheduler import DfkvStoreScheduler + + state = object.__new__(DfkvStoreScheduler) + state.request_level_loads = True + state._unfinished_request_ids = {req.request_id for req in requests} + state._failed_load_req_ids = set() + state._preempted_req_ids = set() + state._allocated_req_ids = set(state._unfinished_request_ids) + state._unfinished_requests = {} + state._request_trackers = {} + state.load_specs = {} + discarded = [] + state.client = SimpleNamespace(discard=discarded.append) + return SimpleNamespace( + state=state, discarded=discarded, + update_connector_output=state.update_connector_output, + get_kv_connector_stats=lambda: None, + take_events=lambda: [], + ) + + +@dataclass +class ForeignMetadata(base.KVConnectorWorkerMetadata): + events: tuple[str, ...] + + def aggregate(self, other): + return ForeignMetadata(self.events + other.events) + + +def foreign_child(): + received = [] + return SimpleNamespace( + received=received, + update_connector_output=lambda output: received.append( + output.kv_connector_worker_meta + ), + get_kv_connector_stats=lambda: None, + take_events=lambda: [], + ) + + +@pytest.mark.parametrize("nested", [False, True]) +@pytest.mark.parametrize("recompute", [False, True]) +@pytest.mark.parametrize("late_foreign", [False, True]) +def test_multi_connector_routes_fenced_failures_without_replaying_foreign_metadata( + legacy_scheduler, nested, recompute, late_foreign, +): + first, second, good = request("first"), request("second"), request("good") + scheduler = legacy_scheduler(first, second, good, recompute=recompute) + # Both dfkv children know every request, so misrouting would quarantine a + # healthy source too. Only the child whose worker failed may be quarantined. + left, right = dfkv_child(first, second, good), dfkv_child(first, second, good) + foreign = foreign_child() + scheduler.connector = ( + multi_connector(left, multi_connector(right, foreign)) + if nested else multi_connector(left, right, foreign) + ) + + def layout(left_meta, right_meta, foreign_meta): + return ( + (left_meta, (right_meta, foreign_meta)) + if nested else (left_meta, right_meta, foreign_meta) + ) + + aggregator = KVOutputAggregator(2) + early = aggregator.aggregate([ + multi_worker_output({"first"}, metadata=layout( + protocol.LegacyReceiveFailures({"first"}), None, ForeignMetadata(("rank0",)), + )), + multi_worker_output({"second"}, metadata=layout( + None, protocol.LegacyReceiveFailures({"second"}), ForeignMetadata(("rank1",)), + )), + ]) + original_metadata = early.kv_connector_output.kv_connector_worker_meta + assert step(scheduler, early) == {} + assert early.kv_connector_output.kv_connector_worker_meta is original_metadata + assert (first.num_computed_tokens, second.num_computed_tokens) == (16, 16) + assert left.discarded == right.discarded == [] + assert foreign.received[0].events == ("rank0", "rank1") + original_foreign = ( + original_metadata.metadata[1].metadata[1] + if nested else original_metadata.metadata[2] + ) + assert foreign.received[0] is original_foreign + + fresh_foreign = ForeignMetadata(("completion",)) if late_foreign else None + late = aggregator.aggregate([ + multi_worker_output({"second"}, metadata=layout(None, None, fresh_foreign)), + multi_worker_output({"first"}, metadata=layout(None, None, None)), + ]) + original_metadata = late.kv_connector_output.kv_connector_worker_meta + outputs = step(scheduler, late) + assert late.kv_connector_output.kv_connector_worker_meta is original_metadata + assert left.discarded == ["first"] + assert right.discarded == ["second"] + assert left.state._failed_load_req_ids == {"first"} + assert right.state._failed_load_req_ids == {"second"} + assert foreign.received == [original_foreign, fresh_foreign] + assert foreign.received[-1] is fresh_foreign + assert good.num_computed_tokens == 16 + if recompute: + assert outputs == {} + assert (first.num_computed_tokens, second.num_computed_tokens) == (0, 0) + assert scheduler.failed_recving_kv_req_ids == {"first", "second"} + scheduler._update_waiting_for_remote_kv(first) + scheduler._update_waiting_for_remote_kv(second) + assert scheduler.freed == ["first", "second"] + assert scheduler.cached == [] + else: + assert first.status == second.status == RequestStatus.FINISHED_ERROR + assert {item.request_id for item in outputs[0].outputs} == {"first", "second"} + assert set(scheduler.freed) == {"first", "second"} + + +def test_nested_pending_failures_complete_independently(legacy_scheduler): + first, second = request("first"), request("second") + scheduler = legacy_scheduler(first, second) + child, foreign = dfkv_child(first, second), foreign_child() + scheduler.connector = multi_connector(foreign, multi_connector(child)) + aggregator = KVOutputAggregator(2) + step(scheduler, aggregator.aggregate([ + multi_worker_output({"first", "second"}, metadata=( + ForeignMetadata(("early",)), + (protocol.LegacyReceiveFailures({"first", "second"}),), + )), + multi_worker_output(metadata=(None, (None,))), + ])) + assert child.discarded == [] + + step(scheduler, aggregator.aggregate([ + multi_worker_output(metadata=(None, (None,))), + multi_worker_output({"first"}, metadata=( + ForeignMetadata(("partial",)), (None,), + )), + ])) + assert child.discarded == ["first"] + assert first.num_computed_tokens == 0 + assert second.num_computed_tokens == 16 + step(scheduler, aggregator.aggregate([ + multi_worker_output(metadata=(None, (None,))), + multi_worker_output({"second"}, metadata=(None, (None,))), + ])) + assert child.discarded == ["first", "second"] + assert second.num_computed_tokens == 0 + assert [ + metadata.events if metadata is not None else None + for metadata in foreign.received + ] == [("early",), ("partial",), None] + + +@pytest.mark.parametrize("retirement", ["cancel", "replace", "remove", "packetless"]) +def test_nested_pending_failures_do_not_poison_retired_requests( + legacy_scheduler, retirement, +): + req = request("retired") + scheduler = legacy_scheduler(req) + child, foreign = dfkv_child(req), foreign_child() + scheduler.connector = multi_connector(foreign, multi_connector(child)) + aggregator = KVOutputAggregator(2) + step(scheduler, aggregator.aggregate([ + multi_worker_output({"retired"}, metadata=( + None, (protocol.LegacyReceiveFailures({"retired"}),), + )), + multi_worker_output(metadata=(None, (None,))), + ])) + if retirement == "cancel": + scheduler.finish_requests({"retired"}, RequestStatus.FINISHED_ABORTED) + assert scheduler.freed == [] + else: + scheduler.requests.clear() + scheduler.skipped_waiting.clear() + if retirement == "packetless": + step(scheduler, ModelRunnerOutput(req_ids=[], req_id_to_index={})) + if retirement == "remove": + reference = weakref.ref(req) + del req + assert reference() is None + replacement = request("retired") + scheduler.requests["retired"] = replacement + scheduler.skipped_waiting.add_request(replacement) + repeat = ( + protocol.LegacyReceiveFailures({"retired"}) + if retirement in {"cancel", "replace"} else None + ) + step(scheduler, aggregator.aggregate([ + multi_worker_output(metadata=(None, (None,))), + multi_worker_output({"retired"}, metadata=(None, (repeat,))), + ])) + assert child.discarded == [] + assert not child.state._failed_load_req_ids + assert not scheduler.failed_recving_kv_req_ids + if retirement == "cancel": + assert scheduler.freed == ["retired"] + assert "retired" not in scheduler.requests + else: + assert replacement.num_computed_tokens == 16 + assert scheduler.freed == [] + if retirement == "replace": + assert not scheduler.finished_recving_kv_req_ids + + +def test_foreign_only_multi_metadata_is_not_intercepted(legacy_scheduler): + scheduler = legacy_scheduler(request("foreign")) + left, right = foreign_child(), foreign_child() + scheduler.connector = multi_connector(left, multi_connector(right)) + output = multi_worker_output(metadata=( + ForeignMetadata(("left",)), (ForeignMetadata(("right",)),), + )) + original = output.kv_connector_output.kv_connector_worker_meta + step(scheduler, output) + assert output.kv_connector_output.kv_connector_worker_meta is original + assert left.received == [original.metadata[0]] + assert left.received[0] is original.metadata[0] + assert right.received == [original.metadata[1].metadata[0]] + assert right.received[0] is original.metadata[1].metadata[0] + + +@pytest.mark.parametrize("recompute", [False, True]) +def test_staggered_workers_wait_for_all_ranks_before_recovery(legacy_scheduler, recompute): + failed, unaffected = request("failed"), request("unaffected") + scheduler = legacy_scheduler(failed, unaffected, recompute=recompute) + aggregator = KVOutputAggregator(expected_finished_count=2) + early = aggregator.aggregate([ + worker_output({"failed"}, {"failed"}), worker_output(), + ]) + assert step(scheduler, early) == {} + assert failed.num_computed_tokens == 16 + assert scheduler.freed == [] + assert scheduler.callbacks == [set()] + assert not scheduler.failed_recving_kv_req_ids + # Rank zero has already drained its failure metadata. Only the final + # worker completes now, so the bridge must retain the earlier failure. + late = aggregator.aggregate([ + worker_output(metadata=False), worker_output({"failed"}, metadata=False), + ]) + outputs = step(scheduler, late) + assert scheduler.callbacks[-1] == {"failed"} + assert unaffected.num_computed_tokens == 16 + assert unaffected.status == RequestStatus.WAITING_FOR_REMOTE_KVS + if recompute: + assert outputs == {} + assert failed.num_computed_tokens == 0 + assert scheduler.failed_recving_kv_req_ids == {"failed"} + assert scheduler.finished_recving_kv_req_ids == {"failed"} + scheduler._update_waiting_for_remote_kv(failed) + assert scheduler.freed == ["failed"] + assert scheduler.cached == [] + assert not scheduler.failed_recving_kv_req_ids + else: + assert failed.status == RequestStatus.FINISHED_ERROR + assert [item.request_id for item in outputs[0].outputs] == ["failed"] + assert outputs[0].outputs[0].finish_reason == failed.get_finished_reason() + assert "failed" not in scheduler.requests + assert scheduler.freed == ["failed"] + + +def test_same_step_worker_metadata_unions_failures(legacy_scheduler): + first, second, good = request("first"), request("second"), request("good") + scheduler = legacy_scheduler(first, second, good) + aggregator = KVOutputAggregator(2) + completed = {"first", "second", "good"} + result = aggregator.aggregate([ + worker_output(completed, {"first"}), + worker_output(completed, {"second"}), + ]) + step(scheduler, result) + assert (first.num_computed_tokens, second.num_computed_tokens) == (0, 0) + assert good.num_computed_tokens == 16 + assert scheduler.callbacks == [{"first", "second"}] + + +def test_request_errors_do_not_scan_or_alias_hybrid_block_ids(legacy_scheduler): + failed, other = request("failed"), request("other") + running = request("running", RequestStatus.RUNNING) + scheduler = legacy_scheduler(failed, other, running) + # No tables are supplied: any block lookup would fail. Request-level + # identities remain unambiguous even if physical pools reuse block IDs. + step(scheduler, worker_output({"failed"}, {"failed", "running", "missing"})) + assert failed.num_computed_tokens == 0 + assert (other.num_computed_tokens, running.num_computed_tokens) == (16, 16) + assert scheduler.failed_recving_kv_req_ids == {"failed"} + assert scheduler.callbacks == [{"failed"}] + + +@pytest.mark.parametrize("recompute", [False, True]) +@pytest.mark.parametrize("request_error", [False, True]) +def test_real_block_errors_keep_engine_prefix_recovery( + legacy_scheduler, recompute, request_error, +): + full, untouched, hybrid = request("full"), request("untouched"), request("hybrid") + scheduler = legacy_scheduler(full, untouched, hybrid, recompute=recompute) + scheduler.tables = { + "full": ([10, 11, 12, 13],), + "untouched": ([20, 21, 22, 23],), + "hybrid": ([30, 31, 32, 33],), + } + result = step(scheduler, worker_output( + {"full", "hybrid"}, {"hybrid"} if request_error else (), blocks={12}, + )) + if recompute: + assert full.num_computed_tokens == 8 + assert hybrid.num_computed_tokens == (0 if request_error else 16) + assert scheduler.failed_recving_kv_req_ids == ( + {"full", "hybrid"} if request_error else {"full"} + ) + else: + assert full.status == RequestStatus.FINISHED_ERROR + assert {item.request_id for item in result[0].outputs} == ( + {"full", "hybrid"} if request_error else {"full"} + ) + assert untouched.num_computed_tokens == 16 + assert untouched.status == RequestStatus.WAITING_FOR_REMOTE_KVS + + +def test_abort_before_last_worker_releases_only_after_completion(legacy_scheduler): + req = request("cancelled") + scheduler = legacy_scheduler(req) + step(scheduler, worker_output(failed={"cancelled"})) + scheduler.finish_requests({"cancelled"}, RequestStatus.FINISHED_ABORTED) + assert scheduler.freed == [] + # The engine intentionally keeps a cancelled receive in requests until + # its all-worker completion. The bridge must not quarantine it again. + step(scheduler, worker_output({"cancelled"}, {"cancelled", "unknown"})) + assert scheduler.freed == ["cancelled"] + assert scheduler.callbacks[-1] == set() + assert "cancelled" not in scheduler.requests + assert not scheduler.failed_recving_kv_req_ids + # Duplicate/late dfkv metadata and completions cannot resurrect state. + step(scheduler, worker_output({"cancelled"}, {"cancelled"})) + replacement = request("cancelled") + scheduler.requests["cancelled"] = replacement + scheduler.skipped_waiting.add_request(replacement) + step(scheduler, worker_output({"cancelled"})) + assert replacement.num_computed_tokens == 16 + assert scheduler.callbacks[-1] == set() + assert scheduler.freed == ["cancelled"] + + +def test_pending_failures_do_not_hold_removed_requests_alive(legacy_scheduler): + req = request("removed") + scheduler = legacy_scheduler(req) + step(scheduler, worker_output(failed={"removed"})) + reference = weakref.ref(req) + scheduler.requests.clear() + scheduler.skipped_waiting.clear() + del req + assert reference() is None + replacement = request("removed") + scheduler.requests["removed"] = replacement + step(scheduler, worker_output({"removed"})) + assert replacement.num_computed_tokens == 16 + assert scheduler.callbacks[-1] == set() + + +def test_replaced_request_does_not_consume_previous_failure_or_completion(legacy_scheduler): + old = request("reused") + scheduler = legacy_scheduler(old) + step(scheduler, worker_output(failed={"reused"})) + replacement = request("reused") + scheduler.requests["reused"] = replacement + scheduler.skipped_waiting.clear() + scheduler.skipped_waiting.add_request(replacement) + step(scheduler, worker_output({"reused"}, {"reused"})) + assert replacement.num_computed_tokens == 16 + assert not scheduler.finished_recving_kv_req_ids + assert not scheduler.failed_recving_kv_req_ids + assert scheduler.callbacks[-1] == set() + + +def test_packetless_step_prunes_retired_failures(legacy_scheduler): + req = request("retired") + scheduler = legacy_scheduler(req) + step(scheduler, worker_output(failed={"retired"})) + scheduler.requests.clear() + scheduler.skipped_waiting.clear() + step(scheduler, ModelRunnerOutput(req_ids=[], req_id_to_index={})) + replacement = request("retired") + scheduler.requests["retired"] = replacement + step(scheduler, worker_output({"retired"})) + assert replacement.num_computed_tokens == 16 + assert scheduler.callbacks[-1] == set() + + +def test_switching_connector_drops_dfkv_pending_scope(legacy_scheduler): + req = request("old") + scheduler = legacy_scheduler(req) + step(scheduler, worker_output(failed={"old"})) + marker = object() + observed = [] + scheduler.connector = SimpleNamespace( + update_connector_output=lambda packet: observed.append( + packet.kv_connector_worker_meta + ), + get_kv_connector_stats=lambda: None, + take_events=lambda: [], + ) + output = worker_output({"old"}, metadata=False) + output.kv_connector_output.kv_connector_worker_meta = marker + step(scheduler, output) + assert observed == [marker] + assert req.num_computed_tokens == 16 + assert not scheduler.failed_recving_kv_req_ids + + +@pytest.mark.parametrize("failure_point", ["callback", "block_handler"]) +def test_temporary_packet_is_restored_when_engine_raises( + legacy_scheduler, failure_point, +): + req = request("load") + scheduler = legacy_scheduler(req) + output = worker_output({"load"}, {"load"}, blocks={12} if failure_point == "block_handler" else ()) + packet = output.kv_connector_output + original_blocks = packet.invalid_block_ids + original_finished = packet.finished_recving + original_metadata = packet.kv_connector_worker_meta + + def fail(*args): + raise RuntimeError("injected engine failure") + + if failure_point == "callback": + scheduler.connector.update_connector_output = fail + else: + scheduler.kv_cache_manager.get_block_ids = fail + with pytest.raises(RuntimeError, match="injected engine failure"): + step(scheduler, output) + assert packet.invalid_block_ids is original_blocks + assert packet.finished_recving is original_finished + assert packet.kv_connector_worker_meta is original_metadata + assert protocol.failed_requests_from_output(packet) == set() + + +def test_unrelated_connector_is_not_intercepted(legacy_scheduler): + req = request("other") + scheduler = legacy_scheduler(req) + marker = object() + output = worker_output(metadata=False) + output.kv_connector_output.kv_connector_worker_meta = marker + observed = [] + scheduler.connector.update_connector_output = lambda packet: observed.append( + packet.kv_connector_worker_meta + ) + step(scheduler, output) + assert observed == [marker] + assert req.num_computed_tokens == 16 + + +def test_native_engine_uses_native_result_without_scheduler_hooks(): + if not protocol.HAS_NATIVE_TRANSFER_RESULTS: + pytest.skip("Native route runs against the native engine installation") + original_update = Scheduler.update_from_output + original_handler = Scheduler._handle_invalid_blocks + protocol.install_legacy_failure_bridge() + protocol.install_legacy_failure_bridge() + assert Scheduler.update_from_output is original_update + assert Scheduler._handle_invalid_blocks is original_handler + assert protocol.TransferResults is base.KVConnectorTransferResults + output = KVConnectorOutput(finished_recving={"load"}, failed_recving={"load"}) + assert protocol.failed_requests_from_output(output) == {"load"} + + +def test_raw_worker_metadata_is_not_a_scheduler_failure(): + metadata = protocol.LegacyReceiveFailures({"load"}) + output = SimpleNamespace(kv_connector_worker_meta=metadata) + assert protocol.failed_requests_from_output(output) == set() + metadata._ready = True + assert protocol.failed_requests_from_output(output) == {"load"} + metadata.aggregate(protocol.LegacyReceiveFailures({"second"})) + assert protocol.failed_requests_from_output(output) == set() diff --git a/test/python/test_dfkv_vllm_connector.py b/test/python/test_dfkv_vllm_connector.py index af05535..3c3cf21 100644 --- a/test/python/test_dfkv_vllm_connector.py +++ b/test/python/test_dfkv_vllm_connector.py @@ -25,7 +25,9 @@ class Event: torch_module = ModuleType("torch") torch_cuda_module = ModuleType("torch.cuda") torch_cuda_module.Event = Event + torch_cuda_module.is_available = lambda: False torch_module.Tensor = Tensor + torch_module.float16 = "float16" torch_module.cuda = torch_cuda_module sys.modules["torch"] = torch_module sys.modules["torch.cuda"] = torch_cuda_module @@ -71,8 +73,34 @@ class Placeholder: class BlockHash(bytes): pass + + @dataclass(frozen=True) + class FullAttentionSpec: + block_size: int + num_kv_heads: int + head_size: int + dtype: object class KVConnectorBase_V1: - pass + def get_finished(self, finished_req_ids): + raise NotImplementedError + + def build_connector_worker_meta(self): + return None + + def update_connector_output(self, connector_output): + pass + + class KVConnectorWorkerMetadata: + def aggregate(self, other): + raise NotImplementedError + + @dataclass + class KVConnectorOutput: + finished_sending: set[str] = field(default_factory=set) + finished_recving: set[str] = field(default_factory=set) + invalid_block_ids: set[int] = field(default_factory=set) + kv_connector_worker_meta: object = None + kv_cache_events: object = None class SupportsHMA: pass @@ -138,6 +166,7 @@ def get_manager_class(cls, _spec): "vllm.distributed.kv_transfer.kv_connector.v1.base", KVConnectorBase_V1=KVConnectorBase_V1, KVConnectorMetadata=Placeholder, + KVConnectorWorkerMetadata=KVConnectorWorkerMetadata, KVConnectorRole=KVConnectorRole, SupportsHMA=SupportsHMA, ) @@ -193,14 +222,14 @@ def get_manager_class(cls, _spec): stub( "vllm.v1.kv_cache_interface", AttentionSpec=Placeholder, - FullAttentionSpec=Placeholder, + FullAttentionSpec=FullAttentionSpec, KVCacheConfig=Placeholder, KVCacheGroupSpec=Placeholder, KVCacheSpec=Placeholder, MambaSpec=Placeholder, UniformTypeKVCacheSpecs=Placeholder, ) - stub("vllm.v1.outputs", KVConnectorOutput=Placeholder) + stub("vllm.v1.outputs", KVConnectorOutput=KVConnectorOutput) stub( "vllm.v1.kv_cache_spec_registry", KVCacheSpecRegistry=KVCacheSpecRegistry, @@ -215,6 +244,9 @@ def get_manager_class(cls, _spec): sys.path.insert(0, str(ROOT / "integration" / "common" / "src")) sys.path.insert(0, str(ROOT / "integration" / "vllm" / "src")) +import torch # noqa: E402 +from vllm.v1.kv_cache_interface import FullAttentionSpec # noqa: E402 + from dfkv_common import VLLM_RAW_V1, pool_key # noqa: E402 from dfkv_vllm.connector import DfkvStoreConnector # noqa: E402 from dfkv_vllm.data import ( # noqa: E402 @@ -227,6 +259,7 @@ def get_manager_class(cls, _spec): DfkvStorePromMetrics, ) from dfkv_vllm.scheduler import DfkvStoreScheduler # noqa: E402 +from dfkv_vllm.transfer_protocol import TransferResults # noqa: E402 class _LookupClient: @@ -271,12 +304,25 @@ def _scheduler( kv_transfer_config=transfer, cache_config=SimpleNamespace(cache_salt="stable"), ) + kv_cache_config = SimpleNamespace( + kv_cache_groups=[ + SimpleNamespace( + layer_names=["model.layers.0.self_attn"], + kv_cache_spec=FullAttentionSpec( + block_size=64, + num_kv_heads=8, + head_size=128, + dtype=torch.float16, + ), + ) + ], + ) with ( patch("dfkv_vllm.scheduler.ensure_deterministic_block_hashing"), patch("dfkv_vllm.scheduler.resolve_kv_cache_block_sizes", return_value=(64, 64)), patch("dfkv_vllm.scheduler.LookupKeyClient", return_value=client), ): - return DfkvStoreScheduler(config, SimpleNamespace()) + return DfkvStoreScheduler(config, kv_cache_config) @staticmethod def _request(request_id: str, blocks: int = 2): @@ -366,9 +412,13 @@ def _block_call(self) -> None: self.call_entered.wait(timeout=5) self.call_release.wait(timeout=5) - def get_finished(self, finished_req_ids, metadata): + def get_transfer_results(self, finished_req_ids, metadata): self._block_call() - return set(finished_req_ids), set(metadata.unfinished_request_ids) + return TransferResults( + finished_sending=set(finished_req_ids), + finished_recving=set(metadata.unfinished_request_ids), + failed_recving=set(), + ) def request_finished(self, request, block_ids): self._block_call() @@ -390,6 +440,7 @@ def _connector( connector._inflight_calls = 0 connector._shutdown = False connector._shutdown_complete = False + connector._legacy_failed_recving = set() connector.connector_worker = worker connector.connector_scheduler = scheduler return connector @@ -431,7 +482,7 @@ def run_shutdown() -> None: self.assertTrue(shutdown_thread.is_alive()) self.assertEqual(backend.close_calls, 0) - with self.assertRaisesRegex(RuntimeError, "connector is shut down"): + with self.assertRaises(RuntimeError): rejected_call() backend.call_release.wait(timeout=5) @@ -444,26 +495,36 @@ def run_shutdown() -> None: connector.shutdown() self.assertEqual(backend.close_calls, 1) - with self.assertRaisesRegex(RuntimeError, "connector is shut down"): + with self.assertRaises(RuntimeError): rejected_call() return call_results - def test_get_finished_admitted_before_shutdown_completes_before_close( + def test_result_calls_admitted_before_shutdown_complete_before_close( self, ) -> None: - backend = _BlockingConnectorBackend() - connector = self._connector(worker=backend) - metadata = DfkvStoreConnectorMetadata({"unfinished"}, set()) - connector._get_connector_metadata = lambda: metadata - - results = self._assert_shutdown_waits_for_call( - connector, - backend, - lambda: connector.get_finished({"finished"}), - lambda: connector.get_finished({"too-late"}), - ) - - self.assertEqual(results, [({"finished"}, {"unfinished"})]) + for legacy in (False, True): + with self.subTest(legacy=legacy): + backend = _BlockingConnectorBackend() + connector = self._connector(worker=backend) + metadata = DfkvStoreConnectorMetadata({"unfinished"}, set()) + connector._get_connector_metadata = lambda: metadata + hook = ( + connector.get_finished if legacy else connector.get_transfer_results + ) + results = self._assert_shutdown_waits_for_call( + connector, + backend, + lambda: hook({"finished"}), + lambda: hook({"too-late"}), + ) + expected = ( + ({"finished"}, {"unfinished"}) if legacy else TransferResults( + finished_sending={"finished"}, + finished_recving={"unfinished"}, + failed_recving=set(), + ) + ) + self.assertEqual(results, [expected]) def test_request_lifecycle_calls_admitted_before_shutdown_finish_first( self, diff --git a/test/python/test_dfkv_vllm_worker.py b/test/python/test_dfkv_vllm_worker.py index 2c69d7e..f5a6bba 100644 --- a/test/python/test_dfkv_vllm_worker.py +++ b/test/python/test_dfkv_vllm_worker.py @@ -27,6 +27,7 @@ class Event: torch_module = ModuleType("torch") torch_cuda_module = ModuleType("torch.cuda") torch_cuda_module.Event = Event + torch_cuda_module.is_available = lambda: False torch_module.Tensor = Tensor torch_module.cuda = torch_cuda_module sys.modules["torch"] = torch_module @@ -73,6 +74,27 @@ class Placeholder: class BlockHash(bytes): pass + class KVConnectorBase_V1: + def get_finished(self, finished_req_ids): + raise NotImplementedError + + def build_connector_worker_meta(self): + return None + + def update_connector_output(self, connector_output): + pass + + class KVConnectorWorkerMetadata: + def aggregate(self, other): + raise NotImplementedError + + @dataclass + class KVConnectorOutput: + finished_sending: set[str] = field(default_factory=set) + finished_recving: set[str] = field(default_factory=set) + invalid_block_ids: set[int] = field(default_factory=set) + kv_connector_worker_meta: object = None + @dataclass class KVCacheBlock: block_id: int @@ -117,6 +139,8 @@ def get_manager_class(cls, _spec): stub( "vllm.distributed.kv_transfer.kv_connector.v1.base", KVConnectorMetadata=Placeholder, + KVConnectorBase_V1=KVConnectorBase_V1, + KVConnectorWorkerMetadata=KVConnectorWorkerMetadata, ) stub( "vllm.distributed.kv_transfer.kv_connector.v1.metrics", @@ -163,6 +187,8 @@ def get_manager_class(cls, _spec): ) stub( "vllm.v1.kv_cache_interface", + AttentionSpec=Placeholder, + MambaSpec=Placeholder, FullAttentionSpec=Placeholder, KVCacheConfig=Placeholder, KVCacheGroupSpec=Placeholder, @@ -173,6 +199,7 @@ def get_manager_class(cls, _spec): "vllm.v1.kv_cache_spec_registry", KVCacheSpecRegistry=KVCacheSpecRegistry, ) + stub("vllm.v1.outputs", KVConnectorOutput=KVConnectorOutput) stub("vllm.v1.request", Request=Placeholder) @@ -193,6 +220,7 @@ def get_manager_class(cls, _spec): ReqMeta, ) from dfkv_vllm.protocol import LOOKUP_MSG # noqa: E402 +from dfkv_vllm.transfer_protocol import TransferResults # noqa: E402 from dfkv_vllm.worker import ( # noqa: E402 DfkvStoreWorker, KVCacheStoreRecvingThread, @@ -556,15 +584,15 @@ def batch_remove(self, keys): kv_role="kv_both", ready_event=threading.Event(), ) - sender.add_stored_request("save") - sender._handle_request( - ReqMeta( - req_id="save", - token_len_chunk=64, - block_ids=([7],), - block_hashes=[self._hash("save")], - ) + request = ReqMeta( + req_id="save", + token_len_chunk=64, + block_ids=([7],), + block_hashes=[self._hash("save")], + can_save=True, ) + sender.add_stored_request(request) + sender._handle_request(request) expected_key = logical_key.to_bytes() self.assertEqual(client.exist_calls, [[expected_key]]) @@ -640,8 +668,9 @@ def close(): token_len_chunk=64, block_ids=([7],), block_hashes=[self._hash("blocked-save")], + can_save=True, ) - sender.add_stored_request(request.req_id) + sender.add_stored_request(request) sender.start() self.assertTrue(sender.ready_event.wait(timeout=1)) self.assertTrue(sender.add_request(request)) @@ -661,6 +690,7 @@ def short_diagnostic_wait(req_id: str) -> bool: sender.wait_for_inflight_put = short_diagnostic_wait worker = DfkvStoreWorker.__new__(DfkvStoreWorker) worker.load_async = True + worker.request_level_loads = False worker.kv_recv_thread = None worker.kv_send_thread = sender worker.lookup_server = None @@ -1158,29 +1188,13 @@ class _Coordinator: lcm_block_size = 128 @staticmethod - def store_mask(_token_len): - return [[True] * 3, [True] * 6] + def load_mask(_block_hashes, token_len): + return [[True] * (token_len // 128), [True] * (token_len // 64)] @staticmethod def block_hashes_for_spec(_block_hashes, spec): return wide_hashes if spec is wide_spec else fine_hashes - @staticmethod - def find_longest_cache_hit(_block_hashes, _token_len, pool): - hit = 0 - for chunk in range(3): - required = ( - (0, wide_hashes[chunk]), - (1, fine_hashes[chunk * 2]), - (1, fine_hashes[chunk * 2 + 1]), - ) - if not all( - pool.get_cached_block(block_hash, [group_id]) - for group_id, block_hash in required - ): - break - hit += 128 - return [[], []], hit class _Client: def __init__(self): @@ -1190,6 +1204,11 @@ def batch_exist(self, keys): self.calls.append(list(keys)) if isinstance(response, Exception): raise response + if response is not None and len(response) == 9: + # Preserve occurrence-specific statuses when admission + # rechecks a shorter prefix, even for repeated hashes. + wide_count = len(keys) // 3 + return response[:wide_count] + response[3:3 + 2 * wide_count] return response client = _Client() @@ -1216,13 +1235,11 @@ def test_stops_at_first_incomplete_cross_group_chunk_and_ignores_later_hits( # Candidate order is three wide-group chunks followed by six fine-group # subchunks. Logical chunk 1 is incomplete only in the fine group; # logical chunk 2 is an isolated later hit and must not be loaded. - worker, client = self._worker( + worker, _client = self._worker( [1, 1, 1, 1, 1, 1, 0, 1, 1] ) block_hashes = [self._hash(f"scheduler-{index}") for index in range(6)] self.assertEqual(worker.lookup(384, block_hashes), 128) - self.assertEqual(len(client.calls), 1) - self.assertEqual(len(client.calls[0]), 9) first_missing_worker, _ = self._worker( [1, 1, 1, 1, 0, 1, 1, 1, 1] @@ -1331,12 +1348,6 @@ def test_recv_worker_configuration_is_positive_and_bounded(self) -> None: *common, recv_workers=invalid # type: ignore[arg-type] ) - receiver = KVCacheStoreRecvingThread( - *common, recv_workers="4" # type: ignore[arg-type] - ) - self.assertEqual(receiver.recv_workers, 4) - self.assertEqual(len(receiver.worker_threads), 4) - receiver.stop() def _requests(self, request_ids: tuple[str, ...]) -> tuple[ dict[str, ReqMeta], dict[bytes, str] @@ -1365,7 +1376,7 @@ def _receiver( self, request_ids: tuple[str, ...], *, - recv_workers: int = 1, + recv_workers: int | str = 1, queue_capacity: int = 8, ) -> tuple[ KVCacheStoreRecvingThread, @@ -1460,7 +1471,7 @@ def test_recv_workers_overlap_complete_out_of_order_and_bound_active_calls( done, counts, order, - ) = self._receiver(request_ids, recv_workers=2, queue_capacity=1) + ) = self._receiver(request_ids, recv_workers="2", queue_capacity=1) try: self.assertTrue(receiver.add_request(requests["slow"])) self._wait(client.started["slow"], "slow native GET") @@ -1574,6 +1585,7 @@ def observed_cancel(req_ids, *, wait=True, **kwargs): ) worker = DfkvStoreWorker.__new__(DfkvStoreWorker) worker.load_async = True + worker.request_level_loads = False worker.kv_send_thread = send_thread worker.kv_recv_thread = receiver returned = threading.Event() @@ -1625,13 +1637,14 @@ def observed_cancel(req_ids, *, wait=True, **kwargs): worker.kv_send_thread = None worker.kv_role = "kv_consumer" worker.load_async = True + worker.request_level_loads = False worker.tp_rank = 0 returned = threading.Event() - result: list[tuple[set[str], set[str]]] = [] + result: list[TransferResults] = [] def finish_abort() -> None: result.append( - worker.get_finished( + worker.get_transfer_results( {"aborted"}, SimpleNamespace(requests=[], preempted_req_ids=set()), ) @@ -1647,7 +1660,11 @@ def finish_abort() -> None: client.release["aborted"].set() self._wait(returned, "finished-request receive fence") thread.join(timeout=1) - self.assertEqual(result, [(set(), {"aborted"})]) + self.assertEqual(result, [TransferResults( + finished_sending=set(), + finished_recving={"aborted"}, + failed_recving=set(), + )]) self.assertEqual(counts["aborted"], 1) self.assertEqual( receiver.get_and_clear_block_ids_with_load_errors(), set() @@ -1765,13 +1782,14 @@ class _Client: _sg_segs_cache = 2 def __init__(self) -> None: - self.keys: list[bytes] = [] + self.calls: list[list[bytes]] = [] def batch_exist(self, keys): - self.keys = list(keys) + self.calls.append(list(keys)) # Both TP objects of chunk 0 are present. Chunk 1 is only # partially present and therefore cannot extend the prefix. - return [1, 1, 1, 0] + missing = PoolKey(replace(metadata, tp_rank=1), hashes[1].hex()) + return [int(key != missing.to_bytes()) for key in keys] client = _Client() records: list[dict[str, object]] = [] @@ -1781,19 +1799,11 @@ def batch_exist(self, keys): _seg_layout=[1, 2, 3], ) - def _find_longest(values, _token_len, pool): - hit = 0 - for value in values: - if pool.get_cached_block(value, [0]) is None: - break - hit += 64 - return None, hit coord = SimpleNamespace( lcm_block_size=64, - store_mask=lambda _token_len: [[True, True]], + load_mask=lambda _hashes, token_len: [[True] * (token_len // 64)], block_hashes_for_spec=lambda values, _spec: values, - find_longest_cache_hit=_find_longest, ) worker = SimpleNamespace( coord=coord, @@ -1822,12 +1832,16 @@ def _find_longest(values, _token_len, pool): for value in hashes for tp in range(2) ] - self.assertEqual(client.keys, expected_keys) - self.assertEqual(len(set(client.keys)), 4) - self.assertEqual(len(records), 1) - self.assertEqual(records[0]["operation"], "lookup_exists") - self.assertEqual(records[0]["num_keys"], 4) - self.assertEqual(records[0]["num_logical_keys"], 2) + self.assertEqual(client.calls[0], expected_keys) + self.assertEqual(len(records), len(client.calls)) + for call, record in zip(client.calls, records, strict=True): + self.assertEqual(record["operation"], "lookup_exists") + self.assertEqual(record["num_keys"], len(call)) + self.assertEqual( + record["num_logical_keys"], + sum(PoolKey(replace(metadata, tp_rank=0), value.hex()).to_bytes() in call + for value in hashes), + ) if __name__ == "__main__": From d8f811ce2ab0e79feed7ad33a8059265c260b49a Mon Sep 17 00:00:00 2001 From: Ketor Date: Thu, 17 Sep 2026 17:48:20 +0800 Subject: [PATCH 2/3] chore: prepare v2.27.2 dual-engine compatibility release --- CHANGELOG.md | 20 ++++++++++++++++++++ VERSION | 2 +- integration/common/pyproject.toml | 2 +- integration/lmcache/pyproject.toml | 4 ++-- integration/vllm/pyproject.toml | 4 ++-- 5 files changed, 26 insertions(+), 6 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 8ba85a5..7a746cb 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,6 +2,26 @@ ## Unreleased +### v2.27.2 — Legacy and native vLLM compatibility + +- Restore one connector package for engines with the legacy `get_finished` + lifecycle and engines with native `get_transfer_results` support. +- On legacy V1 engines, carry failed request identities through supported + worker metadata and enter the existing scheduler recovery/error path only + after all participating workers finish. No fabricated block IDs or installed + engine source-file patches are used. +- Preserve originating-child routing through nested `MultiConnector` metadata; + do not retain or replay other connectors' historical metadata. +- Keep serialized hybrid loads parked and fenced, quarantine failed external + hits, and preserve cancellation/deferred-free ownership. Native engines + continue to use their native request-failure protocol. +- Register CPU connector and worker lifecycle regressions in CTest, and cover + old/native scheduler and executor behavior in their actual engine runtimes. +- Correct v2.27.1's native-only connector compatibility restriction. Legacy + support requires the V1 worker-metadata, completion aggregation and scheduler + recovery hooks; arbitrary V0 engines or partial protocol backports are not + implied. No functional C++ changes, ABI/wire changes or stored-key migration. + ### v2.27.1 — Request-level hybrid KV load recovery - Adopt vLLM's native `KVConnectorTransferResults` protocol for hybrid and diff --git a/VERSION b/VERSION index c089330..8eda5aa 100644 --- a/VERSION +++ b/VERSION @@ -1 +1 @@ -2.27.1 \ No newline at end of file +2.27.2 \ No newline at end of file diff --git a/integration/common/pyproject.toml b/integration/common/pyproject.toml index 7b23af5..10e3147 100644 --- a/integration/common/pyproject.toml +++ b/integration/common/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "dfkv-common" -version = "2.27.1" +version = "2.27.2" description = "Canonical namespace and pool-key schema shared by dfkv connectors" requires-python = ">=3.9" diff --git a/integration/lmcache/pyproject.toml b/integration/lmcache/pyproject.toml index fff4bd5..451be95 100644 --- a/integration/lmcache/pyproject.toml +++ b/integration/lmcache/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "dfkv-connector" -version = "2.27.1" +version = "2.27.2" description = "LMCache RemoteConnector for the dfkv KV cache (ctypes over libdfkv.so)" readme = "README.md" requires-python = ">=3.9" @@ -14,7 +14,7 @@ authors = [{ name = "Wine93", email = "wine93.info@gmail.com" }] # runtime by path (DFKV_LIB / remote_storage_plugin.dfkv.lib). No bundled .so, # no CPython extension, so the wheel is platform-independent. dependencies = [ - "dfkv-common==2.27.1", + "dfkv-common==2.27.2", "lmcache", "torch", ] diff --git a/integration/vllm/pyproject.toml b/integration/vllm/pyproject.toml index 08c96a6..937241d 100644 --- a/integration/vllm/pyproject.toml +++ b/integration/vllm/pyproject.toml @@ -4,12 +4,12 @@ build-backend = "setuptools.build_meta" [project] name = "dfkv-vllm" -version = "2.27.1" +version = "2.27.2" description = "Direct vLLM KVConnectorBase_V1 connector for dfkv (GPUDirect RDMA, no LMCache)" requires-python = ">=3.12" # vllm + torch are provided by the runtime image; not pinned here. The connector # itself is pure Python (ctypes over libdfkv.so), so there is no native build. -dependencies = ["dfkv-common==2.27.1"] +dependencies = ["dfkv-common==2.27.2"] # Telemetry is opt-in: the OTel SDK is only needed when DFKV_METRICS_ENABLED / # DFKV_TRACE_ENABLED is set. Without this extra the connector stays dependency- From 284c457b0288bbd24dca226d869b19b56698b110 Mon Sep 17 00:00:00 2001 From: Ketor Date: Thu, 17 Sep 2026 18:06:52 +0800 Subject: [PATCH 3/3] fix(vllm): register external connector for MultiConnector statistics --- CHANGELOG.md | 3 ++ integration/vllm/README.md | 7 +++ integration/vllm/src/dfkv_vllm/connector.py | 14 ++++++ .../vllm/tests/test_multi_connector_stats.py | 43 +++++++++++++++++++ test/python/test_dfkv_vllm_connector.py | 19 ++++++++ 5 files changed, 86 insertions(+) create mode 100644 integration/vllm/tests/test_multi_connector_stats.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 7a746cb..4f5b1c6 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -12,6 +12,9 @@ engine source-file patches are used. - Preserve originating-child routing through nested `MultiConnector` metadata; do not retain or replay other connectors' historical metadata. +- Register the external connector's canonical name so MultiConnector can + reconstruct cross-process statistics without terminating API output handling. + Preserve an operator's existing registration of the same class. - Keep serialized hybrid loads parked and fenced, quarantine failed external hits, and preserve cancellation/deferred-free ownership. Native engines continue to use their native request-failure protocol. diff --git a/integration/vllm/README.md b/integration/vllm/README.md index 4774013..b59f6cb 100644 --- a/integration/vllm/README.md +++ b/integration/vllm/README.md @@ -65,6 +65,13 @@ modify installed vLLM source files or require a manually applied engine patch; compatibility is based on these capabilities, not a blanket version cutoff. Single-group full-attention retains its existing block-level recovery. +Direct and nested `MultiConnector` compositions preserve failure routing to +the originating dfkv child without replaying other children's metadata. +Importing the external module also registers `DfkvStoreConnector` by name so +MultiConnector can reconstruct serialized statistics in the API process. +An existing registration of the same class is accepted; a conflicting class +binding is not silently replaced. + Failed request IDs and receive completion are published only after native I/O and GPU writes are fenced. These layouts leave block-ID errors empty. With `kv_load_failure_policy="recompute"`, the engine releases the failed allocation diff --git a/integration/vllm/src/dfkv_vllm/connector.py b/integration/vllm/src/dfkv_vllm/connector.py index e2f68a3..b798f60 100644 --- a/integration/vllm/src/dfkv_vllm/connector.py +++ b/integration/vllm/src/dfkv_vllm/connector.py @@ -23,6 +23,7 @@ KVConnectorKVEvents, KVEventAggregator, ) +from vllm.distributed.kv_transfer.kv_connector.factory import KVConnectorFactory from vllm.distributed.kv_transfer.kv_connector.v1.base import ( KVConnectorBase_V1, KVConnectorMetadata, @@ -504,3 +505,16 @@ def build_prom_metrics( return DfkvStorePromMetrics( vllm_config, metric_types, labelnames, per_engine_labelvalues ) + + +# MultiConnector reconstructs cross-process statistics by class name rather +# than module path. Register the same external class for that supported route. +try: + KVConnectorFactory.register_connector( + "DfkvStoreConnector", __name__, "DfkvStoreConnector" + ) +except ValueError: + if KVConnectorFactory.get_connector_class_by_name( + "DfkvStoreConnector" + ) is not DfkvStoreConnector: + raise diff --git a/integration/vllm/tests/test_multi_connector_stats.py b/integration/vllm/tests/test_multi_connector_stats.py new file mode 100644 index 0000000..4eac8d3 --- /dev/null +++ b/integration/vllm/tests/test_multi_connector_stats.py @@ -0,0 +1,43 @@ +"""External connectors must survive MultiConnector statistics deserialization.""" +import subprocess +import sys + +import pytest + + +@pytest.mark.parametrize("pre_registered", [False, True]) +def test_multiconn_reconstructs_serialized_dfkv_metrics(pre_registered): + # Fresh interpreters exercise both module-path discovery and an operator's + # explicit factory registration, without polluting another test's registry. + setup = ( + "KVConnectorFactory.register_connector(" + "'DfkvStoreConnector', 'dfkv_vllm.connector', 'DfkvStoreConnector')\n" + if pre_registered else "" + ) + code = """ +import json +from vllm.distributed.kv_transfer.kv_connector.factory import KVConnectorFactory +""" + setup + """ +from dfkv_vllm.connector import DfkvStoreConnector +from dfkv_vllm.metrics import DfkvStoreConnectorStats +from vllm.distributed.kv_transfer.kv_connector.v1.multi_connector import ( + MultiConnector, MultiKVConnectorStats, +) +first = DfkvStoreConnectorStats() +first.record_operation('load_get', 0.01, 7, num_bytes=512, + num_failed_keys=2, status='partial_failure') +second = DfkvStoreConnectorStats() +second.record_operation('load_get', 0.02, 3, num_bytes=128) +combined = MultiKVConnectorStats(data={'DfkvStoreConnector': first}) +combined.aggregate(MultiKVConnectorStats(data={'DfkvStoreConnector': second})) +payload = json.loads(json.dumps(combined.to_dict())) +restored = MultiConnector.build_kv_connector_stats(payload).reduce()['DfkvStoreConnector'] +assert restored['load_get_count'] == 2, restored +assert restored['load_get_total_keys'] == 10, restored +assert restored['load_get_total_bytes'] == 640, restored +assert restored['load_get_failed_keys'] == 2, restored +""" + completed = subprocess.run( + [sys.executable, "-c", code], capture_output=True, text=True, + ) + assert completed.returncode == 0, completed.stdout + completed.stderr diff --git a/test/python/test_dfkv_vllm_connector.py b/test/python/test_dfkv_vllm_connector.py index 3c3cf21..544e46b 100644 --- a/test/python/test_dfkv_vllm_connector.py +++ b/test/python/test_dfkv_vllm_connector.py @@ -3,6 +3,7 @@ from __future__ import annotations import logging +from importlib import import_module from importlib.util import find_spec import sys import threading @@ -102,6 +103,20 @@ class KVConnectorOutput: kv_connector_worker_meta: object = None kv_cache_events: object = None + class KVConnectorFactory: + _registry = {} + + @classmethod + def register_connector(cls, name, module_path, class_name): + if name in cls._registry: + raise ValueError("connector already registered") + cls._registry[name] = (module_path, class_name) + + @classmethod + def get_connector_class_by_name(cls, name): + module_path, class_name = cls._registry[name] + return getattr(import_module(module_path), class_name) + class SupportsHMA: pass @@ -161,6 +176,10 @@ def get_manager_class(cls, _spec): ) stub("vllm.distributed.kv_transfer") stub("vllm.distributed.kv_transfer.kv_connector") + stub( + "vllm.distributed.kv_transfer.kv_connector.factory", + KVConnectorFactory=KVConnectorFactory, + ) stub("vllm.distributed.kv_transfer.kv_connector.v1") stub( "vllm.distributed.kv_transfer.kv_connector.v1.base",