[diffusion] feat: add performance mode server args (#24491)

This commit is contained in:
Mick
2026-05-14 00:57:46 +08:00
committed by GitHub
parent e2290b155a
commit ff70aeac30
79 changed files with 1734 additions and 532 deletions
@@ -59,12 +59,15 @@ def test_mooncake_te_condition(server_args: ServerArgs) -> bool:
ib_device=ib_device,
)
with patch(
"sglang.srt.distributed.device_communicators.mooncake_transfer_engine.init_mooncake_transfer_engine",
side_effect=_fake_init_mooncake_transfer_engine,
), patch(
"sglang.srt.model_executor.model_runner.get_local_ip_auto",
return_value="127.0.0.1",
with (
patch(
"sglang.srt.distributed.device_communicators.mooncake_transfer_engine.init_mooncake_transfer_engine",
side_effect=_fake_init_mooncake_transfer_engine,
),
patch(
"sglang.srt.model_executor.model_runner.get_local_ip_auto",
return_value="127.0.0.1",
),
):
ModelRunner.init_shared_mooncake_transfer_engine(dummy_runner)
+7 -5
View File
@@ -30,11 +30,13 @@ class TestTokenizerBatchEncode(unittest.TestCase):
)
self.port_args = PortArgs.init_new(self.server_args)
with patch("zmq.asyncio.Context"), patch(
"sglang.srt.utils.get_zmq_socket"
), patch(
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
) as mock_tokenizer:
with (
patch("zmq.asyncio.Context"),
patch("sglang.srt.utils.get_zmq_socket"),
patch(
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
) as mock_tokenizer,
):
mock_tokenizer.return_value = Mock(vocab_size=32000)
self.tokenizer_manager = TokenizerManager(self.server_args, self.port_args)
+28 -20
View File
@@ -37,11 +37,13 @@ class TestInputFormatDetection(unittest.TestCase):
self.server_args = ServerArgs(model_path=DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
self.port_args = PortArgs.init_new(self.server_args)
with patch("zmq.asyncio.Context"), patch(
"sglang.srt.utils.network.get_zmq_socket"
), patch(
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
) as mock_tokenizer:
with (
patch("zmq.asyncio.Context"),
patch("sglang.srt.utils.network.get_zmq_socket"),
patch(
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
) as mock_tokenizer,
):
mock_tokenizer.return_value = Mock(vocab_size=32000)
self.tokenizer_manager = TokenizerManager(self.server_args, self.port_args)
@@ -133,11 +135,13 @@ class TestTokenizerInputPreparation(unittest.TestCase):
self.server_args = ServerArgs(model_path=DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
self.port_args = PortArgs.init_new(self.server_args)
with patch("zmq.asyncio.Context"), patch(
"sglang.srt.utils.network.get_zmq_socket"
), patch(
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
) as mock_tokenizer:
with (
patch("zmq.asyncio.Context"),
patch("sglang.srt.utils.network.get_zmq_socket"),
patch(
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
) as mock_tokenizer,
):
mock_tokenizer.return_value = Mock(vocab_size=32000)
self.tokenizer_manager = TokenizerManager(self.server_args, self.port_args)
@@ -191,11 +195,13 @@ class TestTokenizerResultExtraction(unittest.TestCase):
self.server_args = ServerArgs(model_path=DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
self.port_args = PortArgs.init_new(self.server_args)
with patch("zmq.asyncio.Context"), patch(
"sglang.srt.utils.network.get_zmq_socket"
), patch(
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
) as mock_tokenizer:
with (
patch("zmq.asyncio.Context"),
patch("sglang.srt.utils.network.get_zmq_socket"),
patch(
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
) as mock_tokenizer,
):
mock_tokenizer.return_value = Mock(vocab_size=32000)
self.tokenizer_manager = TokenizerManager(self.server_args, self.port_args)
@@ -313,11 +319,13 @@ class TestTokenizerManagerIntegration(unittest.TestCase):
self.server_args = ServerArgs(model_path=DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
self.port_args = PortArgs.init_new(self.server_args)
with patch("zmq.asyncio.Context"), patch(
"sglang.srt.utils.network.get_zmq_socket"
), patch(
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
) as mock_tokenizer:
with (
patch("zmq.asyncio.Context"),
patch("sglang.srt.utils.network.get_zmq_socket"),
patch(
"sglang.srt.utils.hf_transformers_utils.get_tokenizer"
) as mock_tokenizer,
):
mock_tokenizer.return_value = Mock(vocab_size=32000)
self.tokenizer_manager = TokenizerManager(self.server_args, self.port_args)
@@ -405,12 +405,15 @@ class TestBenchmarkDatasetsAPI(unittest.TestCase):
fake_mmmu_dataset = _FakeMMMUDataset(
[{"image_1": Image.new("RGB", (4, 4), color="white"), "question": "q"}]
)
with patch(
"sglang.benchmark.datasets.mmmu.get_processor",
return_value=self.processor,
), patch(
"sglang.benchmark.datasets.mmmu.load_dataset",
return_value=fake_mmmu_dataset,
with (
patch(
"sglang.benchmark.datasets.mmmu.get_processor",
return_value=self.processor,
),
patch(
"sglang.benchmark.datasets.mmmu.load_dataset",
return_value=fake_mmmu_dataset,
),
):
mmmu_args = make_args(dataset_name="mmmu", num_prompts=1)
mmmu_rows = get_dataset(mmmu_args, self.tokenizer, model_id="dummy-model")
@@ -73,20 +73,16 @@ def test_parallel_group_construction_tp8_attn_cp2():
# Mock the distributed backend
# Note: get_rank() returns 0 because we're testing from a single process,
# but initialize_model_parallel() still creates all groups for all ranks
with patch.object(parallel_state, "_WORLD", None), patch.object(
parallel_state, "_TP", None
), patch.object(parallel_state, "_ATTN_CP", None), patch.object(
parallel_state, "_ATTN_TP", None
), patch.object(
parallel_state, "_PP", None
), patch(
"torch.distributed.is_initialized", return_value=True
), patch(
"torch.distributed.get_world_size", return_value=world_size
), patch(
"torch.distributed.get_rank", return_value=0
), patch(
"torch.distributed.get_backend", return_value="nccl"
with (
patch.object(parallel_state, "_WORLD", None),
patch.object(parallel_state, "_TP", None),
patch.object(parallel_state, "_ATTN_CP", None),
patch.object(parallel_state, "_ATTN_TP", None),
patch.object(parallel_state, "_PP", None),
patch("torch.distributed.is_initialized", return_value=True),
patch("torch.distributed.get_world_size", return_value=world_size),
patch("torch.distributed.get_rank", return_value=0),
patch("torch.distributed.get_backend", return_value="nccl"),
):
# Mock init_model_parallel_group to capture the groups being created
@@ -101,11 +97,14 @@ def test_parallel_group_construction_tp8_attn_cp2():
mock_group.device_group = Mock()
return mock_group
with patch.object(
parallel_state,
"init_model_parallel_group",
side_effect=mock_init_model_parallel_group,
), patch.object(parallel_state, "get_world_group") as mock_world_group:
with (
patch.object(
parallel_state,
"init_model_parallel_group",
side_effect=mock_init_model_parallel_group,
),
patch.object(parallel_state, "get_world_group") as mock_world_group,
):
# Mock world group
mock_world = Mock()
@@ -173,22 +172,17 @@ def test_parallel_group_construction_tp8_moe_ep4_cp2():
world_size = 8
# Mock the distributed backend
with patch.object(parallel_state, "_WORLD", None), patch.object(
parallel_state, "_TP", None
), patch.object(parallel_state, "_MOE_EP", None), patch.object(
parallel_state, "_MOE_DP", None
), patch.object(
parallel_state, "_MOE_TP", None
), patch.object(
parallel_state, "_PP", None
), patch(
"torch.distributed.is_initialized", return_value=True
), patch(
"torch.distributed.get_world_size", return_value=world_size
), patch(
"torch.distributed.get_rank", return_value=0
), patch(
"torch.distributed.get_backend", return_value="nccl"
with (
patch.object(parallel_state, "_WORLD", None),
patch.object(parallel_state, "_TP", None),
patch.object(parallel_state, "_MOE_EP", None),
patch.object(parallel_state, "_MOE_DP", None),
patch.object(parallel_state, "_MOE_TP", None),
patch.object(parallel_state, "_PP", None),
patch("torch.distributed.is_initialized", return_value=True),
patch("torch.distributed.get_world_size", return_value=world_size),
patch("torch.distributed.get_rank", return_value=0),
patch("torch.distributed.get_backend", return_value="nccl"),
):
# Mock init_model_parallel_group to capture the groups being created
@@ -203,11 +197,14 @@ def test_parallel_group_construction_tp8_moe_ep4_cp2():
mock_group.device_group = Mock()
return mock_group
with patch.object(
parallel_state,
"init_model_parallel_group",
side_effect=mock_init_model_parallel_group,
), patch.object(parallel_state, "get_world_group") as mock_world_group:
with (
patch.object(
parallel_state,
"init_model_parallel_group",
side_effect=mock_init_model_parallel_group,
),
patch.object(parallel_state, "get_world_group") as mock_world_group,
):
# Mock world group
mock_world = Mock()
@@ -113,11 +113,15 @@ class TestLoRAQwen3_8BLogprobDiff(CustomTestCase):
of internal param names that would break LoRA auto-detection."""
model = _build_qwen3_mock()
with patch("sglang.srt.layers.linear.LinearBase", _MockLinearBase), patch(
"sglang.srt.layers.moe.fused_moe_triton.layer.FusedMoE", _MockFusedMoE
), patch(
"sglang.srt.layers.vocab_parallel_embedding.ParallelLMHead",
_MockParallelLMHead,
with (
patch("sglang.srt.layers.linear.LinearBase", _MockLinearBase),
patch(
"sglang.srt.layers.moe.fused_moe_triton.layer.FusedMoE", _MockFusedMoE
),
patch(
"sglang.srt.layers.vocab_parallel_embedding.ParallelLMHead",
_MockParallelLMHead,
),
):
detected = auto_detect_lora_target_modules(model)
@@ -161,14 +161,17 @@ class TestGenerationModels(CustomTestCase):
) as hf_runner:
hf_outputs = hf_runner.forward(prompts, max_new_tokens=max_new_tokens)
with env_ctx, SRTRunner(
model_path,
tp_size=model_case.tp_size,
torch_dtype=torch_dtype,
model_type="generation",
trust_remote_code=model_case.trust_remote_code,
attention_backend=model_case.attention_backend,
) as srt_runner:
with (
env_ctx,
SRTRunner(
model_path,
tp_size=model_case.tp_size,
torch_dtype=torch_dtype,
model_type="generation",
trust_remote_code=model_case.trust_remote_code,
attention_backend=model_case.attention_backend,
) as srt_runner,
):
srt_outputs = srt_runner.forward(prompts, max_new_tokens=max_new_tokens)
check_close_model_outputs(
+3 -2
View File
@@ -86,8 +86,9 @@ class TestLMHeadFP32(unittest.TestCase):
state.update(called=True, ooperationp="linear", a=x.dtype, b=w.dtype)
return original_linear(x, w, bias)
with patch("torch.matmul", new=probe_matmul), patch(
"torch.nn.functional.linear", new=probe_linear
with (
patch("torch.matmul", new=probe_matmul),
patch("torch.nn.functional.linear", new=probe_linear),
):
logits = logprocessor._get_logits(hidden_state, head, meta)
self.assertEqual(hidden_state.dtype, hidden_state_dtype)
@@ -31,9 +31,10 @@ class TestRetractDecode(CustomTestCase):
cls.model = DEFAULT_MODEL_NAME_FOR_TEST
cls.base_url = DEFAULT_URL_FOR_TEST
launch_args = ["--chunked-prefill-size", "128"] + cls.other_args
with envs.SGLANG_TEST_RETRACT.override(
True
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(1):
with (
envs.SGLANG_TEST_RETRACT.override(True),
envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(1),
):
cls.process = popen_launch_server(
cls.model,
cls.base_url,
@@ -331,9 +331,10 @@ class TestAbortWithRunningTimeout(CustomTestCase):
def setUpClass(cls):
cls.model = DEFAULT_MODEL_NAME_FOR_TEST
cls.base_url = DEFAULT_URL_FOR_TEST
with envs.SGLANG_REQ_RUNNING_TIMEOUT.override(
0.001
), envs.SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION.override(False):
with (
envs.SGLANG_REQ_RUNNING_TIMEOUT.override(0.001),
envs.SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION.override(False),
):
cls.process = popen_launch_server(
cls.model,
cls.base_url,
@@ -796,9 +796,10 @@ class TestStreamingSessionRetractMixedChunk(TestStreamingSession):
def setUpClass(cls):
cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
cls.base_url = DEFAULT_URL_FOR_TEST
with envs.SGLANG_TEST_RETRACT.override(
True
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2):
with (
envs.SGLANG_TEST_RETRACT.override(True),
envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2),
):
cls.process = popen_launch_server(
cls.model,
cls.base_url,
@@ -825,9 +826,10 @@ class TestStreamingSessionRetractLargePage(TestStreamingSession):
def setUpClass(cls):
cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
cls.base_url = DEFAULT_URL_FOR_TEST
with envs.SGLANG_TEST_RETRACT.override(
True
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2):
with (
envs.SGLANG_TEST_RETRACT.override(True),
envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2),
):
cls.process = popen_launch_server(
cls.model,
cls.base_url,
@@ -856,9 +858,10 @@ class TestStreamingSessionEagle(TestStreamingSession):
def setUpClass(cls):
cls.model = DEFAULT_TARGET_MODEL_EAGLE3
cls.base_url = DEFAULT_URL_FOR_TEST
with envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(
2
), envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(True):
with (
envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2),
envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(True),
):
cls.process = popen_launch_server(
cls.model,
cls.base_url,
@@ -897,12 +900,10 @@ class TestStreamingSessionEagleV2(TestStreamingSession):
def setUpClass(cls):
cls.model = DEFAULT_TARGET_MODEL_EAGLE3
cls.base_url = DEFAULT_URL_FOR_TEST
with envs.SGLANG_ENABLE_SPEC_V2.override(
True
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(
2
), envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(
True
with (
envs.SGLANG_ENABLE_SPEC_V2.override(True),
envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2),
envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(True),
):
cls.process = popen_launch_server(
cls.model,
@@ -944,12 +945,10 @@ class TestStreamingSessionEagleRetractLargePage(TestStreamingSession):
def setUpClass(cls):
cls.model = DEFAULT_TARGET_MODEL_EAGLE3
cls.base_url = DEFAULT_URL_FOR_TEST
with envs.SGLANG_TEST_RETRACT.override(
True
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(
2
), envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(
True
with (
envs.SGLANG_TEST_RETRACT.override(True),
envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2),
envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(True),
):
cls.process = popen_launch_server(
cls.model,
@@ -991,14 +990,11 @@ class TestStreamingSessionEagleV2RetractLargePage(TestStreamingSession):
def setUpClass(cls):
cls.model = DEFAULT_TARGET_MODEL_EAGLE3
cls.base_url = DEFAULT_URL_FOR_TEST
with envs.SGLANG_ENABLE_SPEC_V2.override(
True
), envs.SGLANG_TEST_RETRACT.override(
True
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(
2
), envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(
True
with (
envs.SGLANG_ENABLE_SPEC_V2.override(True),
envs.SGLANG_TEST_RETRACT.override(True),
envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2),
envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(True),
):
cls.process = popen_launch_server(
cls.model,
@@ -70,9 +70,10 @@ class TestStreamingSessionSWARetractLargePage(TestStreamingSession):
def setUpClass(cls):
cls.model = SWA_MODEL
cls.base_url = DEFAULT_URL_FOR_TEST
with envs.SGLANG_TEST_RETRACT.override(
True
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2):
with (
envs.SGLANG_TEST_RETRACT.override(True),
envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2),
):
cls.process = popen_launch_server(
cls.model,
cls.base_url,
@@ -100,9 +101,10 @@ class TestStreamingSessionSWARetractMixedChunk(TestStreamingSession):
def setUpClass(cls):
cls.model = SWA_MODEL
cls.base_url = DEFAULT_URL_FOR_TEST
with envs.SGLANG_TEST_RETRACT.override(
True
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2):
with (
envs.SGLANG_TEST_RETRACT.override(True),
envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2),
):
cls.process = popen_launch_server(
cls.model,
cls.base_url,
+4 -6
View File
@@ -53,12 +53,10 @@ class TestDFlashServerBase(CustomTestCase, MatchedStopMixin, GSM8KMixin):
old_value = os.environ.get("SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN")
os.environ["SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN"] = "1"
try:
with envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(
1
), envs.SGLANG_SPEC_NAN_DETECTION.override(
True
), envs.SGLANG_SPEC_OOB_DETECTION.override(
True
with (
envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(1),
envs.SGLANG_SPEC_NAN_DETECTION.override(True),
envs.SGLANG_SPEC_OOB_DETECTION.override(True),
):
cls.process = popen_launch_server(
cls.model,
@@ -49,9 +49,10 @@ class TestDeepseekV3FP4MTP(CustomTestCase):
"--model-loader-extra-config",
'{"enable_multithread_load": true,"num_threads": 64}',
]
with envs.SGLANG_SPEC_NAN_DETECTION.override(
True
), envs.SGLANG_SPEC_OOB_DETECTION.override(True):
with (
envs.SGLANG_SPEC_NAN_DETECTION.override(True),
envs.SGLANG_SPEC_OOB_DETECTION.override(True),
):
cls.process = popen_launch_server(
cls.model,
cls.base_url,
@@ -59,12 +59,10 @@ class TestEagleConstrainedDecoding(
cls.grammar_backend,
]
launch_args.extend(cls.other_launch_args)
with envs.SGLANG_ENABLE_SPEC_V2.override(
cls.spec_v2
), envs.SGLANG_SPEC_NAN_DETECTION.override(
True
), envs.SGLANG_SPEC_OOB_DETECTION.override(
True
with (
envs.SGLANG_ENABLE_SPEC_V2.override(cls.spec_v2),
envs.SGLANG_SPEC_NAN_DETECTION.override(True),
envs.SGLANG_SPEC_OOB_DETECTION.override(True),
):
cls.process = popen_launch_server(
cls.model,
@@ -57,9 +57,10 @@ class TestEAGLE3EngineDPAttention(CustomTestCase):
"--cuda-graph-max-bs",
"64",
]
with envs.SGLANG_SPEC_NAN_DETECTION.override(
True
), envs.SGLANG_SPEC_OOB_DETECTION.override(True):
with (
envs.SGLANG_SPEC_NAN_DETECTION.override(True),
envs.SGLANG_SPEC_OOB_DETECTION.override(True),
):
cls.process = popen_launch_server(
cls.model,
cls.base_url,
@@ -63,14 +63,11 @@ class TestEagle3ServerBase(CustomTestCase, MatchedStopMixin):
*[str(i) for i in range(1, cls.max_running_requests + 1)],
]
launch_args.extend(cls.other_launch_args)
with envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(
1
), envs.SGLANG_SPEC_NAN_DETECTION.override(
True
), envs.SGLANG_SPEC_OOB_DETECTION.override(
True
), envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(
True
with (
envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(1),
envs.SGLANG_SPEC_NAN_DETECTION.override(True),
envs.SGLANG_SPEC_OOB_DETECTION.override(True),
envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(True),
):
cls.process = popen_launch_server(
cls.model,
@@ -65,9 +65,10 @@ class TestEagleDPAttnServerSmall(CustomTestCase):
"--speculative-num-draft-tokens",
"4",
]
with envs.SGLANG_SPEC_NAN_DETECTION.override(
True
), envs.SGLANG_SPEC_OOB_DETECTION.override(True):
with (
envs.SGLANG_SPEC_NAN_DETECTION.override(True),
envs.SGLANG_SPEC_OOB_DETECTION.override(True),
):
cls.process = popen_launch_server(
cls.model,
cls.base_url,
@@ -73,9 +73,10 @@ class TestEagleDPAttnServerLarge(CustomTestCase):
"--model-loader-extra-config",
'{"enable_multithread_load": true,"num_threads": 64}',
]
with envs.SGLANG_SPEC_NAN_DETECTION.override(
True
), envs.SGLANG_SPEC_OOB_DETECTION.override(True):
with (
envs.SGLANG_SPEC_NAN_DETECTION.override(True),
envs.SGLANG_SPEC_OOB_DETECTION.override(True),
):
cls.process = popen_launch_server(
cls.model,
cls.base_url,
@@ -51,9 +51,10 @@ class ServerWithGrammar(CustomTestCase):
"--speculative-num-draft-tokens=8",
]
with envs.SGLANG_SPEC_NAN_DETECTION.override(
True
), envs.SGLANG_SPEC_OOB_DETECTION.override(True):
with (
envs.SGLANG_SPEC_NAN_DETECTION.override(True),
envs.SGLANG_SPEC_OOB_DETECTION.override(True),
):
cls.process = popen_launch_server(
cls.model,
cls.base_url,
@@ -119,9 +119,12 @@ class ServingChatTestCase(unittest.TestCase):
# ------------- conversion tests -------------
def test_convert_to_internal_request_single(self):
with patch(
"sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv"
) as conv_mock, patch.object(self.chat, "_process_messages") as proc_mock:
with (
patch(
"sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv"
) as conv_mock,
patch.object(self.chat, "_process_messages") as proc_mock,
):
conv_ins = Mock()
conv_ins.get_prompt.return_value = "Test prompt"
conv_ins.image_data = conv_ins.audio_data = None
@@ -123,32 +123,38 @@ class TestTraceCustomIdGenerator(unittest.TestCase):
# __get_host_id
class TestGetHostId(unittest.TestCase):
def test_from_machine_id_file(self):
with patch("os.path.exists", return_value=True), patch(
"builtins.open",
unittest.mock.mock_open(read_data="abc123\n"),
with (
patch("os.path.exists", return_value=True),
patch(
"builtins.open",
unittest.mock.mock_open(read_data="abc123\n"),
),
):
self.assertEqual(_get_host_id(), "abc123")
def test_from_machine_id_file_error(self):
"""Falls back to MAC address when file read fails."""
with patch("os.path.exists", return_value=True), patch(
"builtins.open", side_effect=IOError("read error")
with (
patch("os.path.exists", return_value=True),
patch("builtins.open", side_effect=IOError("read error")),
):
result = _get_host_id()
self.assertIsInstance(result, str)
self.assertGreater(len(result), 0)
def test_from_mac_address(self):
with patch("os.path.exists", return_value=False), patch(
"uuid.getnode", return_value=0x112233445566
with (
patch("os.path.exists", return_value=False),
patch("uuid.getnode", return_value=0x112233445566),
):
result = _get_host_id()
self.assertIsInstance(result, str)
self.assertGreater(len(result), 0)
def test_unknown_fallback(self):
with patch("os.path.exists", return_value=False), patch(
"uuid.getnode", return_value=0
with (
patch("os.path.exists", return_value=False),
patch("uuid.getnode", return_value=0),
):
self.assertEqual(_get_host_id(), "unknown")
@@ -50,11 +50,14 @@ class TestGetVersionTag(unittest.TestCase):
)
def test_exact_version_tag_takes_precedence_over_latest_tag(self):
with patch.object(
self.version_helper, "get_exact_version_tag", return_value="v0.5.9"
), patch.object(
self.version_helper, "get_latest_version_tag_describe"
) as latest_describe:
with (
patch.object(
self.version_helper, "get_exact_version_tag", return_value="v0.5.9"
),
patch.object(
self.version_helper, "get_latest_version_tag_describe"
) as latest_describe,
):
self.assertEqual(self.version_helper.get_version_describe(), "v0.5.9")
latest_describe.assert_not_called()
@@ -68,15 +71,16 @@ class TestGetVersionTag(unittest.TestCase):
self.assertIn(FALLBACK_VERSION, content)
def test_tag_only_cli_mode_remains_available_for_callers_that_need_latest_tag(self):
with patch.object(
sys, "argv", ["get_version_tag.py", "--tag-only"]
), patch.object(
self.version_helper, "get_latest_version_tag", return_value="v0.5.10"
), patch.object(
self.version_helper, "get_version_describe"
) as version_describe, patch(
"builtins.print"
) as print_mock:
with (
patch.object(sys, "argv", ["get_version_tag.py", "--tag-only"]),
patch.object(
self.version_helper, "get_latest_version_tag", return_value="v0.5.10"
),
patch.object(
self.version_helper, "get_version_describe"
) as version_describe,
patch("builtins.print") as print_mock,
):
self.version_helper.main()
version_describe.assert_not_called()
@@ -466,11 +466,14 @@ class TestCompare(_WeightCheckerTestBase):
class TestHandle(_WeightCheckerTestBase):
def test_routes_to_actions(self):
with patch.object(self.checker, "_snapshot") as m_snap, patch.object(
self.checker, "_reset_tensors"
) as m_reset, patch.object(self.checker, "_compare") as m_compare, patch.object(
self.checker, "_compute_checksum", return_value={"checksums": {}}
) as m_checksum:
with (
patch.object(self.checker, "_snapshot") as m_snap,
patch.object(self.checker, "_reset_tensors") as m_reset,
patch.object(self.checker, "_compare") as m_compare,
patch.object(
self.checker, "_compute_checksum", return_value={"checksums": {}}
) as m_checksum,
):
self.checker.handle("snapshot")
self.checker.handle("reset_tensors")
self.checker.handle("compare")