Add the KV-canary mock-model end-to-end test harness (#26811)

This commit is contained in:
fzyzcjy
2026-05-31 09:56:00 +08:00
committed by GitHub
parent e188745ee0
commit 27eb139ef7
7 changed files with 401 additions and 0 deletions
@@ -0,0 +1,57 @@
from __future__ import annotations
import io
import os
from typing import ClassVar
from sglang.test.kv_canary.e2e_base import CapturedServerE2EBase
from sglang.test.kv_canary.utils import post_parallel_generate
from sglang.test.mock_model.utils import (
MOCK_MODEL_PATH,
mock_model_server_args,
mock_model_server_env,
)
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
popen_launch_server,
)
class MockModelPerturbE2EBase(CapturedServerE2EBase):
extra_env: ClassVar[dict[str, str]] = {}
extra_server_args: ClassVar[tuple[str, ...]] = ()
@classmethod
def setUpClass(cls) -> None:
server_env = os.environ.copy()
server_env.update(mock_model_server_env())
server_env.update(cls.extra_env)
cls._stdout_buf = io.StringIO()
cls._stderr_buf = io.StringIO()
cls.process = popen_launch_server(
MOCK_MODEL_PATH,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=mock_model_server_args(
*cls.extra_server_args, canary_mode="log"
),
env=server_env,
return_stdout_stderr=(cls._stdout_buf, cls._stderr_buf),
)
def send_parallel_requests(
self,
n: int = 4,
*,
max_new_tokens: int = 256,
timeout: float = 30.0,
) -> list[dict]:
prompts = ["hello world " * 50] * n
return post_parallel_generate(
url=self.base_url + "/generate",
prompts=prompts,
max_new_tokens=max_new_tokens,
timeout=timeout,
)
+131
View File
@@ -0,0 +1,131 @@
from __future__ import annotations
import io
import subprocess
from dataclasses import dataclass
from typing import Any, Sequence
from sglang.bench_serving import run_benchmark
from sglang.srt.utils import kill_process_tree
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
get_benchmark_args,
popen_launch_server,
)
MOCK_MODEL_PATH = "Qwen/Qwen3-0.6B"
_MOCK_MODEL_SERVER_ARGS_NO_CANARY: list[str] = [
"--load-format",
"dummy",
"--sampling-backend",
"pytorch",
"--disable-piecewise-cuda-graph",
]
@dataclass(frozen=True, slots=True, kw_only=True)
class MockModelBenchResult:
result: dict[str, Any]
stdout: str
stderr: str
server_return_code: int | None
@property
def log_text(self) -> str:
return self.stdout + self.stderr
def log_tail(self, length: int = 2000) -> str:
return self.log_text[-length:]
def mock_model_server_args(*extra_args: str, canary_mode: str = "raise") -> list[str]:
return [
*_MOCK_MODEL_SERVER_ARGS_NO_CANARY,
"--kv-canary",
canary_mode,
*extra_args,
]
def mock_model_server_env(*, input_check_enabled: bool = True) -> dict[str, str]:
"""Return env overrides for popen_launch_server in mock-model + canary mode."""
return {}
def run_mock_model_bench_serving(
*,
extra_server_args: Sequence[str],
input_check_enabled: bool = True,
num_prompts: int = 32,
random_input_len: int = 6144,
random_output_len: int = 1024,
) -> MockModelBenchResult:
stdout_buf = io.StringIO()
stderr_buf = io.StringIO()
process: subprocess.Popen[Any] | None = None
try:
process = popen_launch_server(
MOCK_MODEL_PATH,
DEFAULT_URL_FOR_TEST,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=mock_model_server_args(*extra_server_args),
env=mock_model_server_env(input_check_enabled=input_check_enabled),
return_stdout_stderr=(stdout_buf, stderr_buf),
)
args = get_benchmark_args(
base_url=DEFAULT_URL_FOR_TEST,
dataset_name="random",
tokenizer=MOCK_MODEL_PATH,
num_prompts=num_prompts,
random_input_len=random_input_len,
random_output_len=random_output_len,
request_rate=float("inf"),
max_concurrency=num_prompts,
)
args.random_range_ratio = 1.0
args.warmup_requests = 0
args.disable_tqdm = True
result = run_benchmark(args)
server_return_code = process.poll()
bench_result = MockModelBenchResult(
result=result,
stdout=stdout_buf.getvalue(),
stderr=stderr_buf.getvalue(),
server_return_code=server_return_code,
)
_assert_mock_model_bench_succeeded(
bench_result=bench_result,
expected_completed=num_prompts,
)
return bench_result
finally:
if process is not None:
kill_process_tree(process.pid)
def _assert_mock_model_bench_succeeded(
*,
bench_result: MockModelBenchResult,
expected_completed: int,
) -> None:
completed = bench_result.result.get("completed")
if completed != expected_completed:
raise AssertionError(
f"Expected {expected_completed} completed requests, got {completed}"
)
if bench_result.server_return_code is not None:
raise AssertionError(
f"Mock-model server exited with code {bench_result.server_return_code}.\n{bench_result.log_tail()}"
)
if "kv_canary violation:" in bench_result.log_text:
raise AssertionError(
f"Unexpected kv_canary violation in mock-model server log.\n{bench_result.log_tail()}"
)
@@ -1,8 +1,10 @@
import io
import logging
import os
import shlex
import time
import warnings
from typing import ClassVar, Optional
from urllib.parse import urlparse
from sglang.srt.environ import envs
@@ -22,6 +24,14 @@ logger = logging.getLogger(__name__)
class PDDisaggregationServerBase(CustomTestCase):
capture_per_side_logs: ClassVar[bool] = False
extra_prefill_env: ClassVar[dict[str, str]] = {}
extra_decode_env: ClassVar[dict[str, str]] = {}
_prefill_stdout_buf: ClassVar[Optional[io.StringIO]] = None
_prefill_stderr_buf: ClassVar[Optional[io.StringIO]] = None
_decode_stdout_buf: ClassVar[Optional[io.StringIO]] = None
_decode_stderr_buf: ClassVar[Optional[io.StringIO]] = None
@classmethod
def setUpClass(cls):
os.environ["MC_TCP_ENABLE_CONNECTION_POOL"] = "true"
@@ -40,6 +50,11 @@ class PDDisaggregationServerBase(CustomTestCase):
f"{cls.base_host=} {cls.lb_port=} {cls.prefill_port=} {cls.decode_port=} {cls.bootstrap_port=}"
)
cls.process_lb, cls.process_decode, cls.process_prefill = None, None, None
if cls.capture_per_side_logs:
cls._prefill_stdout_buf = io.StringIO()
cls._prefill_stderr_buf = io.StringIO()
cls._decode_stdout_buf = io.StringIO()
cls._decode_stderr_buf = io.StringIO()
cls._fail_fast_stop = None
# config transfer backend and rdma devices
@@ -81,6 +96,12 @@ class PDDisaggregationServerBase(CustomTestCase):
cls.prefill_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=prefill_args,
env=dict(cls.extra_prefill_env),
return_stdout_stderr=(
(cls._prefill_stdout_buf, cls._prefill_stderr_buf)
if cls.capture_per_side_logs
else None
),
)
@classmethod
@@ -102,6 +123,12 @@ class PDDisaggregationServerBase(CustomTestCase):
cls.decode_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=decode_args,
env=dict(cls.extra_decode_env),
return_stdout_stderr=(
(cls._decode_stdout_buf, cls._decode_stderr_buf)
if cls.capture_per_side_logs
else None
),
)
@classmethod
@@ -163,6 +190,20 @@ class PDDisaggregationServerBase(CustomTestCase):
except Exception as e:
print(f"Error killing process {process.pid}: {e}")
if cls.capture_per_side_logs:
for buf in (
cls._prefill_stdout_buf,
cls._prefill_stderr_buf,
cls._decode_stdout_buf,
cls._decode_stderr_buf,
):
if buf is not None:
buf.close()
cls._prefill_stdout_buf = None
cls._prefill_stderr_buf = None
cls._decode_stdout_buf = None
cls._decode_stderr_buf = None
# wait for 5 seconds
time.sleep(5)