Skip to content
Closed
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
84 changes: 84 additions & 0 deletions tests/v1/engine/test_connector_poller.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,84 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""ConnectorPoller: services a connector only while the engine waits."""

import threading
import time
from concurrent.futures import Future

import pytest

from vllm.v1.engine.core import ConnectorPoller


class _CountingConnector:
def __init__(self, fail: bool = False) -> None:
self.calls = 0
self.fail = fail
self.thread_names: set[str] = set()

def poll_pending_work(self) -> None:
self.calls += 1
self.thread_names.add(threading.current_thread().name)
if self.fail:
raise RuntimeError("sweep failed")


def _resolve_after(future: Future, delay_s: float, value: object = "out") -> None:
def run():
time.sleep(delay_s)
future.set_result(value)

threading.Thread(target=run, daemon=True).start()


def test_poller_sweeps_only_while_waiting():
connector = _CountingConnector()
poller = ConnectorPoller(connector, interval_s=0.001)
try:
time.sleep(0.02)
assert connector.calls == 0

future: Future = Future()
_resolve_after(future, 0.05)
assert poller.wait(future) == "out"
assert connector.calls > 0
assert connector.thread_names == {"kv-connector-poller"}

settled = connector.calls
time.sleep(0.02)
assert connector.calls == settled
finally:
poller.close()


def test_poller_reraises_sweep_error_on_engine_thread():
connector = _CountingConnector(fail=True)
poller = ConnectorPoller(connector, interval_s=0.001)
try:
future: Future = Future()
_resolve_after(future, 0.03)
with pytest.raises(RuntimeError, match="sweep failed"):
poller.wait(future)
assert connector.calls == 1

# A later wait runs sweeps again; the stored error was consumed.
connector.fail = False
future = Future()
_resolve_after(future, 0.03)
assert poller.wait(future) == "out"
assert connector.calls > 1
finally:
poller.close()


def test_poller_propagates_future_exception():
connector = _CountingConnector()
poller = ConnectorPoller(connector, interval_s=0.001)
try:
future: Future = Future()
future.set_exception(ValueError("model failed"))
with pytest.raises(ValueError, match="model failed"):
poller.wait(future)
finally:
poller.close()
36 changes: 33 additions & 3 deletions tests/v1/kv_offload/tiering/p2p/test_data_transport.py
Original file line number Diff line number Diff line change
Expand Up @@ -324,7 +324,9 @@ def test_non_ucx_backends_passes_backends_kwarg(self):
):
NixlTransport("test:1", self._make_view(), backends=["MOONCAKE"])

config_fn.assert_called_once_with(backends=["MOONCAKE"], capture_telemetry=True)
config_fn.assert_called_once_with(
backends=["MOONCAKE"], capture_telemetry=True, sync_mode=None
)
# num_threads must NOT be passed on the non-UCX branch.
assert "num_threads" not in config_fn.call_args.kwargs

Expand All @@ -340,7 +342,9 @@ def test_ucx_only_passes_num_threads(self):
):
NixlTransport("test:1", self._make_view(), num_threads=8)

config_fn.assert_called_once_with(num_threads=8, capture_telemetry=True)
config_fn.assert_called_once_with(
num_threads=8, capture_telemetry=True, sync_mode=None
)
assert "backends" not in config_fn.call_args.kwargs

def test_default_backends_is_ucx_only(self):
Expand All @@ -356,4 +360,30 @@ def test_default_backends_is_ucx_only(self):
NixlTransport("test:1", self._make_view())

# Default num_threads=4, no backends kwarg.
config_fn.assert_called_once_with(num_threads=4, capture_telemetry=True)
config_fn.assert_called_once_with(
num_threads=4, capture_telemetry=True, sync_mode=None
)

def test_agent_runs_in_strict_thread_sync_when_available(self):
"""With a NIXL that exposes the sync enum, the agent serializes its
API so peer registration may run on the worker thread."""
import sys
import types

fake_module = types.ModuleType("fake_nixl_api")
fake_module.nixl_thread_sync_t = types.SimpleNamespace(
NIXL_THREAD_SYNC_STRICT="strict"
)
config_fn = MagicMock(return_value=MagicMock(name="cfg"))
config_fn.__module__ = fake_module.__name__
agent_cls = MagicMock()
with (
patch.dict(sys.modules, {fake_module.__name__: fake_module}),
patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", agent_cls),
patch(
"vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgentConfig", config_fn
),
):
NixlTransport("test:1", self._make_view())

assert config_fn.call_args.kwargs["sync_mode"] == "strict"
10 changes: 10 additions & 0 deletions tests/v1/kv_offload/tiering/p2p/test_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -1073,6 +1073,16 @@ def add_remote_peer(
"block_len": block_len,
}

def add_remote_peer_async(
self, peer_id, agent_metadata, base_addr, num_blocks, block_len
):
from concurrent.futures import Future

self.add_remote_peer(peer_id, agent_metadata, base_addr, num_blocks, block_len)
future: Future[None] = Future()
future.set_result(None)
return future

def remove_remote_peer(self, peer_id: str) -> None:
self._remote_peers.pop(peer_id, None)

Expand Down
59 changes: 59 additions & 0 deletions tests/v1/kv_offload/tiering/p2p/test_sessions.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@

import time
from collections.abc import Sequence
from concurrent.futures import Future

import numpy as np
import pytest
Expand Down Expand Up @@ -86,6 +87,10 @@ def __init__(
self._poll_failed: list[int] = []
self._cancel_still_inflight: set[int] = set()
self._cancel_calls: list[tuple[list[int], str]] = []
# When True, add_remote_peer_async parks the registration in
# pending_registrations for the test to complete or fail.
self.defer_registration = False
self.pending_registrations: dict[str, Future[None]] = {}

@property
def base_addr(self) -> int:
Expand Down Expand Up @@ -116,6 +121,25 @@ def add_remote_peer(
"block_len": block_len,
}

def add_remote_peer_async(
self, peer_id, agent_metadata, base_addr, num_blocks, block_len
) -> Future[None]:
future: Future[None] = Future()
if self.defer_registration:
self.pending_registrations[peer_id] = future
future.add_done_callback(
lambda f: (
f.exception() is None
and self.add_remote_peer(
peer_id, agent_metadata, base_addr, num_blocks, block_len
)
)
)
return future
self.add_remote_peer(peer_id, agent_metadata, base_addr, num_blocks, block_len)
future.set_result(None)
return future

def remove_remote_peer(self, peer_id: str) -> None:
self._remote_peers.pop(peer_id, None)

Expand Down Expand Up @@ -390,6 +414,41 @@ def test_peer_connect_triggers_add_remote_and_ack(self):
ack = next(m for m in conn._sent if m[TYPE_KEY] == ConnectAckMsg.TYPE)
assert ack[ConnectAckMsg.PEER_ID] == "local:9000"

def test_connect_ack_waits_for_peer_registration(self):
"""ConnectAck is deferred until the transport has registered the peer."""
transport = FakeDataTransport()
transport.defer_registration = True
session, conn, _ = _make_session(transport=transport)
conn.enqueue(_peer_connect_msg())
session.poll()
assert not any(m[TYPE_KEY] == ConnectAckMsg.TYPE for m in conn._sent)
assert "peer:8000" not in transport._remote_peers
assert session.has_pending_work
session.poll()
assert not any(m[TYPE_KEY] == ConnectAckMsg.TYPE for m in conn._sent)

transport.pending_registrations["peer:8000"].set_result(None)
session.poll()
assert "peer:8000" in transport._remote_peers
ack = next(m for m in conn._sent if m[TYPE_KEY] == ConnectAckMsg.TYPE)
assert ack[ConnectAckMsg.PEER_ID] == "local:9000"
assert not session.has_pending_work

def test_failed_peer_registration_rejects_peer(self):
"""A registration error rejects the peer like a validation failure."""
transport = FakeDataTransport()
transport.defer_registration = True
session, conn, _ = _make_session(transport=transport)
conn.enqueue(_peer_connect_msg())
session.poll()
transport.pending_registrations["peer:8000"].set_exception(
RuntimeError("agent metadata rejected")
)
session.poll()
assert not session.alive
assert not any(m[TYPE_KEY] == ConnectAckMsg.TYPE for m in conn._sent)
assert not session.has_pending_work

def test_connect_ack_makes_session_ready(self):
"""Session.ready becomes True after ConnectAckMsg."""
session, conn, _ = _make_session()
Expand Down
40 changes: 40 additions & 0 deletions tests/v1/kv_offload/tiering/test_tiering_offloading.py
Original file line number Diff line number Diff line change
Expand Up @@ -148,6 +148,46 @@ def get_stats(self) -> OffloadingConnectorStats | None:
return stats


class ServeRecordingSecondaryTierManager(MetricsSecondaryTierManager):
"""Test-only secondary tier that records service sweeps."""

def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.finished_polls = 0
self.serve_calls = 0

def get_finished_jobs(self) -> Iterable[JobResult]:
self.finished_polls += 1
return ()

def serve_external_requests(self, parent) -> None:
self.serve_calls += 1


def test_poll_pending_work_services_tiers_between_steps():
"""poll_pending_work() collects finished jobs and serves every tier
without a scheduler step and without consuming the per-step gate."""
mock_region = _mock_mmap_region(5)
primary = CPUPrimaryTierOffloadingManager(num_chunks=5, mmap_region=mock_region)
tier = ServeRecordingSecondaryTierManager(
offloading_spec=_MOCK_OFFLOADING_SPEC,
primary_kv_view=mock_region.create_kv_memoryview(),
tier_type="recording",
)
manager = TieringOffloadingManager(primary_tier=primary, secondary_tiers=[tier])

manager.poll_pending_work()
manager.poll_pending_work()
assert tier.finished_polls == 2
assert tier.serve_calls == 2
assert manager._processed_jobs_this_step is False

ctx = ScheduleEndContext(new_req_ids=[], preempted_req_ids=())
manager.on_schedule_end(ctx)
assert tier.finished_polls == 3
assert tier.serve_calls == 3


def test_tiering_spec_collects_secondary_metric_definitions(monkeypatch):
monkeypatch.setitem(
SecondaryTierFactory._registry,
Expand Down
12 changes: 12 additions & 0 deletions vllm/distributed/kv_transfer/kv_connector/v1/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -638,6 +638,18 @@ def has_pending_push_work(self) -> bool:
# scheduler alive (e.g. extend has_unfinished_requests).
return False

def poll_pending_work(self) -> None:
"""Advance scheduler-side work while the engine waits on a model step.

The engine core calls this at a bounded interval between scheduling
iterations, never concurrently with any other scheduler-side method,
so a connector that services peers from the scheduler side (control
messages, transfer completions) is not limited to one service window
per step. Implementations must not touch scheduler state and must
return promptly when there is nothing to do.
"""
return

@classmethod
def get_required_kvcache_layout(cls, vllm_config: "VllmConfig") -> str | None:
"""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -588,6 +588,10 @@ def has_pending_block_frees(self) -> bool:
def has_pending_push_work(self) -> bool:
return any(c.has_pending_push_work() for c in self._connectors)

def poll_pending_work(self) -> None:
for c in self._connectors:
c.poll_pending_work()

@classmethod
def get_required_kvcache_layout(cls, vllm_config: "VllmConfig") -> str | None:
"""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1839,6 +1839,9 @@ def has_pending_push_work(self) -> bool:
"""
return bool(self._jobs) or self.manager.has_pending_work()

def poll_pending_work(self) -> None:
self.manager.poll_pending_work()

def update_connector_output(self, connector_output: KVConnectorOutput):
"""
Update KVConnector state from worker-side connectors output.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -202,6 +202,10 @@ def has_pending_push_work(self) -> bool:
assert self.connector_scheduler is not None
return self.connector_scheduler.has_pending_push_work()

def poll_pending_work(self) -> None:
assert self.connector_scheduler is not None
self.connector_scheduler.poll_pending_work()

def update_connector_output(self, connector_output: KVConnectorOutput):
assert self.connector_scheduler is not None
self.connector_scheduler.update_connector_output(connector_output)
Expand Down
Loading
Loading