fix unit test

Signed-off-by: yewentao256 <zhyanwentao@126.com>
This commit is contained in:
yewentao256
2026-03-02 22:11:52 +00:00
parent b9685019bd
commit 09cd90a196
2 changed files with 71 additions and 14 deletions
+39 -6
View File
@@ -40,6 +40,26 @@ from .utils import EOS_TOKEN_ID, create_requests, create_scheduler, mock_kv
pytestmark = pytest.mark.cpu_test
def _get_remote_waiting_queue(scheduler: Scheduler):
return getattr(scheduler, "waiting_for_remote_kvs", None)
def _num_waiting_requests(scheduler: Scheduler) -> int:
remote_waiting = _get_remote_waiting_queue(scheduler)
return len(scheduler.waiting) + (len(remote_waiting) if remote_waiting else 0)
def _get_remote_waiting_requests(scheduler: Scheduler) -> list[Request]:
remote_waiting = _get_remote_waiting_queue(scheduler)
if remote_waiting is not None:
return list(remote_waiting)
return [
req
for req in scheduler.waiting
if req.status == RequestStatus.WAITING_FOR_REMOTE_KVS
]
def test_add_requests():
scheduler = create_scheduler()
requests = create_requests(num_requests=10)
@@ -1120,7 +1140,8 @@ def _step_until_kv_transfer_finished(scheduler: Scheduler, req_ids: list[str]):
# Requests should first transition to WAITING_FOR_REMOTE_KVS
output = scheduler.schedule()
assert len(scheduler.waiting) == len(req_ids)
assert _num_waiting_requests(scheduler) == len(req_ids)
assert len(_get_remote_waiting_requests(scheduler)) == len(req_ids)
assert len(scheduler.running) == 0
assert len(output.scheduled_new_reqs) == 0
for req in scheduler.requests.values():
@@ -1139,7 +1160,8 @@ def _step_until_kv_transfer_finished(scheduler: Scheduler, req_ids: list[str]):
# Simulate KV transfer completion using KVConnectorOutput.finished_recving
output = scheduler.schedule()
assert len(scheduler.waiting) == len(req_ids)
assert _num_waiting_requests(scheduler) == len(req_ids)
assert len(_get_remote_waiting_requests(scheduler)) == len(req_ids)
assert len(scheduler.running) == 0
MODEL_RUNNER_OUTPUT = ModelRunnerOutput(
@@ -1546,7 +1568,7 @@ def test_kv_connector_handles_preemption(is_async, use_ec_connector, ec_role):
# All can be scheduled - 1st token.
output = scheduler.schedule()
if is_async:
assert len(scheduler.waiting) == 2
assert _num_waiting_requests(scheduler) == 2
assert scheduler.running == []
_step_until_kv_transfer_finished(scheduler, req_ids)
output = scheduler.schedule()
@@ -1604,7 +1626,9 @@ def test_kv_connector_handles_preemption(is_async, use_ec_connector, ec_role):
# This will have a local and remote cache hit.
output = scheduler.schedule()
if is_async:
waiting_req_ids = [req.request_id for req in scheduler.waiting]
waiting_req_ids = [
req.request_id for req in _get_remote_waiting_requests(scheduler)
]
assert len(waiting_req_ids) == 1
_step_until_kv_transfer_finished(scheduler, waiting_req_ids)
output = scheduler.schedule()
@@ -3614,6 +3638,10 @@ def test_prepend_skipped_requests_order():
# simulate first 2 waiting requests are waiting for remote KVs
for req in expected_waiting_reqs[:2]:
req.status = RequestStatus.WAITING_FOR_REMOTE_KVS
remote_waiting = _get_remote_waiting_queue(scheduler)
if remote_waiting is not None:
scheduler.waiting.remove_request(req)
remote_waiting.add_request(req)
# schedule step
# expect the first 2 waiting to be skipped, the third running,
@@ -3623,8 +3651,13 @@ def test_prepend_skipped_requests_order():
# pop the third request which is expected to be running
expected_waiting_reqs.pop(2)
# verify waiting order is preserved
assert list(scheduler.waiting) == expected_waiting_reqs
# verify waiting order is preserved for schedulable requests.
remote_waiting = _get_remote_waiting_queue(scheduler)
if remote_waiting is not None:
assert list(scheduler.waiting) == expected_waiting_reqs[2:]
assert list(remote_waiting) == expected_waiting_reqs[:2]
else:
assert list(scheduler.waiting) == expected_waiting_reqs
def test_abort_request_waiting_for_remote_kvs():
@@ -17,6 +17,26 @@ from .utils import (
)
def _get_remote_waiting_queue(scheduler: Scheduler):
return getattr(scheduler, "waiting_for_remote_kvs", None)
def _num_waiting_requests(scheduler: Scheduler) -> int:
remote_waiting = _get_remote_waiting_queue(scheduler)
return len(scheduler.waiting) + (len(remote_waiting) if remote_waiting else 0)
def _get_remote_waiting_requests(scheduler: Scheduler) -> list[Request]:
remote_waiting = _get_remote_waiting_queue(scheduler)
if remote_waiting is not None:
return list(remote_waiting)
return [
req
for req in scheduler.waiting
if req.status == RequestStatus.WAITING_FOR_REMOTE_KVS
]
def _make_get_num_new_matched_tokens(
req_num_new_matched_tokens: dict[str, int],
async_load,
@@ -76,8 +96,8 @@ def test_async_load_failure(
scheduler_output = scheduler.schedule()
assert len(scheduler.waiting) == 3
for request in scheduler.waiting:
assert _num_waiting_requests(scheduler) == 3
for request in _get_remote_waiting_requests(scheduler):
assert request.num_computed_tokens == 0
assert request.status == RequestStatus.WAITING_FOR_REMOTE_KVS
assert scheduler.connector.get_num_new_matched_tokens.call_count == 3
@@ -96,8 +116,8 @@ def test_async_load_failure(
min_invalid_block_idx = min(invalid_block_idxs)
assert len(scheduler.waiting) == 3
for request in scheduler.waiting:
assert _num_waiting_requests(scheduler) == 3
for request in _get_remote_waiting_requests(scheduler):
if request.request_id == request2.request_id:
assert request.num_computed_tokens == (
min_invalid_block_idx * scheduler.block_size
@@ -303,8 +323,10 @@ def test_async_progressive_load_failure(
scheduler_output = scheduler.schedule()
assert len(scheduler.waiting) == 1
assert scheduler.waiting.peek_request().request_id == request.request_id
assert _num_waiting_requests(scheduler) == 1
remote_waiting_reqs = _get_remote_waiting_requests(scheduler)
assert len(remote_waiting_reqs) == 1
assert remote_waiting_reqs[0].request_id == request.request_id
assert request.num_computed_tokens == 0
assert request.status == RequestStatus.WAITING_FOR_REMOTE_KVS
assert scheduler.connector.get_num_new_matched_tokens.call_count == 1
@@ -325,8 +347,10 @@ def test_async_progressive_load_failure(
min_invalid_block_idx = min(min_invalid_block_idx, invalid_block_idx)
assert len(scheduler.waiting) == 1
assert scheduler.waiting.peek_request().request_id == request.request_id
assert _num_waiting_requests(scheduler) == 1
remote_waiting_reqs = _get_remote_waiting_requests(scheduler)
assert len(remote_waiting_reqs) == 1
assert remote_waiting_reqs[0].request_id == request.request_id
assert request.num_computed_tokens == (
min_invalid_block_idx * scheduler.block_size
)