Each of these wrote a value after resolution so a later reader would find it on the instance. None of them needed the instance: one write was redundant, and the two that carry a value the resolved-config readback reports move to get_context().override, which the readback overlays. - The SM100 GDN prefill default was written onto ServerArgs and read back one line later by initialize_linear_attn_config. It is now the return value of flashinfer_gdn_prefill_default, threaded into initialize_linear_attn_config (an explicit --linear-attn-prefill-backend still wins) and recorded with get_context().override so /server_info reports the backend in effect. - The XGrammar fallback recorded grammar_backend="none" on the instance. No code reads the field after the factory reads it once, but get_internal_state reports the whole resolved config, so the fallback now lands there instead: the readback tells the truth and the seed keeps the requested backend. - UnifiedRadixCache.init_hicache re-applied the direct-IO layout fixup that __post_init__ already applies: init_hicache only runs when hierarchical cache is on, which is exactly when _handle_hicache normalizes page_first to page_first_direct (pinned by test_hicache_io_backend_and_mem_layout_ compatibility::direct_with_page_first). Three fixtures reached the fixup by building ServerArgs(model_path="dummy"), whose resolution is skipped, so they now declare the layout resolution would have produced. Writer ratchet 34 -> 31.
445 lines
18 KiB
Python
445 lines
18 KiB
Python
"""
|
|
Unit tests for sglang.srt.constrained.base_grammar_backend.
|
|
|
|
Test Coverage:
|
|
- GrammarStats: default values, mutable default isolation
|
|
- BaseGrammarObject: default behavior
|
|
- InvalidGrammarObject: error message
|
|
- BaseGrammarBackend: caching, dispatch routing, unsupported fallback,
|
|
thread pool execution, cache hit/miss
|
|
- create_grammar_backend: factory routing, "none" backend, invalid name,
|
|
custom registry, reasoner wrapping
|
|
- register_grammar_backend: registration and lookup
|
|
|
|
Usage:
|
|
python -m pytest test_base_grammar_backend.py -v
|
|
"""
|
|
|
|
import unittest
|
|
from concurrent.futures import Future
|
|
from unittest.mock import MagicMock, patch
|
|
|
|
from sglang.srt.constrained.base_grammar_backend import (
|
|
GRAMMAR_BACKEND_REGISTRY,
|
|
BaseGrammarBackend,
|
|
BaseGrammarObject,
|
|
GrammarStats,
|
|
InvalidGrammarObject,
|
|
create_grammar_backend,
|
|
register_grammar_backend,
|
|
)
|
|
from sglang.test.ci.ci_register import register_cpu_ci
|
|
|
|
register_cpu_ci(2.0, "base-a-test-cpu")
|
|
|
|
|
|
class TestGrammarStats(unittest.TestCase):
|
|
"""Test GrammarStats dataclass."""
|
|
|
|
def test_defaults(self):
|
|
stats = GrammarStats()
|
|
self.assertIsNone(stats.compilation_time)
|
|
self.assertIsNone(stats.schema_count)
|
|
self.assertIsNone(stats.ebnf_size)
|
|
self.assertFalse(stats.is_cache_hit)
|
|
self.assertFalse(stats.is_grammar_aborted)
|
|
self.assertEqual(stats.tree_traversal_time, [])
|
|
self.assertIsNone(stats.dispatch_type)
|
|
self.assertEqual(stats.num_timeout, 0)
|
|
|
|
def test_tree_traversal_time_mutable_default(self):
|
|
"""Ensure each instance gets its own list."""
|
|
s1 = GrammarStats()
|
|
s2 = GrammarStats()
|
|
s1.tree_traversal_time.append(0.1)
|
|
self.assertEqual(len(s2.tree_traversal_time), 0)
|
|
|
|
|
|
class TestBaseGrammarObject(unittest.TestCase):
|
|
"""Test BaseGrammarObject base class."""
|
|
|
|
|
|
class TestInvalidGrammarObject(unittest.TestCase):
|
|
"""Test InvalidGrammarObject."""
|
|
|
|
def test_default_error_message(self):
|
|
obj = InvalidGrammarObject()
|
|
self.assertEqual(obj.error_message, "Unknown grammar error")
|
|
|
|
def test_custom_error_message(self):
|
|
obj = InvalidGrammarObject("Regex compilation failed")
|
|
self.assertEqual(obj.error_message, "Regex compilation failed")
|
|
|
|
|
|
class TestBaseGrammarBackend(unittest.TestCase):
|
|
"""Test BaseGrammarBackend caching and dispatch."""
|
|
|
|
def setUp(self):
|
|
self.backend = BaseGrammarBackend()
|
|
|
|
def tearDown(self):
|
|
self.backend.executor.shutdown(wait=True)
|
|
|
|
def test_cache_hit_returns_copy(self):
|
|
"""Cache hit should return a copy of the cached object."""
|
|
mock_copy = BaseGrammarObject()
|
|
obj = MagicMock(spec=BaseGrammarObject)
|
|
obj.copy.return_value = mock_copy
|
|
|
|
key = ("json", "schema")
|
|
self.backend.set_cache(key, obj)
|
|
result, cache_hit = self.backend.get_cached_or_future_value(key, False)
|
|
|
|
self.assertTrue(cache_hit)
|
|
obj.copy.assert_called_once()
|
|
self.assertIs(result, mock_copy)
|
|
|
|
def test_cache_hit_inits_reasoning(self):
|
|
obj = MagicMock(spec=BaseGrammarObject)
|
|
copied = MagicMock(spec=BaseGrammarObject)
|
|
obj.copy.return_value = copied
|
|
|
|
key = ("json", "schema")
|
|
self.backend.set_cache(key, obj)
|
|
self.backend.get_cached_or_future_value(key, True)
|
|
copied.maybe_init_reasoning.assert_called_once_with(True)
|
|
|
|
def test_cache_miss_returns_future(self):
|
|
key = ("json", "schema")
|
|
result, cache_hit = self.backend.get_cached_or_future_value(key, False)
|
|
self.assertFalse(cache_hit)
|
|
self.assertIsInstance(result, Future)
|
|
# The future should complete (dispatch_json returns InvalidGrammarObject)
|
|
value = result.result(timeout=5)
|
|
self.assertIsInstance(value, InvalidGrammarObject)
|
|
|
|
def test_dispatch_fallback_raises(self):
|
|
with self.assertRaises(ValueError):
|
|
self.backend.dispatch_fallback("unknown", "value")
|
|
|
|
def test_init_value_dispatch_routes_all_types(self):
|
|
"""_init_value_dispatch routes all grammar types to their dispatch methods."""
|
|
cases = [
|
|
("json", "schema"),
|
|
("regex", "[a-z]+"),
|
|
("ebnf", "root ::= 'x'"),
|
|
("structural_tag", "{}"),
|
|
]
|
|
for grammar_type, value in cases:
|
|
with self.subTest(grammar_type=grammar_type):
|
|
result = self.backend._init_value_dispatch((grammar_type, value), False)
|
|
self.assertIsInstance(result, InvalidGrammarObject)
|
|
|
|
def test_init_value_dispatch_unknown_type_raises(self):
|
|
with self.assertRaises(ValueError):
|
|
self.backend._init_value_dispatch(("unknown_type", "value"), False)
|
|
|
|
def test_init_value_dispatch_sets_compilation_time(self):
|
|
"""When grammar has stats, compilation_time should be set."""
|
|
mock_grammar = MagicMock(spec=BaseGrammarObject)
|
|
mock_grammar.grammar_stats = GrammarStats()
|
|
self.backend.dispatch_json = MagicMock(return_value=mock_grammar)
|
|
|
|
result = self.backend._init_value_dispatch(("json", "schema"), False)
|
|
self.assertIsNotNone(result.grammar_stats.compilation_time)
|
|
self.assertGreater(result.grammar_stats.compilation_time, 0)
|
|
|
|
def test_init_value_dispatch_no_stats(self):
|
|
"""When grammar has no stats, should not crash."""
|
|
mock_grammar = MagicMock(spec=BaseGrammarObject)
|
|
mock_grammar.grammar_stats = None
|
|
self.backend.dispatch_json = MagicMock(return_value=mock_grammar)
|
|
# Should not raise
|
|
self.backend._init_value_dispatch(("json", "schema"), False)
|
|
|
|
def test_reset_then_miss(self):
|
|
"""After reset, previously cached keys should be misses."""
|
|
key = ("json", "schema")
|
|
obj = MagicMock(spec=BaseGrammarObject)
|
|
obj.copy.return_value = obj
|
|
self.backend.set_cache(key, obj)
|
|
|
|
_, hit = self.backend.get_cached_or_future_value(key, False)
|
|
self.assertTrue(hit)
|
|
|
|
self.backend.reset()
|
|
result, hit = self.backend.get_cached_or_future_value(key, False)
|
|
self.assertFalse(hit)
|
|
self.assertIsInstance(result, Future)
|
|
|
|
def test_dispatch_fallback_error_message_content(self):
|
|
"""dispatch_fallback error should include the key type and value."""
|
|
with self.assertRaises(ValueError) as ctx:
|
|
self.backend.dispatch_fallback("custom_type", "custom_value")
|
|
self.assertIn("custom_type", str(ctx.exception))
|
|
self.assertIn("custom_value", str(ctx.exception))
|
|
|
|
def test_init_value_dispatch_none_grammar(self):
|
|
"""When dispatch returns None, should not crash on stats check."""
|
|
self.backend.dispatch_json = MagicMock(return_value=None)
|
|
result = self.backend._init_value_dispatch(("json", "schema"), False)
|
|
self.assertIsNone(result)
|
|
|
|
def test_cache_miss_duplicate_key_submits_separate_futures(self):
|
|
"""Two cache misses for the same key each get their own Future.
|
|
|
|
The backend does not deduplicate in-flight compilations — that is
|
|
handled at the GrammarManager level via grammar_queue. Each call
|
|
to get_cached_or_future_value with an uncached key submits a new
|
|
task to the executor."""
|
|
key = ("json", "schema")
|
|
result1, hit1 = self.backend.get_cached_or_future_value(key, False)
|
|
result2, hit2 = self.backend.get_cached_or_future_value(key, False)
|
|
|
|
self.assertFalse(hit1)
|
|
self.assertFalse(hit2)
|
|
self.assertIsInstance(result1, Future)
|
|
self.assertIsInstance(result2, Future)
|
|
# They are independent futures, not shared
|
|
self.assertIsNot(result1, result2)
|
|
|
|
# Both should complete successfully
|
|
self.assertIsInstance(result1.result(timeout=5), InvalidGrammarObject)
|
|
self.assertIsInstance(result2.result(timeout=5), InvalidGrammarObject)
|
|
|
|
|
|
class TestRegisterGrammarBackend(unittest.TestCase):
|
|
"""Test grammar backend registry."""
|
|
|
|
def setUp(self):
|
|
self._saved = dict(GRAMMAR_BACKEND_REGISTRY)
|
|
|
|
def tearDown(self):
|
|
GRAMMAR_BACKEND_REGISTRY.clear()
|
|
GRAMMAR_BACKEND_REGISTRY.update(self._saved)
|
|
|
|
def test_overwrite_registration(self):
|
|
register_grammar_backend("dup", lambda *a: "first")
|
|
register_grammar_backend("dup", lambda *a: "second")
|
|
self.assertEqual(
|
|
GRAMMAR_BACKEND_REGISTRY["dup"](None, None, None, None), "second"
|
|
)
|
|
|
|
|
|
class TestCreateGrammarBackend(unittest.TestCase):
|
|
"""Test create_grammar_backend factory function."""
|
|
|
|
def setUp(self):
|
|
self._saved = dict(GRAMMAR_BACKEND_REGISTRY)
|
|
|
|
def tearDown(self):
|
|
GRAMMAR_BACKEND_REGISTRY.clear()
|
|
GRAMMAR_BACKEND_REGISTRY.update(self._saved)
|
|
|
|
def _make_server_args(
|
|
self, backend="none", reasoning_parser=None, enable_strict_thinking=False
|
|
):
|
|
args = MagicMock()
|
|
args.override = lambda source, **updates: [
|
|
setattr(args, key, value) for key, value in updates.items()
|
|
]
|
|
args.grammar_backend = backend
|
|
args.reasoning_parser = reasoning_parser
|
|
args.enable_strict_thinking = enable_strict_thinking
|
|
args.constrained_json_whitespace_pattern = None
|
|
args.constrained_json_disable_any_whitespace = False
|
|
return args
|
|
|
|
def test_none_backend_returns_none(self):
|
|
args = self._make_server_args("none")
|
|
result = create_grammar_backend(args, None, 32000)
|
|
self.assertIsNone(result)
|
|
|
|
def test_none_backend_with_strict_thinking_raises(self):
|
|
args = self._make_server_args("none", enable_strict_thinking=True)
|
|
with self.assertRaisesRegex(ValueError, "enable-strict-thinking"):
|
|
create_grammar_backend(args, None, 32000)
|
|
|
|
def test_invalid_backend_raises(self):
|
|
args = self._make_server_args("nonexistent_backend")
|
|
with self.assertRaises(ValueError):
|
|
create_grammar_backend(args, None, 32000)
|
|
|
|
def test_custom_backend_receives_args(self):
|
|
received = {}
|
|
|
|
def init_fn(server_args, tokenizer, vocab_size, eos_token_ids):
|
|
received["server_args"] = server_args
|
|
received["tokenizer"] = tokenizer
|
|
received["vocab_size"] = vocab_size
|
|
received["eos_token_ids"] = eos_token_ids
|
|
return MagicMock()
|
|
|
|
register_grammar_backend("capture", init_fn)
|
|
args = self._make_server_args("capture")
|
|
create_grammar_backend(args, "my_tok", 50000, {0, 2})
|
|
self.assertEqual(received["tokenizer"], "my_tok")
|
|
self.assertEqual(received["vocab_size"], 50000)
|
|
self.assertEqual(received["eos_token_ids"], {0, 2})
|
|
|
|
def test_custom_backend_skips_reasoner_wrapping(self):
|
|
"""Custom registered backends return directly, bypassing reasoner wrapping."""
|
|
mock_inner = MagicMock(spec=BaseGrammarBackend)
|
|
register_grammar_backend("inner_r", lambda *a: mock_inner)
|
|
|
|
args = self._make_server_args("inner_r", reasoning_parser="deepseek-r1")
|
|
tokenizer = MagicMock()
|
|
|
|
result = create_grammar_backend(args, tokenizer, 32000)
|
|
# Custom backends return early, no reasoner wrapping applied
|
|
self.assertIs(result, mock_inner)
|
|
|
|
@patch("sglang.srt.constrained.outlines_backend.OutlinesGrammarBackend")
|
|
def test_outlines_backend(self, mock_outlines_cls):
|
|
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
|
mock_outlines_cls.return_value = mock_backend
|
|
args = self._make_server_args("outlines")
|
|
args.constrained_json_whitespace_pattern = r"\s*"
|
|
|
|
result = create_grammar_backend(args, "tok", 32000)
|
|
mock_outlines_cls.assert_called_once_with("tok", whitespace_pattern=r"\s*")
|
|
self.assertIs(result, mock_backend)
|
|
|
|
@patch("sglang.srt.constrained.xgrammar_backend.XGrammarGrammarBackend")
|
|
def test_xgrammar_backend(self, mock_xgrammar_cls):
|
|
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
|
mock_xgrammar_cls.return_value = mock_backend
|
|
args = self._make_server_args("xgrammar")
|
|
args.constrained_json_disable_any_whitespace = True
|
|
|
|
result = create_grammar_backend(args, "tok", 32000, {1, 2})
|
|
mock_xgrammar_cls.assert_called_once_with(
|
|
"tok", vocab_size=32000, model_eos_token_ids=[1, 2], any_whitespace=False
|
|
)
|
|
self.assertIs(result, mock_backend)
|
|
|
|
@patch("sglang.srt.constrained.xgrammar_backend.XGrammarGrammarBackend")
|
|
def test_xgrammar_unsupported_tokenizer_falls_back_to_none(self, mock_xgrammar_cls):
|
|
from sglang.srt.constrained.xgrammar_backend import TokenizerNotSupportedError
|
|
from sglang.srt.runtime_context import get_context, get_exec
|
|
|
|
mock_xgrammar_cls.side_effect = TokenizerNotSupportedError(
|
|
"unsupported tokenizer"
|
|
)
|
|
override = get_context().override_server_args(grammar_backend="xgrammar")
|
|
server_args = override.install()
|
|
self.addCleanup(override.restore)
|
|
|
|
self.assertIsNone(create_grammar_backend(server_args, "tok", 32000, {1}))
|
|
self.assertEqual(get_exec().kernel.grammar_backend, "none")
|
|
self.assertEqual(
|
|
get_context().resolved_server_args_dict()["grammar_backend"], "none"
|
|
)
|
|
self.assertEqual(server_args.grammar_backend, "xgrammar")
|
|
|
|
@patch("sglang.srt.constrained.llguidance_backend.GuidanceBackend")
|
|
def test_llguidance_backend(self, mock_guidance_cls):
|
|
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
|
mock_guidance_cls.return_value = mock_backend
|
|
args = self._make_server_args("llguidance")
|
|
args.constrained_json_disable_any_whitespace = False
|
|
args.constrained_json_whitespace_pattern = r"\s+"
|
|
|
|
result = create_grammar_backend(args, "tok", 32000, {1, 2})
|
|
mock_guidance_cls.assert_called_once_with(
|
|
tokenizer="tok",
|
|
any_whitespace=True,
|
|
whitespace_pattern=r"\s+",
|
|
n_vocab=32000,
|
|
eos_token_ids={1, 2},
|
|
)
|
|
self.assertIs(result, mock_backend)
|
|
|
|
@patch("sglang.srt.constrained.outlines_backend.OutlinesGrammarBackend")
|
|
def test_reasoner_wrapping_on_builtin_backend(self, mock_outlines_cls):
|
|
"""Non-custom backends get wrapped with ReasonerGrammarBackend."""
|
|
from sglang.srt.constrained.reasoner_grammar_backend import (
|
|
ReasonerGrammarBackend,
|
|
)
|
|
|
|
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
|
mock_backend.is_support_token_filter = False
|
|
mock_outlines_cls.return_value = mock_backend
|
|
args = self._make_server_args("outlines", reasoning_parser="deepseek-r1")
|
|
tokenizer = MagicMock()
|
|
# encode must return a single-token list for think_start/end tokens
|
|
tokenizer.encode.return_value = [42]
|
|
|
|
result = create_grammar_backend(args, tokenizer, 32000, think_end_ids=[42])
|
|
self.assertIsInstance(result, ReasonerGrammarBackend)
|
|
self.assertIs(result.grammar_backend, mock_backend)
|
|
|
|
@patch("sglang.srt.constrained.outlines_backend.OutlinesGrammarBackend")
|
|
def test_no_reasoner_wrapping_without_think_end_ids(self, mock_outlines_cls):
|
|
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
|
mock_outlines_cls.return_value = mock_backend
|
|
args = self._make_server_args("outlines", reasoning_parser="deepseek-r1")
|
|
tokenizer = MagicMock(spec=[])
|
|
|
|
result = create_grammar_backend(args, tokenizer, 32000, think_end_ids=None)
|
|
self.assertIs(result, mock_backend)
|
|
|
|
@patch("sglang.srt.constrained.outlines_backend.OutlinesGrammarBackend")
|
|
def test_no_reasoner_wrapping_without_reasoning_parser(self, mock_outlines_cls):
|
|
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
|
mock_outlines_cls.return_value = mock_backend
|
|
args = self._make_server_args("outlines", reasoning_parser=None)
|
|
tokenizer = MagicMock()
|
|
|
|
result = create_grammar_backend(args, tokenizer, 32000, think_end_ids=[42])
|
|
self.assertIs(result, mock_backend)
|
|
|
|
@patch("sglang.srt.constrained.xgrammar_backend.XGrammarGrammarBackend")
|
|
def test_xgrammar_eos_none(self, mock_xgrammar_cls):
|
|
"""eos_token_ids=None should pass None, not an empty list."""
|
|
mock_xgrammar_cls.return_value = MagicMock(spec=BaseGrammarBackend)
|
|
args = self._make_server_args("xgrammar")
|
|
|
|
create_grammar_backend(args, "tok", 32000, None)
|
|
_, kwargs = mock_xgrammar_cls.call_args
|
|
self.assertIsNone(kwargs["model_eos_token_ids"])
|
|
|
|
|
|
class TestLlguidanceStructuralTagTriggerPairing(unittest.TestCase):
|
|
"""Bug regression: dispatch_structural_tag paired EVERY structure with
|
|
triggers[0]. Detectors with per-tool triggers (Inkling emits
|
|
<|message_model|>{name}<|content_invoke_tool_json|> per tool) produce
|
|
multiple distinct triggers, and llguidance's StructTag asserts
|
|
begin.startswith(trigger) — so any multi-tool constrained request
|
|
compiled to InvalidGrammarObject."""
|
|
|
|
def test_each_structure_pairs_with_its_own_trigger(self):
|
|
import json
|
|
|
|
from sglang.srt.constrained.llguidance_backend import GuidanceBackend
|
|
|
|
backend = object.__new__(GuidanceBackend)
|
|
backend._from_serialized = lambda serialized: serialized
|
|
begins = [
|
|
'<|message_model|>alpha<|content_invoke_tool_json|>{"name":"alpha","args":',
|
|
'<|message_model|>beta<|content_invoke_tool_json|>{"name":"beta","args":',
|
|
]
|
|
key = json.dumps(
|
|
{
|
|
"type": "structural_tag",
|
|
"structures": [
|
|
{
|
|
"begin": begin,
|
|
"schema": {"type": "object"},
|
|
"end": "<|end_message|>",
|
|
}
|
|
for begin in begins
|
|
],
|
|
"triggers": [
|
|
"<|message_model|>alpha<|content_invoke_tool_json|>",
|
|
"<|message_model|>beta<|content_invoke_tool_json|>",
|
|
],
|
|
}
|
|
)
|
|
result = backend.dispatch_structural_tag(key)
|
|
self.assertNotIsInstance(result, InvalidGrammarObject)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|