[misc] Remove unit test cases that fail the admission criteria (round 2) (#30703)

This commit is contained in:
Liangsheng Yin
2026-07-09 16:35:59 -07:00
committed by GitHub
parent b86466d54b
commit a36c873147
24 changed files with 29 additions and 770 deletions
@@ -535,23 +535,6 @@ class TestGetReadyGrammarRequests(unittest.TestCase):
req.set_finish_with_abort.assert_called_once()
self.assertIn("timed out", req.set_finish_with_abort.call_args[0][0])
def test_future_exception_creates_invalid_grammar_object(self):
"""A future that raised an exception should create InvalidGrammarObject, not crash."""
mgr = self._make_mgr()
future = Future()
future.set_exception(RuntimeError("compilation crashed"))
req = _make_req(json_schema="crash")
req.grammar = future
req.grammar_key = ("json", "crash")
mgr.grammar_queue.append(req)
result = mgr.get_ready_grammar_requests()
self.assertEqual(len(result), 1)
self.assertIsInstance(result[0].grammar, InvalidGrammarObject)
req.set_finish_with_abort.assert_called_once()
def test_ready_future_applies_request_budget_without_polluting_cache(self):
mgr = self._make_mgr()
@@ -515,41 +515,10 @@ class TestValidationEdgeCases(unittest.TestCase):
with self.assertRaises(ValidationError):
CompletionRequest(model="test-model", prompt="Hello", max_tokens=-1)
def test_model_serialization_roundtrip(self):
"""Test that models can be serialized and deserialized"""
original_request = ChatCompletionRequest(
model="test-model",
messages=[{"role": "user", "content": "Hello"}],
temperature=0.7,
max_tokens=100,
)
# Serialize to dict
data = original_request.model_dump()
# Deserialize back
restored_request = ChatCompletionRequest(**data)
self.assertEqual(restored_request.model, original_request.model)
self.assertEqual(restored_request.temperature, original_request.temperature)
self.assertEqual(restored_request.max_tokens, original_request.max_tokens)
self.assertEqual(len(restored_request.messages), len(original_request.messages))
class TestParsedResponseFieldsProtocol(unittest.TestCase):
"""Test ParsedResponseFields protocol."""
def test_parsed_response_fields_protocol(self):
"""ParsedResponseFields protocol works with isinstance."""
from sglang.srt.entrypoints.openai.protocol import ParsedResponseFields
class MockFields:
content = "hello"
tool_calls = None
reasoning_content = None
self.assertIsInstance(MockFields(), ParsedResponseFields)
if __name__ == "__main__":
unittest.main(verbosity=2)
@@ -604,77 +604,6 @@ class TestHunyuanDetectorStructureInfo(CustomTestCase):
self.assertFalse(self.detector.supports_structural_tag())
class TestHunyuanDetectorAccuracy(CustomTestCase):
"""Accuracy tests for realistic HYV3 output patterns."""
def setUp(self):
self.tools = _make_tools()
self.detector = HunyuanDetector()
def test_reference_zero_arg_inline(self):
out = (
"<tool_calls><tool_call>get_current_date<tool_sep></tool_call></tool_calls>"
)
r = self.detector.detect_and_parse(out, self.tools)
self.assertEqual(len(r.calls), 1)
self.assertEqual(r.calls[0].name, "get_current_date")
self.assertEqual(json.loads(r.calls[0].parameters), {})
self.assertEqual(r.normal_text, "")
def test_reference_zero_arg_newline(self):
out = "<tool_calls>\n<tool_call>get_current_date<tool_sep>\n</tool_call>\n</tool_calls>"
r = self.detector.detect_and_parse(out, self.tools)
self.assertEqual(len(r.calls), 1)
self.assertEqual(r.calls[0].name, "get_current_date")
def test_reference_args_same_line(self):
out = (
"<tool_calls><tool_call>get_weather<tool_sep><arg_key>city</arg_key><arg_value>Beijing"
"</arg_value><arg_key>date</arg_key><arg_value>2026-03-30</arg_value></tool_call></tool_calls>"
)
r = self.detector.detect_and_parse(out, self.tools)
self.assertEqual(len(r.calls), 1)
args = json.loads(r.calls[0].parameters)
self.assertEqual(args, {"city": "Beijing", "date": "2026-03-30"})
def test_reference_args_with_newlines(self):
out = (
"<tool_calls>\n<tool_call>get_weather<tool_sep>\n<arg_key>city</arg_key>\n<arg_value>Beijing"
"</arg_value>\n<arg_key>date</arg_key>\n<arg_value>2026-03-30</arg_value>\n</tool_call>\n</tool_calls>"
)
r = self.detector.detect_and_parse(out, self.tools)
self.assertEqual(len(r.calls), 1)
args = json.loads(r.calls[0].parameters)
self.assertEqual(args, {"city": "Beijing", "date": "2026-03-30"})
def test_reference_content_before(self):
out = "Checking.<tool_calls>\n<tool_call>get_current_date<tool_sep>\n</tool_call>\n</tool_calls>"
r = self.detector.detect_and_parse(out, self.tools)
self.assertEqual(len(r.calls), 1)
self.assertEqual(r.normal_text, "Checking.")
def test_reference_multiple(self):
out = (
"<tool_calls>\n<tool_call>get_weather<tool_sep>\n<arg_key>city</arg_key>\n<arg_value>Beijing"
"</arg_value>\n<arg_key>date</arg_key>\n<arg_value>2026-03-30</arg_value>\n</tool_call>\n"
"<tool_call>get_weather<tool_sep>\n<arg_key>city</arg_key>\n<arg_value>Hangzhou</arg_value>\n"
"<arg_key>date</arg_key>\n<arg_value>2026-03-30</arg_value>\n</tool_call>\n</tool_calls>"
)
r = self.detector.detect_and_parse(out, self.tools)
self.assertEqual(len(r.calls), 2)
def test_reference_empty_content_none(self):
out = "<tool_calls>\n<tool_call>get_current_date<tool_sep>\n</tool_call>\n</tool_calls>"
r = self.detector.detect_and_parse(out, self.tools)
self.assertEqual(r.normal_text, "")
def test_reference_no_tool_call(self):
out = "This is a plain response."
r = self.detector.detect_and_parse(out, self.tools)
self.assertEqual(len(r.calls), 0)
self.assertEqual(r.normal_text, out)
class TestHunyuanDetectorFunctionCallParser(CustomTestCase):
"""Test through the FunctionCallParser interface."""
@@ -29,7 +29,6 @@ from sglang.srt.lora.mem_pool import (
LoRAMemoryPool,
_get_moe_ep_context,
_get_moe_tp_context,
_moe_runner_keeps_global_expert_ids,
)
@@ -134,16 +133,6 @@ class TestNumExpertHelpers(unittest.TestCase):
class TestGlobalToLocalExpertId(unittest.TestCase):
"""`_global_to_local_expert_id` — the per-rank filter + remap."""
def test_passthrough_without_ep(self):
pool = _make_pool(
num_experts_global=8,
moe_ep_size=1,
moe_ep_rank=0,
moe_use_local_expert_ids=False,
)
for gid in range(8):
self.assertEqual(pool._global_to_local_expert_id(gid), gid)
def test_rank0_of_ep4_owns_first_quarter(self):
pool = _make_pool(
num_experts_global=8,
@@ -368,10 +357,6 @@ class TestModuleLevelHelpers(unittest.TestCase):
self.assertEqual(tp_size, 1)
self.assertEqual(tp_rank, 0)
def test_keeps_global_expert_ids_defaults_to_false(self):
# Without a specific flashinfer backend selected, default is False.
self.assertFalse(_moe_runner_keeps_global_expert_ids())
class TestPoolInitPicksUpEpContext(unittest.TestCase):
"""`LoRAMemoryPool.__init__` should read EP context from the module-
@@ -359,33 +359,6 @@ class TestGenerateReqInputNormalization(CustomTestCase):
with self.assertRaises(ValueError):
req.normalize_batch_and_arguments()
def test_input_embeds_single_to_batch_conversion(self):
"""Test that single input_embeds are properly converted to batch when using parallel sampling."""
# Test the specific case that was fixed: single input_embeds with n > 1
req = GenerateReqInput(
input_embeds=[[0.1, 0.2, 0.3]], sampling_params={"n": 2} # Single embedding
)
req.normalize_batch_and_arguments()
# Should convert single to batch and then expand
self.assertFalse(req.is_single)
self.assertEqual(len(req.input_embeds), 2)
# Both should be the same single embedding
self.assertEqual(req.input_embeds[0], [[0.1, 0.2, 0.3]])
self.assertEqual(req.input_embeds[1], [[0.1, 0.2, 0.3]])
# Test with higher n value
req = GenerateReqInput(input_embeds=[[0.1, 0.2, 0.3]], sampling_params={"n": 5})
req.normalize_batch_and_arguments()
self.assertFalse(req.is_single)
self.assertEqual(len(req.input_embeds), 5)
# All should be the same
for i in range(5):
self.assertEqual(req.input_embeds[i], [[0.1, 0.2, 0.3]])
def test_lora_path_normalization(self):
"""Test normalization of lora_path."""
# Test single lora_path with batch input
@@ -646,23 +619,6 @@ class TestGenerateReqInputNormalization(CustomTestCase):
)
req.normalize_batch_and_arguments()
def test_multiple_input_formats(self):
"""Test different combinations of input formats."""
# Test with text only
req = GenerateReqInput(text="Hello")
req.normalize_batch_and_arguments()
self.assertTrue(req.is_single)
# Test with input_ids only
req = GenerateReqInput(input_ids=[1, 2, 3])
req.normalize_batch_and_arguments()
self.assertTrue(req.is_single)
# Test with input_embeds only
req = GenerateReqInput(input_embeds=[[0.1, 0.2]])
req.normalize_batch_and_arguments()
self.assertTrue(req.is_single)
if __name__ == "__main__":
unittest.main()
@@ -266,14 +266,6 @@ class TestFactoryFunctions(CustomTestCase):
args = SimpleNamespace(enable_dp_attention=True, nnodes=2)
self.assertTrue(should_use_zmq(args))
def test_should_use_zmq_single_node(self):
args = SimpleNamespace(enable_dp_attention=False, nnodes=1)
self.assertFalse(should_use_zmq(args))
def test_should_use_zmq_dp_attention_single_node(self):
args = SimpleNamespace(enable_dp_attention=True, nnodes=1)
self.assertFalse(should_use_zmq(args))
class TestZmqReaderOwner(CustomTestCase):
"""At most one process binds the zmq PULL socket across all callers."""
@@ -136,18 +136,6 @@ class TestEntryStateAndPins(unittest.TestCase):
self.assertTrue(entry.is_evictable())
def test_filling_entry_is_not_evictable(self):
entry = EmbeddingCacheEntry(
hash="h",
modality=Modality.IMAGE,
num_tokens=2,
dim=4,
page_runs=[PageRun(0, 1)],
state=EntryState.FILLING,
)
self.assertFalse(entry.is_evictable())
def test_ready_entry_with_pin_is_not_evictable(self):
entry = EmbeddingCacheEntry(
hash="h",
@@ -19,7 +19,6 @@ register_cpu_ci(est_time=10, suite="base-a-test-cpu")
import os
import shutil
import tempfile
import threading
import time
import unittest
from unittest import mock
@@ -401,32 +400,6 @@ class TestMinFreeSpaceWatermark(HiCacheFileLRUTestBase):
class TestPreReservationConcurrency(HiCacheFileLRUTestBase):
def test_concurrent_sets_keep_total_consistent_with_lru(self):
"""Under concurrent writes, _total_bytes stays consistent with _lru."""
b = self.make_backend(max_size="300", eviction_ratio=1.0)
n_threads = 8
per_size = 60
errors = []
def writer(i):
try:
b.set(f"k{i}", _t(per_size, fill=i % 256))
except Exception as e:
errors.append(e)
threads = [threading.Thread(target=writer, args=(i,)) for i in range(n_threads)]
for t in threads:
t.start()
for t in threads:
t.join()
self.assertEqual(errors, [])
# Invariant 1: _total_bytes equals the sum of tracked LRU sizes.
tracked_sum = sum(b._evictor._lru.values())
self.assertEqual(b._evictor._total_bytes, tracked_sum)
# Invariant 2: _total_bytes does not exceed the cap.
self.assertLessEqual(b._evictor._total_bytes, 300)
def test_pre_reservation_visible_during_write(self):
"""An in-flight reservation must not be evicted by a concurrent set()."""
b = self.make_backend(max_size="100", eviction_ratio=1.0)
@@ -197,33 +197,8 @@ class TestStoreCache4D(unittest.TestCase):
# ---- Test 4: int64 loc dtype (already exercised, explicit) ----
def test_store_cache_4d_int64_loc(self):
"""The full-side path passes int64 loc (matches the v2p table
dtype)."""
self._check_parity(
num_pages=32,
page_size=1,
head_num=4,
head_dim=64,
v_head_dim=64,
N=10,
loc_dtype=torch.int64,
)
# ---- Test 5: bf16 dtype (the production case) ----
def test_store_cache_4d_dtype_bf16(self):
"""bf16 is the production K/V dtype for gpt-oss-20b, Falcon-H1."""
self._check_parity(
num_pages=16,
page_size=64,
head_num=4,
head_dim=128,
v_head_dim=128,
N=64,
dtype=torch.bfloat16,
)
# ---- Test 6: fp8_e5m2 dtype ----
def test_store_cache_4d_dtype_fp8_e5m2(self):
@@ -433,37 +433,6 @@ class TestSWALockReleaseLifecycle(CustomTestCase):
)
tree.sanity_check()
def test_full_lifecycle_inc_dec_swa_dec_lock_balances(self):
tree, allocator, _ = _build_tree(sliding_window_size=4)
leaf = _insert_chain(tree, allocator, [1, 2, 3, 4, 5, 6, 7, 8])
full_protected0 = tree.full_protected_size_
swa_protected0 = tree.swa_protected_size_
full_avail0 = allocator.full_available_size()
swa_avail0 = allocator.swa_available_size()
inc_res = tree.inc_lock_ref(leaf)
swa_uuid = inc_res.swa_uuid_for_lock
self.assertGreater(tree.full_protected_size_, full_protected0)
self.assertGreater(tree.swa_protected_size_, swa_protected0)
tree.dec_swa_lock_only(leaf, swa_uuid_for_lock=swa_uuid)
self.assertEqual(tree.swa_protected_size_, swa_protected0)
self.assertGreater(tree.full_protected_size_, full_protected0)
tree.dec_lock_ref(
leaf, DecLockRefParams(swa_uuid_for_lock=swa_uuid), skip_swa=True
)
self.assertEqual(tree.full_protected_size_, full_protected0)
self.assertEqual(tree.swa_protected_size_, swa_protected0)
self.assertEqual(allocator.full_available_size(), full_avail0)
self.assertEqual(allocator.swa_available_size(), swa_avail0 + len(leaf.value))
tree.sanity_check()
if __name__ == "__main__":
unittest.main()
@@ -201,53 +201,6 @@ class TestSWA(unittest.TestCase):
self.assertEqual(list(second_insert_events[0].token_ids), [5])
self.assertEqual(second_insert_events[0].parent_block_hash, split_parent_hash)
def test_swa_memory_pool(self):
size = 16
size_swa = 16
page_size = 1
head_num = 8
head_dim = 128
num_layers = 48
global_interval = 4
dtype = torch.bfloat16
device = get_device()
full_attention_layer_ids = [i for i in range(0, num_layers, global_interval)]
full_attention_layer_ids_set = set(full_attention_layer_ids)
swa_attention_layer_ids = [
i for i in range(num_layers) if i not in full_attention_layer_ids_set
]
pool = SWAKVPool(
size=size,
size_swa=size_swa,
page_size=page_size,
dtype=dtype,
head_num=head_num,
head_dim=head_dim,
swa_attention_layer_ids=swa_attention_layer_ids,
full_attention_layer_ids=full_attention_layer_ids,
device=device,
)
alloc = SWATokenToKVPoolAllocator(
size=size,
size_swa=size_swa,
page_size=page_size,
dtype=dtype,
device=device,
kvcache=pool,
need_sort=False,
)
self.assertEqual(
alloc.full_available_size() + alloc.swa_available_size(), size + size_swa
)
index = alloc.alloc(1)
self.assertEqual(
alloc.full_available_size() + alloc.swa_available_size(),
size_swa + size_swa - 2,
)
alloc.free_swa(index)
result = alloc.translate_loc_from_full_to_swa(index)
print(result)
def test_swa_memory_pool_paged_free_clears_full_page_mapping(self):
page_size = 4
_, allocator, _ = _build_swa_tree(
@@ -204,24 +204,6 @@ class TestModelOptExport(unittest.TestCase):
)
mock_export.assert_called_once_with(self.mock_model, self.export_dir, None)
@unittest.skipIf(not MODELOPT_AVAILABLE, "nvidia-modelopt not available")
def test_setup_quantization_without_export(self):
"""Test quantization setup without export path specified."""
with patch("modelopt.torch.quantization.utils.is_quantized", return_value=True):
# Act
with patch.object(
self.model_loader, "_export_modelopt_checkpoint"
) as mock_export:
self.model_loader._setup_modelopt_quantization(
self.mock_model,
self.mock_tokenizer,
self.mock_quant_cfg,
export_path=None, # No export path
)
# Assert
mock_export.assert_not_called()
def test_quantize_and_serve_config_validation(self):
"""Test that quantize_and_serve is properly disabled."""
# Test that quantize-and-serve mode raises NotImplementedError
@@ -274,25 +256,6 @@ class TestModelOptExport(unittest.TestCase):
# Assert
mock_standard.assert_called_once_with(model_config, device_config)
def _get_export_info(self, export_dir: str) -> dict:
"""Get information about an exported model."""
if not self._validate_export(export_dir):
return None
try:
config_path = os.path.join(export_dir, "config.json")
with open(config_path, "r") as f:
config = json.load(f)
return {
"model_type": config.get("model_type", "unknown"),
"architectures": config.get("architectures", []),
"quantization_config": config.get("quantization_config", {}),
"export_dir": export_dir,
}
except Exception:
return None
@unittest.skipIf(not MODELOPT_AVAILABLE, "nvidia-modelopt not available")
class TestModelOptExportIntegration(unittest.TestCase):
@@ -101,88 +101,6 @@ class TestModelOptModelLoader(CustomTestCase):
self.mock_get_tp_group.stop()
self.mock_mp_is_initialized.stop()
@patch("sglang.srt.model_loader.loader.QUANT_CFG_CHOICES", QUANT_CFG_CHOICES)
@patch("sglang.srt.model_loader.loader.logger")
def test_successful_fp8_quantization(self, mock_logger):
"""Test successful FP8 quantization workflow."""
# Create loader instance
loader = ModelOptModelLoader(self.load_config)
# Mock modelopt modules
mock_mtq = MagicMock()
# Configure mtq mock with FP8_DEFAULT_CFG
mock_fp8_cfg = MagicMock()
mock_mtq.FP8_DEFAULT_CFG = mock_fp8_cfg
mock_mtq.quantize.return_value = self.mock_base_model
mock_mtq.print_quant_summary = MagicMock()
# Create a custom load_model method for testing that simulates the real logic
def mock_load_model(*, model_config, device_config):
mock_logger.info("ModelOptModelLoader: Loading base model...")
# Simulate loading base model (this is already mocked)
model = self.mock_base_model
# Simulate the quantization config lookup
quant_choice_str = model_config._get_modelopt_quant_type()
quant_cfg_name = QUANT_CFG_CHOICES.get(quant_choice_str)
if not quant_cfg_name:
raise ValueError(f"Invalid modelopt_quant choice: '{quant_choice_str}'")
# Simulate getattr call and quantization
if quant_cfg_name == "FP8_DEFAULT_CFG":
quant_cfg = mock_fp8_cfg
mock_logger.info(
f"Quantizing model with ModelOpt using config attribute: mtq.{quant_cfg_name}"
)
# Simulate mtq.quantize call
quantized_model = mock_mtq.quantize(model, quant_cfg, forward_loop=None)
mock_logger.info("Model successfully quantized with ModelOpt.")
# Simulate print_quant_summary call
mock_mtq.print_quant_summary(quantized_model)
return quantized_model.eval()
return model.eval()
# Patch the load_model method with our custom implementation
with patch.object(loader, "load_model", side_effect=mock_load_model):
# Execute the load_model method
result_model = loader.load_model(
model_config=self.model_config, device_config=self.device_config
)
# Verify the quantization process
mock_mtq.quantize.assert_called_once_with(
self.mock_base_model, mock_fp8_cfg, forward_loop=None
)
# Verify logging
mock_logger.info.assert_any_call(
"ModelOptModelLoader: Loading base model..."
)
mock_logger.info.assert_any_call(
"Quantizing model with ModelOpt using config attribute: mtq.FP8_DEFAULT_CFG"
)
mock_logger.info.assert_any_call(
"Model successfully quantized with ModelOpt."
)
# Verify print_quant_summary was called
mock_mtq.print_quant_summary.assert_called_once_with(self.mock_base_model)
# Verify eval() was called on the returned model
self.mock_base_model.eval.assert_called()
# Verify we get back the expected model
self.assertEqual(result_model, self.mock_base_model)
@patch("sglang.srt.model_loader.loader.logger")
def test_missing_modelopt_import(self, mock_logger):
"""Test error handling when modelopt library is not available."""
@@ -486,49 +404,6 @@ class TestModelOptModelLoader(CustomTestCase):
class TestModelOptLoaderIntegration(CustomTestCase):
"""Integration tests for ModelOptModelLoader with Engine API."""
@patch("sglang.srt.model_loader.loader.get_model_loader")
@patch("sglang.srt.entrypoints.engine.Engine.__init__")
def test_engine_with_modelopt_quant_parameter(
self, mock_engine_init, mock_get_model_loader
):
"""Test that Engine properly handles modelopt_quant parameter."""
# Mock the Engine.__init__ to avoid actual initialization
mock_engine_init.return_value = None
# Mock get_model_loader to return our ModelOptModelLoader
mock_loader = MagicMock(spec=ModelOptModelLoader)
mock_get_model_loader.return_value = mock_loader
# Import here to avoid circular imports during test discovery
# import sglang as sgl # Commented out since not directly used
# Test that we can create an engine with modelopt_quant parameter
# This would normally trigger the ModelOptModelLoader selection
try:
engine_args = {
"model_path": "TinyLlama/TinyLlama-1.1B-Chat-v1.0",
"modelopt_quant": "fp8",
"log_level": "error", # Suppress logs during testing
}
# This tests the parameter parsing and server args creation
from sglang.srt.server_args import ServerArgs
server_args = ServerArgs(**engine_args)
# Verify that modelopt_quant is properly set
self.assertEqual(server_args.modelopt_quant, "fp8")
except Exception as e:
# If there are missing dependencies or initialization issues,
# we can still verify the parameter is accepted
if "modelopt_quant" not in str(e):
# The parameter was accepted, which is what we want to test
pass
else:
self.fail(f"modelopt_quant parameter not properly handled: {e}")
@patch("sglang.srt.model_loader.loader.get_model_loader")
@patch("sglang.srt.entrypoints.engine.Engine.__init__")
def test_engine_with_modelopt_quant_cli_argument(
@@ -743,12 +618,6 @@ class TestModelOptMixedPrecisionConfig(CustomTestCase):
)
)
def test_mixed_precision_uses_nvfp4_min_capability(self):
self.assertEqual(
ModelOptMixedPrecisionConfig.get_min_capability(),
ModelOptFp4Config.get_min_capability(),
)
def test_mixed_precision_quant_layer_resolution_after_mapping(self):
quant_config = ModelOptMixedPrecisionConfig.from_config(
{
@@ -250,10 +250,6 @@ class TestFileRequestMetricsExporter(unittest.TestCase):
self.assertIsNone(exporter._current_file_handler)
self.assertIsNone(exporter._current_hour_suffix)
def test_close_noop_when_no_handler(self):
exporter = self._make_exporter()
exporter.close() # should not raise
def test_close_error(self):
"""Close failure is logged but state is still reset."""
exporter = self._make_exporter()
@@ -13,12 +13,10 @@ from unittest.mock import patch
import sglang.srt.observability.trace as mod
from sglang.srt.observability.trace import (
SpanAttributes,
TraceCustomIdGenerator,
TraceEvent,
TraceNullContext,
TraceReqContext,
TraceSliceContext,
TraceThreadContext,
TraceThreadInfo,
extract_trace_headers,
get_global_trace_level,
@@ -85,25 +83,6 @@ class TestTraceFunctions(unittest.TestCase):
self.assertGreater(ts, 0)
class TestDataclasses(unittest.TestCase):
def test_trace_thread_info(self):
info = TraceThreadInfo("host", 123, "label", 0, 1, 0)
self.assertEqual(info.thread_label, "label")
def test_trace_event(self):
evt = TraceEvent("name", 100, {"k": "v"})
self.assertEqual(evt.event_name, "name")
def test_trace_slice_context(self):
s = TraceSliceContext("slice", 100, end_time_ns=200, level=2, attrs={"a": 1})
self.assertEqual(s.slice_name, "slice")
def test_trace_thread_context(self):
info = TraceThreadInfo("h", 1, "l", 0, 0, 0)
ctx = TraceThreadContext(thread_info=info, cur_slice_stack=[])
self.assertEqual(len(ctx.cur_slice_stack), 0)
class TestTraceNullContext(unittest.TestCase):
def test_null_object_pattern(self):
ctx = TraceNullContext()
@@ -122,15 +101,6 @@ class TestSpanAttributes(unittest.TestCase):
self.assertIsInstance(SpanAttributes.GEN_AI_USAGE_COMPLETION_TOKENS, str)
class TestTraceCustomIdGenerator(unittest.TestCase):
def test_generates_nonzero_ids(self):
gen = TraceCustomIdGenerator()
trace_id = gen.generate_trace_id()
span_id = gen.generate_span_id()
self.assertIsInstance(trace_id, int)
self.assertIsInstance(span_id, int)
# __get_host_id
class TestGetHostId(unittest.TestCase):
def test_from_machine_id_file(self):
@@ -219,19 +189,6 @@ class TestTraceReqContextDisabled(unittest.TestCase):
self.assertFalse(ctx.tracing_enable)
self.assertFalse(ctx.is_tracing_enabled())
def test_all_methods_noop(self):
ctx = TraceReqContext(rid="req-1")
ctx.trace_req_start()
ctx.trace_req_finish()
ctx.trace_slice_start("s", 1)
ctx.trace_slice_end("s", 1)
ctx.trace_slice(TraceSliceContext("s", 100))
ctx.trace_event("e", 1)
ctx.trace_set_root_attrs({"k": "v"})
ctx.trace_set_thread_attrs({"k": "v"})
ctx.abort()
ctx.rebuild_thread_context()
def test_getstate_disabled(self):
ctx = TraceReqContext(rid="req-1")
state = ctx.__getstate__()
@@ -243,8 +200,6 @@ class TestTraceReqContextDisabled(unittest.TestCase):
# opentelemetry_initialized is False → tracing forced off
self.assertFalse(ctx.tracing_enable)
def test_trace_set_thread_info_disabled(self):
trace_set_thread_info("test_label")
# Should not register anything
@@ -332,13 +287,6 @@ class TestTraceReqContextEnabled(unittest.TestCase):
self.assertIsNotNone(ctx.root_span)
ctx.trace_req_finish(ts=2000)
def test_trace_req_finish_without_start(self):
"""finish without start is a no-op."""
ctx = TraceReqContext(rid="req-1")
ctx.trace_req_start(ts=1000)
ctx.root_span = None
ctx.trace_req_finish(ts=2000)
def test_trace_slice_combined(self):
"""trace_slice() creates and ends a span in one call."""
ctx = TraceReqContext(rid="req-1")
@@ -456,24 +404,6 @@ class TestTraceReqContextEnabled(unittest.TestCase):
ctx.trace_req_finish(ts=5000)
def test_trace_set_root_attrs(self):
ctx = TraceReqContext(rid="req-1")
ctx.trace_req_start(ts=1000)
ctx.trace_set_root_attrs({"model": "llama"})
ctx.trace_req_finish(ts=2000)
def test_trace_set_root_attrs_no_span(self):
ctx = TraceReqContext(rid="req-1")
ctx.trace_req_start(ts=1000)
ctx.root_span = None
ctx.trace_set_root_attrs({"model": "llama"}) # no crash
def test_trace_set_thread_attrs(self):
ctx = TraceReqContext(rid="req-1")
ctx.trace_req_start(ts=1000)
ctx.trace_set_thread_attrs({"batch_size": 32})
ctx.trace_req_finish(ts=2000)
def test_abort_with_unclosed_slices(self):
ctx = TraceReqContext(rid="req-1")
ctx.trace_req_start(ts=1000)
@@ -22,27 +22,6 @@ register_cpu_ci(est_time=7, suite="base-a-test-cpu")
register_cpu_ci(est_time=7, suite="base-c-test-cpu")
class TestFimPosition(CustomTestCase):
def test_middle_and_end_are_distinct(self):
"""Test that MIDDLE and END are different enum values."""
self.assertNotEqual(FimPosition.MIDDLE, FimPosition.END)
class TestCompletionTemplate(CustomTestCase):
def test_dataclass_fields(self):
"""Test creating a CompletionTemplate with all fields."""
t = CompletionTemplate(
name="test",
fim_begin_token="<begin>",
fim_middle_token="<middle>",
fim_end_token="<end>",
fim_position=FimPosition.MIDDLE,
)
self.assertEqual(t.name, "test")
self.assertEqual(t.fim_begin_token, "<begin>")
self.assertEqual(t.fim_position, FimPosition.MIDDLE)
class TestRegisterCompletionTemplate(CustomTestCase):
def test_builtin_templates_registered(self):
"""Test that deepseek_coder, star_coder, qwen_coder are pre-registered."""
@@ -102,46 +102,6 @@ class TestTemplateContentFormatDetection(CustomTestCase):
result = detect_jinja_template_content_format(msg_content_pattern)
self.assertEqual(result, "openai")
def test_detect_m_content_pattern(self):
"""Test detection of template with m.content pattern (should be 'openai' format)."""
msg_content_pattern = """
[gMASK]<sop>
{%- for m in messages %}
{%- if m.role == 'system' %}
<|system|>
{{ m.content }}
{%- elif m.role == 'user' %}
<|user|>{{ '\n' }}
{%- if m.content is string %}
{{ m.content }}
{%- else %}
{%- for item in m.content %}
{%- if item.type == 'video' or 'video' in item %}
<|begin_of_video|><|video|><|end_of_video|>
{%- elif item.type == 'image' or 'image' in item %}
<|begin_of_image|><|image|><|end_of_image|>
{%- elif item.type == 'text' %}
{{ item.text }}
{%- endif %}
{%- endfor %}
{%- endif %}
{%- elif m.role == 'assistant' %}
{%- if m.metadata %}
<|assistant|>{{ m.metadata }}
{{ m.content }}
{%- else %}
<|assistant|>
{{ m.content }}
{%- endif %}
{%- endif %}
{%- endfor %}
{% if add_generation_prompt %}<|assistant|>
{% endif %}
"""
result = detect_jinja_template_content_format(msg_content_pattern)
self.assertEqual(result, "openai")
def test_process_content_openai_format(self):
"""Test content processing for openai format."""
msg_dict = {
@@ -583,19 +583,6 @@ class TestToolCallParserDetection(unittest.TestCase):
self.assertLess(minicpm5_idx, rule_names.index("mimo"))
self.assertLess(minicpm5_idx, rule_names.index("qwen"))
def test_minicpm5_not_misclassified_as_qwen(self):
template = (
"{% set enable_thinking = enable_thinking if enable_thinking is defined else true %}"
'\n<function name="{{ tool.name }}">'
'\n<param name="{{ param.name }}">{{ param.value }}</param>'
"\n</function>"
)
force, config = detect_reasoning_pattern(template)
result = detect_tool_call_parser(
template, _DummyTokenizer(["<function", "<param"]), config, force
)
self.assertEqual(result, "minicpm5")
class TestResolveAutoParsers(unittest.TestCase):
"""Tests for resolve_auto_parsers()."""
@@ -62,13 +62,6 @@ class TestHfStoreConfig(CustomTestCase):
cfg = hfs.HfStoreConfig.from_env()
self.assertEqual(cfg.revision, "dev")
def test_from_env_default_revision(self):
with patch.dict(
os.environ, {"SGLANG_PRECISION_HF_REPO": "my/repo"}, clear=False
):
cfg = hfs.HfStoreConfig.from_env()
self.assertEqual(cfg.revision, "main")
def test_from_env_raises_when_missing(self):
with patch.dict(os.environ, {}, clear=True):
with self.assertRaises(RuntimeError):
@@ -626,19 +619,6 @@ class TestWithRetries(CustomTestCase):
self.assertEqual(result, "ok")
mock_time.sleep.assert_called_once()
@patch("sglang.test.precision_baseline_store.time")
def test_retries_on_5xx(self, mock_time):
from huggingface_hub.errors import HfHubHTTPError
resp_500 = MagicMock()
resp_500.status_code = 500
exc_500 = HfHubHTTPError("server error", response=resp_500)
mock_op = MagicMock(side_effect=[exc_500, "ok"])
result = hfs._with_retries(mock_op, what="test", base_delay=0.01)
self.assertEqual(result, "ok")
mock_time.sleep.assert_called_once()
@patch("sglang.test.precision_baseline_store.time")
def test_raises_on_auth_error(self, mock_time):
from huggingface_hub.errors import HfHubHTTPError
@@ -178,11 +178,6 @@ class TestServerArgsOwnership(_IsolatedServerArgs):
self.assertIs(get_server_args(), sentinel)
self.assertIs(get_context().server_args, sentinel)
def test_context_publish_visible_through_legacy_getter(self):
sentinel = object()
get_context().set_server_args(sentinel)
self.assertIs(server_args_module.get_global_server_args(), sentinel)
def test_tokenizer_alias_is_same_function(self):
self.assertIs(
server_args_module.set_global_server_args_for_tokenizer,
-17
View File
@@ -15,11 +15,6 @@ register_cpu_ci(1.0, "base-a-test-cpu")
class TestAuthDecision(CustomTestCase):
def test_allowed_default(self):
decision = AuthDecision(allowed=True)
self.assertTrue(decision.allowed)
self.assertEqual(decision.error_status_code, 401)
def test_not_allowed_with_custom_status(self):
decision = AuthDecision(allowed=False, error_status_code=403)
self.assertFalse(decision.allowed)
@@ -32,11 +27,6 @@ class TestAuthDecision(CustomTestCase):
class TestAuthLevel(CustomTestCase):
def test_enum_values(self):
self.assertEqual(AuthLevel.NORMAL.value, "normal")
self.assertEqual(AuthLevel.ADMIN_OPTIONAL.value, "admin_optional")
self.assertEqual(AuthLevel.ADMIN_FORCE.value, "admin_force")
def test_is_string_enum(self):
self.assertIsInstance(AuthLevel.NORMAL, str)
# str mixin allows direct comparison with string values
@@ -51,13 +41,6 @@ class TestAuthLevelDecorator(CustomTestCase):
self.assertEqual(my_endpoint._auth_level, AuthLevel.ADMIN_FORCE)
def test_decorator_preserves_function(self):
@auth_level(AuthLevel.NORMAL)
def my_endpoint():
return 42
self.assertEqual(my_endpoint(), 42)
class TestDecideRequestAuth(CustomTestCase):
"""Tests for the pure decide_request_auth function."""
@@ -45,11 +45,6 @@ class TestNormalizeRopeScalingCompat(unittest.TestCase):
normalize_rope_scaling_compat(cfg)
self.assertEqual(cfg.rope_scaling["type"], "custom")
def test_no_op_when_no_rope_scaling(self):
cfg = PretrainedConfig()
normalize_rope_scaling_compat(cfg)
self.assertIsNone(getattr(cfg, "rope_scaling", None))
def test_no_op_when_rope_scaling_is_none(self):
cfg = PretrainedConfig()
cfg.rope_scaling = None
@@ -479,59 +474,6 @@ class TestPatchRemovedSymbols(unittest.TestCase):
# ---------------------------------------------------------------------------
class TestPatchRopeParametersValidation(unittest.TestCase):
# -----------------------------------------------------------------------
# Test ``rope_theta`` injection into ``rope_scaling``.
#
# Upstream `transformers.PretrainedConfig` now natively handles this
# logic. While the manual injection patch has been removed, these
# test cases are retained to ensure regression testing of the
# configuration's injection behavior.
# -----------------------------------------------------------------------
def test_injects_rope_theta_into_rope_scaling(self):
config_dict = {
"model_type": "llama",
"rope_theta": 500000.0,
"max_position_embeddings": 131072,
"rope_scaling": {
"rope_type": "llama3",
"factor": 8.0,
"low_freq_factor": 1.0,
"high_freq_factor": 4.0,
"original_max_position_embeddings": 8192,
},
}
config = PretrainedConfig.from_dict(config_dict)
rope_params = getattr(config, "rope_parameters", None)
if rope_params is not None:
self.assertIn("rope_theta", rope_params)
def test_no_injection_when_rope_theta_already_in_scaling(self):
config_dict = {
"model_type": "llama",
"rope_theta": 500000.0,
"max_position_embeddings": 131072,
"rope_scaling": {
"rope_type": "llama3",
"factor": 8.0,
"rope_theta": 999.0,
"low_freq_factor": 1.0,
"high_freq_factor": 4.0,
"original_max_position_embeddings": 8192,
},
}
config = PretrainedConfig.from_dict(config_dict)
rope_params = getattr(config, "rope_parameters", None)
if rope_params is not None:
self.assertEqual(rope_params["rope_theta"], 999.0)
def test_no_crash_without_rope_scaling(self):
config_dict = {"model_type": "llama", "rope_theta": 10000.0}
config = PretrainedConfig.from_dict(config_dict)
self.assertIsNotNone(config)
# ---------------------------------------------------------------------------
# compat: _ensure_clean_up_tokenization_compat
# ---------------------------------------------------------------------------
@@ -217,32 +217,6 @@ class TestProfileMergerIntegration(CustomTestCase):
req = ProfileReq(merge_profiles=True)
self.assertTrue(req.merge_profiles)
def test_integration_parameters(self):
import inspect
# Test TokenizerManager
from sglang.srt.managers.tokenizer_control_mixin import (
TokenizerControlMixin,
)
sig = inspect.signature(TokenizerControlMixin.start_profile)
self.assertIn("req", sig.parameters)
self.assertNotIn("merge_profiles", sig.parameters)
# Test SchedulerProfilerMixin
from sglang.srt.managers.scheduler_components.profiler_manager import (
SchedulerProfilerManager,
)
sig = inspect.signature(SchedulerProfilerManager._init_profile)
self.assertIn("merge_profiles", sig.parameters)
# Test CLI profiler
from sglang.profiler import run_profile
sig = inspect.signature(run_profile)
self.assertIn("merge_profiles", sig.parameters)
class TestProfileMergerEdgeCases(CustomTestCase):
def setUp(self):