[Test] Add unit tests for srt/constrained module (#21010)
This commit is contained in:
@@ -0,0 +1,431 @@
|
|||||||
|
"""
|
||||||
|
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, "stage-a-cpu-only")
|
||||||
|
|
||||||
|
|
||||||
|
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."""
|
||||||
|
|
||||||
|
def test_is_terminated_default(self):
|
||||||
|
obj = BaseGrammarObject()
|
||||||
|
self.assertFalse(obj.is_terminated())
|
||||||
|
|
||||||
|
def test_maybe_init_reasoning_noop(self):
|
||||||
|
obj = BaseGrammarObject()
|
||||||
|
obj.maybe_init_reasoning(True) # Should not raise
|
||||||
|
|
||||||
|
|
||||||
|
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_set_and_get_cache(self):
|
||||||
|
obj = BaseGrammarObject()
|
||||||
|
key = ("json", '{"type": "object"}')
|
||||||
|
self.backend.set_cache(key, obj)
|
||||||
|
self.assertIn(key, self.backend.cache)
|
||||||
|
self.assertIs(self.backend.cache[key], obj)
|
||||||
|
|
||||||
|
def test_reset_clears_cache(self):
|
||||||
|
self.backend.set_cache(("json", "schema"), BaseGrammarObject())
|
||||||
|
self.backend.reset()
|
||||||
|
self.assertEqual(len(self.backend.cache), 0)
|
||||||
|
|
||||||
|
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_all_dispatch_methods_unsupported(self):
|
||||||
|
"""All dispatch methods on base class return InvalidGrammarObject."""
|
||||||
|
cases = [
|
||||||
|
("dispatch_json", ("schema",)),
|
||||||
|
("dispatch_regex", ("[a-z]+",)),
|
||||||
|
("dispatch_ebnf", ("root ::= 'hello'",)),
|
||||||
|
("dispatch_structural_tag", ("{}",)),
|
||||||
|
]
|
||||||
|
for method_name, args in cases:
|
||||||
|
with self.subTest(method=method_name):
|
||||||
|
result = getattr(self.backend, method_name)(*args)
|
||||||
|
self.assertIsInstance(result, 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_register_and_use(self):
|
||||||
|
mock_init = MagicMock(return_value="custom_backend")
|
||||||
|
register_grammar_backend("my_backend", mock_init)
|
||||||
|
self.assertIn("my_backend", GRAMMAR_BACKEND_REGISTRY)
|
||||||
|
|
||||||
|
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):
|
||||||
|
args = MagicMock()
|
||||||
|
args.grammar_backend = backend
|
||||||
|
args.reasoning_parser = reasoning_parser
|
||||||
|
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_invalid_backend_raises(self):
|
||||||
|
args = self._make_server_args("nonexistent_backend")
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
create_grammar_backend(args, None, 32000)
|
||||||
|
|
||||||
|
def test_custom_registered_backend(self):
|
||||||
|
mock_backend = MagicMock()
|
||||||
|
register_grammar_backend("test_custom", lambda *a: mock_backend)
|
||||||
|
args = self._make_server_args("test_custom")
|
||||||
|
result = create_grammar_backend(args, "tok", 32000, {1, 2})
|
||||||
|
self.assertIs(result, mock_backend)
|
||||||
|
|
||||||
|
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")
|
||||||
|
tokenizer = MagicMock()
|
||||||
|
tokenizer.think_end_id = 42
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
mock_xgrammar_cls.side_effect = TokenizerNotSupportedError(
|
||||||
|
"unsupported tokenizer"
|
||||||
|
)
|
||||||
|
args = self._make_server_args("xgrammar")
|
||||||
|
|
||||||
|
result = create_grammar_backend(args, "tok", 32000, {1})
|
||||||
|
self.assertIsNone(result)
|
||||||
|
self.assertEqual(args.grammar_backend, "none")
|
||||||
|
|
||||||
|
@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)
|
||||||
|
mock_guidance_cls.assert_called_once_with(
|
||||||
|
tokenizer="tok", any_whitespace=True, whitespace_pattern=r"\s+"
|
||||||
|
)
|
||||||
|
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_outlines_cls.return_value = mock_backend
|
||||||
|
args = self._make_server_args("outlines", reasoning_parser="deepseek")
|
||||||
|
tokenizer = MagicMock()
|
||||||
|
tokenizer.think_end_id = 42
|
||||||
|
|
||||||
|
result = create_grammar_backend(args, tokenizer, 32000)
|
||||||
|
self.assertIsInstance(result, ReasonerGrammarBackend)
|
||||||
|
self.assertEqual(result.think_end_id, 42)
|
||||||
|
self.assertIs(result.grammar_backend, mock_backend)
|
||||||
|
|
||||||
|
@patch("sglang.srt.constrained.outlines_backend.OutlinesGrammarBackend")
|
||||||
|
def test_no_reasoner_wrapping_without_think_end_id(self, mock_outlines_cls):
|
||||||
|
"""Without think_end_id on tokenizer, no reasoner wrapping."""
|
||||||
|
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
||||||
|
mock_outlines_cls.return_value = mock_backend
|
||||||
|
args = self._make_server_args("outlines", reasoning_parser="deepseek")
|
||||||
|
tokenizer = MagicMock(spec=[]) # No think_end_id attribute
|
||||||
|
|
||||||
|
result = create_grammar_backend(args, tokenizer, 32000)
|
||||||
|
self.assertIs(result, mock_backend)
|
||||||
|
|
||||||
|
@patch("sglang.srt.constrained.outlines_backend.OutlinesGrammarBackend")
|
||||||
|
def test_no_reasoner_wrapping_without_reasoning_parser(self, mock_outlines_cls):
|
||||||
|
"""Without reasoning_parser, no reasoner wrapping even with think_end_id."""
|
||||||
|
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
||||||
|
mock_outlines_cls.return_value = mock_backend
|
||||||
|
args = self._make_server_args("outlines", reasoning_parser=None)
|
||||||
|
tokenizer = MagicMock()
|
||||||
|
tokenizer.think_end_id = 42
|
||||||
|
|
||||||
|
result = create_grammar_backend(args, tokenizer, 32000)
|
||||||
|
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"])
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,583 @@
|
|||||||
|
"""
|
||||||
|
Unit tests for sglang.srt.constrained.grammar_manager.
|
||||||
|
|
||||||
|
Test Coverage:
|
||||||
|
- GrammarManager initialization, queue management, len, clear
|
||||||
|
- process_req_with_grammar: dispatch by constraint type (json, regex, ebnf,
|
||||||
|
structural_tag), no-constraint requests, no-backend error, cache hits,
|
||||||
|
cached invalid grammar abort
|
||||||
|
- abort_requests: single abort, abort all, future cancellation
|
||||||
|
- get_ready_grammar_requests: future completion, invalid grammar handling,
|
||||||
|
timeout with max poll iterations, aborted request handling, queue cleanup
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python -m pytest test_grammar_manager.py -v
|
||||||
|
"""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from concurrent.futures import Future
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
from sglang.srt.constrained.base_grammar_backend import (
|
||||||
|
BaseGrammarBackend,
|
||||||
|
BaseGrammarObject,
|
||||||
|
InvalidGrammarObject,
|
||||||
|
)
|
||||||
|
from sglang.srt.constrained.grammar_manager import GrammarManager
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(2.0, "stage-a-cpu-only")
|
||||||
|
|
||||||
|
|
||||||
|
def _make_scheduler(grammar_backend_name="none", skip_tokenizer=False):
|
||||||
|
"""Create a mock scheduler with necessary attributes."""
|
||||||
|
scheduler = MagicMock()
|
||||||
|
scheduler.server_args.grammar_backend = grammar_backend_name
|
||||||
|
scheduler.server_args.skip_tokenizer_init = skip_tokenizer
|
||||||
|
scheduler.server_args.reasoning_parser = None
|
||||||
|
scheduler.server_args.constrained_json_whitespace_pattern = None
|
||||||
|
scheduler.server_args.constrained_json_disable_any_whitespace = False
|
||||||
|
|
||||||
|
# Distributed group mocks
|
||||||
|
scheduler.dp_tp_cpu_group = MagicMock()
|
||||||
|
scheduler.dp_tp_group.world_size = 1
|
||||||
|
scheduler.dp_tp_group.first_rank = 0
|
||||||
|
scheduler.dp_tp_group.is_first_rank = True
|
||||||
|
|
||||||
|
return scheduler
|
||||||
|
|
||||||
|
|
||||||
|
def _make_req(
|
||||||
|
json_schema=None, regex=None, ebnf=None, structural_tag=None, rid="req-1"
|
||||||
|
):
|
||||||
|
"""Create a mock request with sampling params."""
|
||||||
|
req = MagicMock()
|
||||||
|
req.rid = rid
|
||||||
|
req.sampling_params.json_schema = json_schema
|
||||||
|
req.sampling_params.regex = regex
|
||||||
|
req.sampling_params.ebnf = ebnf
|
||||||
|
req.sampling_params.structural_tag = structural_tag
|
||||||
|
req.require_reasoning = False
|
||||||
|
req.grammar = None
|
||||||
|
req.grammar_key = None
|
||||||
|
req.grammar_wait_ct = 0
|
||||||
|
req.finished.return_value = False
|
||||||
|
return req
|
||||||
|
|
||||||
|
|
||||||
|
class TestGrammarManagerInit(unittest.TestCase):
|
||||||
|
"""Test GrammarManager initialization."""
|
||||||
|
|
||||||
|
@patch("sglang.srt.constrained.grammar_manager.create_grammar_backend")
|
||||||
|
def test_init_with_backend(self, mock_create):
|
||||||
|
mock_create.return_value = MagicMock(spec=BaseGrammarBackend)
|
||||||
|
scheduler = _make_scheduler("xgrammar")
|
||||||
|
scheduler.server_args.skip_tokenizer_init = False
|
||||||
|
|
||||||
|
mgr = GrammarManager(scheduler)
|
||||||
|
self.assertIsNotNone(mgr.grammar_backend)
|
||||||
|
self.assertEqual(len(mgr), 0)
|
||||||
|
|
||||||
|
def test_init_skip_tokenizer(self):
|
||||||
|
scheduler = _make_scheduler(skip_tokenizer=True)
|
||||||
|
mgr = GrammarManager(scheduler)
|
||||||
|
self.assertIsNone(mgr.grammar_backend)
|
||||||
|
|
||||||
|
@patch("sglang.srt.constrained.grammar_manager.create_grammar_backend")
|
||||||
|
def test_len_and_has_waiting(self, mock_create):
|
||||||
|
mock_create.return_value = None
|
||||||
|
scheduler = _make_scheduler()
|
||||||
|
mgr = GrammarManager(scheduler)
|
||||||
|
self.assertEqual(len(mgr), 0)
|
||||||
|
self.assertFalse(mgr.has_waiting_grammars())
|
||||||
|
|
||||||
|
@patch("sglang.srt.constrained.grammar_manager.create_grammar_backend")
|
||||||
|
def test_clear_resets_backend(self, mock_create):
|
||||||
|
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
||||||
|
mock_create.return_value = mock_backend
|
||||||
|
scheduler = _make_scheduler()
|
||||||
|
scheduler.server_args.skip_tokenizer_init = False
|
||||||
|
|
||||||
|
mgr = GrammarManager(scheduler)
|
||||||
|
mgr.clear()
|
||||||
|
mock_backend.reset.assert_called_once()
|
||||||
|
|
||||||
|
@patch("sglang.srt.constrained.grammar_manager.create_grammar_backend")
|
||||||
|
def test_clear_no_backend(self, mock_create):
|
||||||
|
mock_create.return_value = None
|
||||||
|
scheduler = _make_scheduler()
|
||||||
|
mgr = GrammarManager(scheduler)
|
||||||
|
mgr.clear() # Should not raise
|
||||||
|
|
||||||
|
|
||||||
|
class TestProcessReqWithGrammar(unittest.TestCase):
|
||||||
|
"""Test process_req_with_grammar dispatch and caching."""
|
||||||
|
|
||||||
|
def _make_mgr(self):
|
||||||
|
scheduler = _make_scheduler()
|
||||||
|
scheduler.server_args.skip_tokenizer_init = True
|
||||||
|
mgr = GrammarManager(scheduler)
|
||||||
|
mgr.grammar_backend = MagicMock(spec=BaseGrammarBackend)
|
||||||
|
return mgr
|
||||||
|
|
||||||
|
def test_no_constraint_returns_false(self):
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
req = _make_req() # No constraints
|
||||||
|
result = mgr.process_req_with_grammar(req)
|
||||||
|
self.assertFalse(result)
|
||||||
|
self.assertEqual(len(mgr.grammar_queue), 0)
|
||||||
|
|
||||||
|
def test_json_schema_cache_miss(self):
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
future = Future()
|
||||||
|
mgr.grammar_backend.get_cached_or_future_value.return_value = (future, False)
|
||||||
|
|
||||||
|
req = _make_req(json_schema='{"type": "object"}')
|
||||||
|
result = mgr.process_req_with_grammar(req)
|
||||||
|
|
||||||
|
self.assertTrue(result)
|
||||||
|
self.assertEqual(len(mgr.grammar_queue), 1)
|
||||||
|
self.assertEqual(req.grammar_key, ("json", '{"type": "object"}'))
|
||||||
|
|
||||||
|
def test_regex_cache_miss(self):
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
future = Future()
|
||||||
|
mgr.grammar_backend.get_cached_or_future_value.return_value = (future, False)
|
||||||
|
|
||||||
|
req = _make_req(regex="[a-z]+")
|
||||||
|
result = mgr.process_req_with_grammar(req)
|
||||||
|
|
||||||
|
self.assertTrue(result)
|
||||||
|
self.assertEqual(req.grammar_key, ("regex", "[a-z]+"))
|
||||||
|
|
||||||
|
def test_ebnf_cache_miss(self):
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
future = Future()
|
||||||
|
mgr.grammar_backend.get_cached_or_future_value.return_value = (future, False)
|
||||||
|
|
||||||
|
req = _make_req(ebnf="root ::= 'hello'")
|
||||||
|
result = mgr.process_req_with_grammar(req)
|
||||||
|
|
||||||
|
self.assertTrue(result)
|
||||||
|
self.assertEqual(req.grammar_key, ("ebnf", "root ::= 'hello'"))
|
||||||
|
|
||||||
|
def test_structural_tag_cache_miss(self):
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
future = Future()
|
||||||
|
mgr.grammar_backend.get_cached_or_future_value.return_value = (future, False)
|
||||||
|
|
||||||
|
req = _make_req(structural_tag='{"structures": [], "triggers": []}')
|
||||||
|
result = mgr.process_req_with_grammar(req)
|
||||||
|
|
||||||
|
self.assertTrue(result)
|
||||||
|
self.assertEqual(
|
||||||
|
req.grammar_key,
|
||||||
|
("structural_tag", '{"structures": [], "triggers": []}'),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_cache_hit_returns_false(self):
|
||||||
|
"""Cache hit should NOT add to grammar queue."""
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
grammar_obj = MagicMock(spec=BaseGrammarObject)
|
||||||
|
mgr.grammar_backend.get_cached_or_future_value.return_value = (
|
||||||
|
grammar_obj,
|
||||||
|
True,
|
||||||
|
)
|
||||||
|
|
||||||
|
req = _make_req(json_schema='{"type": "object"}')
|
||||||
|
result = mgr.process_req_with_grammar(req)
|
||||||
|
|
||||||
|
self.assertFalse(result)
|
||||||
|
self.assertEqual(len(mgr.grammar_queue), 0)
|
||||||
|
self.assertIs(req.grammar, grammar_obj)
|
||||||
|
|
||||||
|
def test_cache_hit_invalid_grammar_aborts(self):
|
||||||
|
"""Cache hit with InvalidGrammarObject should abort the request."""
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
invalid = InvalidGrammarObject("bad schema")
|
||||||
|
mgr.grammar_backend.get_cached_or_future_value.return_value = (invalid, True)
|
||||||
|
|
||||||
|
req = _make_req(json_schema="bad")
|
||||||
|
result = mgr.process_req_with_grammar(req)
|
||||||
|
|
||||||
|
self.assertFalse(result)
|
||||||
|
req.set_finish_with_abort.assert_called_once()
|
||||||
|
self.assertIn("bad schema", req.set_finish_with_abort.call_args[0][0])
|
||||||
|
|
||||||
|
def test_no_backend_aborts(self):
|
||||||
|
"""No grammar backend should abort request."""
|
||||||
|
scheduler = _make_scheduler()
|
||||||
|
scheduler.server_args.skip_tokenizer_init = True
|
||||||
|
mgr = GrammarManager(scheduler)
|
||||||
|
mgr.grammar_backend = None
|
||||||
|
|
||||||
|
req = _make_req(json_schema='{"type": "object"}')
|
||||||
|
result = mgr.process_req_with_grammar(req)
|
||||||
|
|
||||||
|
self.assertFalse(result)
|
||||||
|
req.set_finish_with_abort.assert_called_once()
|
||||||
|
self.assertIn("not supported", req.set_finish_with_abort.call_args[0][0])
|
||||||
|
|
||||||
|
def test_json_takes_priority_over_other_constraints(self):
|
||||||
|
"""When json_schema is set, it should be used regardless of other fields."""
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
future = Future()
|
||||||
|
mgr.grammar_backend.get_cached_or_future_value.return_value = (future, False)
|
||||||
|
|
||||||
|
req = _make_req(json_schema='{"type": "object"}', regex="[a-z]+")
|
||||||
|
mgr.process_req_with_grammar(req)
|
||||||
|
self.assertEqual(req.grammar_key, ("json", '{"type": "object"}'))
|
||||||
|
|
||||||
|
def test_require_reasoning_forwarded_to_backend(self):
|
||||||
|
"""require_reasoning from the request should be passed to the backend."""
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
grammar_obj = MagicMock(spec=BaseGrammarObject)
|
||||||
|
mgr.grammar_backend.get_cached_or_future_value.return_value = (
|
||||||
|
grammar_obj,
|
||||||
|
True,
|
||||||
|
)
|
||||||
|
|
||||||
|
req = _make_req(json_schema="schema")
|
||||||
|
req.require_reasoning = True
|
||||||
|
mgr.process_req_with_grammar(req)
|
||||||
|
|
||||||
|
mgr.grammar_backend.get_cached_or_future_value.assert_called_once_with(
|
||||||
|
("json", "schema"), True
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_has_waiting_grammars_after_enqueue(self):
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
future = Future()
|
||||||
|
mgr.grammar_backend.get_cached_or_future_value.return_value = (future, False)
|
||||||
|
|
||||||
|
self.assertFalse(mgr.has_waiting_grammars())
|
||||||
|
req = _make_req(json_schema="schema")
|
||||||
|
mgr.process_req_with_grammar(req)
|
||||||
|
self.assertTrue(mgr.has_waiting_grammars())
|
||||||
|
self.assertEqual(len(mgr), 1)
|
||||||
|
|
||||||
|
|
||||||
|
class TestAbortRequests(unittest.TestCase):
|
||||||
|
"""Test abort_requests handling."""
|
||||||
|
|
||||||
|
def _make_mgr_with_queue(self):
|
||||||
|
scheduler = _make_scheduler()
|
||||||
|
scheduler.server_args.skip_tokenizer_init = True
|
||||||
|
mgr = GrammarManager(scheduler)
|
||||||
|
mgr.grammar_backend = MagicMock(spec=BaseGrammarBackend)
|
||||||
|
return mgr
|
||||||
|
|
||||||
|
def test_abort_by_rid_prefix(self):
|
||||||
|
mgr = self._make_mgr_with_queue()
|
||||||
|
req = _make_req(rid="req-123")
|
||||||
|
future = MagicMock(spec=Future)
|
||||||
|
req.grammar = future
|
||||||
|
mgr.grammar_queue.append(req)
|
||||||
|
|
||||||
|
abort_req = MagicMock()
|
||||||
|
abort_req.abort_all = False
|
||||||
|
abort_req.rid = "req-123"
|
||||||
|
|
||||||
|
mgr.abort_requests(abort_req)
|
||||||
|
future.cancel.assert_called_once()
|
||||||
|
req.set_finish_with_abort.assert_called_once()
|
||||||
|
|
||||||
|
def test_abort_non_matching_rid(self):
|
||||||
|
mgr = self._make_mgr_with_queue()
|
||||||
|
req = _make_req(rid="req-999")
|
||||||
|
req.grammar = MagicMock(spec=Future)
|
||||||
|
mgr.grammar_queue.append(req)
|
||||||
|
|
||||||
|
abort_req = MagicMock()
|
||||||
|
abort_req.abort_all = False
|
||||||
|
abort_req.rid = "req-123"
|
||||||
|
|
||||||
|
mgr.abort_requests(abort_req)
|
||||||
|
req.set_finish_with_abort.assert_not_called()
|
||||||
|
|
||||||
|
def test_abort_all(self):
|
||||||
|
mgr = self._make_mgr_with_queue()
|
||||||
|
reqs = []
|
||||||
|
for i in range(3):
|
||||||
|
req = _make_req(rid=f"req-{i}")
|
||||||
|
req.grammar = MagicMock(spec=Future)
|
||||||
|
mgr.grammar_queue.append(req)
|
||||||
|
reqs.append(req)
|
||||||
|
|
||||||
|
abort_req = MagicMock()
|
||||||
|
abort_req.abort_all = True
|
||||||
|
abort_req.rid = ""
|
||||||
|
|
||||||
|
mgr.abort_requests(abort_req)
|
||||||
|
for req in reqs:
|
||||||
|
req.set_finish_with_abort.assert_called_once()
|
||||||
|
|
||||||
|
def test_abort_empty_queue(self):
|
||||||
|
"""Aborting on an empty queue should not raise."""
|
||||||
|
mgr = self._make_mgr_with_queue()
|
||||||
|
abort_req = MagicMock()
|
||||||
|
abort_req.abort_all = True
|
||||||
|
abort_req.rid = ""
|
||||||
|
mgr.abort_requests(abort_req) # Should not raise
|
||||||
|
|
||||||
|
def test_abort_prefix_match(self):
|
||||||
|
"""rid.startswith means prefix matching, not exact matching."""
|
||||||
|
mgr = self._make_mgr_with_queue()
|
||||||
|
req = _make_req(rid="req-123-suffix")
|
||||||
|
req.grammar = MagicMock(spec=Future)
|
||||||
|
mgr.grammar_queue.append(req)
|
||||||
|
|
||||||
|
abort_req = MagicMock()
|
||||||
|
abort_req.abort_all = False
|
||||||
|
abort_req.rid = "req-123"
|
||||||
|
|
||||||
|
mgr.abort_requests(abort_req)
|
||||||
|
req.set_finish_with_abort.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
class TestGetReadyGrammarRequests(unittest.TestCase):
|
||||||
|
"""Test get_ready_grammar_requests polling and result handling."""
|
||||||
|
|
||||||
|
def _make_mgr(self):
|
||||||
|
scheduler = _make_scheduler()
|
||||||
|
scheduler.server_args.skip_tokenizer_init = True
|
||||||
|
mgr = GrammarManager(scheduler)
|
||||||
|
mgr.grammar_backend = MagicMock(spec=BaseGrammarBackend)
|
||||||
|
# Use very short poll interval for tests
|
||||||
|
mgr.SGLANG_GRAMMAR_POLL_INTERVAL = 0.01
|
||||||
|
mgr.SGLANG_GRAMMAR_MAX_POLL_ITERATIONS = 3
|
||||||
|
return mgr
|
||||||
|
|
||||||
|
def test_ready_future_returns_req(self):
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
|
||||||
|
grammar_obj = MagicMock(spec=BaseGrammarObject)
|
||||||
|
grammar_obj.copy.return_value = grammar_obj
|
||||||
|
future = Future()
|
||||||
|
future.set_result(grammar_obj)
|
||||||
|
|
||||||
|
req = _make_req(json_schema="schema")
|
||||||
|
req.grammar = future
|
||||||
|
req.grammar_key = ("json", "schema")
|
||||||
|
mgr.grammar_queue.append(req)
|
||||||
|
|
||||||
|
result = mgr.get_ready_grammar_requests()
|
||||||
|
self.assertEqual(len(result), 1)
|
||||||
|
self.assertIs(result[0], req)
|
||||||
|
self.assertIs(req.grammar, grammar_obj)
|
||||||
|
# Cache should be set
|
||||||
|
mgr.grammar_backend.set_cache.assert_called_once()
|
||||||
|
# Queue should be empty
|
||||||
|
self.assertEqual(len(mgr.grammar_queue), 0)
|
||||||
|
|
||||||
|
def test_invalid_grammar_aborts_req(self):
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
|
||||||
|
invalid = InvalidGrammarObject("compile error")
|
||||||
|
invalid_copy = InvalidGrammarObject("compile error")
|
||||||
|
invalid.copy = MagicMock(return_value=invalid_copy)
|
||||||
|
future = Future()
|
||||||
|
future.set_result(invalid)
|
||||||
|
|
||||||
|
req = _make_req(json_schema="bad")
|
||||||
|
req.grammar = future
|
||||||
|
req.grammar_key = ("json", "bad")
|
||||||
|
mgr.grammar_queue.append(req)
|
||||||
|
|
||||||
|
result = mgr.get_ready_grammar_requests()
|
||||||
|
self.assertEqual(len(result), 1)
|
||||||
|
req.set_finish_with_abort.assert_called_once()
|
||||||
|
self.assertIn("compile error", req.set_finish_with_abort.call_args[0][0])
|
||||||
|
|
||||||
|
def test_aborted_req_removed_from_queue(self):
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
|
||||||
|
req = _make_req(json_schema="schema")
|
||||||
|
req.finished.return_value = True # Already aborted
|
||||||
|
req.grammar = None
|
||||||
|
mgr.grammar_queue.append(req)
|
||||||
|
|
||||||
|
result = mgr.get_ready_grammar_requests()
|
||||||
|
self.assertEqual(len(result), 1)
|
||||||
|
self.assertEqual(len(mgr.grammar_queue), 0)
|
||||||
|
|
||||||
|
def test_timeout_aborts_req(self):
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
mgr.SGLANG_GRAMMAR_MAX_POLL_ITERATIONS = 1
|
||||||
|
|
||||||
|
future = Future() # Never completes
|
||||||
|
req = _make_req(json_schema="slow")
|
||||||
|
req.grammar = future
|
||||||
|
req.grammar_key = ("json", "slow")
|
||||||
|
req.grammar_wait_ct = 0
|
||||||
|
mgr.grammar_queue.append(req)
|
||||||
|
|
||||||
|
# First call: not ready, increments wait_ct to 1 (== max_poll)
|
||||||
|
result = mgr.get_ready_grammar_requests()
|
||||||
|
# Should timeout and abort
|
||||||
|
self.assertEqual(len(result), 1)
|
||||||
|
req.set_finish_with_abort.assert_called_once()
|
||||||
|
self.assertIn("timed out", req.set_finish_with_abort.call_args[0][0])
|
||||||
|
# Cache should store InvalidGrammarObject for timeout
|
||||||
|
mgr.grammar_backend.set_cache.assert_called_once()
|
||||||
|
cached_key, cached_val = mgr.grammar_backend.set_cache.call_args[0]
|
||||||
|
self.assertEqual(cached_key, ("json", "slow"))
|
||||||
|
self.assertIsInstance(cached_val, InvalidGrammarObject)
|
||||||
|
|
||||||
|
def test_pending_future_stays_in_queue(self):
|
||||||
|
"""Futures that aren't done stay in the queue."""
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
mgr.SGLANG_GRAMMAR_MAX_POLL_ITERATIONS = 100 # High to avoid timeout
|
||||||
|
|
||||||
|
future = Future() # Never completes
|
||||||
|
req = _make_req(json_schema="pending")
|
||||||
|
req.grammar = future
|
||||||
|
req.grammar_key = ("json", "pending")
|
||||||
|
req.grammar_wait_ct = 0
|
||||||
|
mgr.grammar_queue.append(req)
|
||||||
|
|
||||||
|
result = mgr.get_ready_grammar_requests()
|
||||||
|
self.assertEqual(len(result), 0)
|
||||||
|
self.assertEqual(len(mgr.grammar_queue), 1)
|
||||||
|
self.assertEqual(req.grammar_wait_ct, 1)
|
||||||
|
|
||||||
|
def test_mixed_ready_and_pending(self):
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
mgr.SGLANG_GRAMMAR_MAX_POLL_ITERATIONS = 100
|
||||||
|
|
||||||
|
# Ready request
|
||||||
|
grammar_obj = MagicMock(spec=BaseGrammarObject)
|
||||||
|
grammar_obj.copy.return_value = grammar_obj
|
||||||
|
done_future = Future()
|
||||||
|
done_future.set_result(grammar_obj)
|
||||||
|
ready_req = _make_req(json_schema="ready", rid="r1")
|
||||||
|
ready_req.grammar = done_future
|
||||||
|
ready_req.grammar_key = ("json", "ready")
|
||||||
|
|
||||||
|
# Pending request
|
||||||
|
pending_future = Future()
|
||||||
|
pending_req = _make_req(json_schema="pending", rid="r2")
|
||||||
|
pending_req.grammar = pending_future
|
||||||
|
pending_req.grammar_key = ("json", "pending")
|
||||||
|
pending_req.grammar_wait_ct = 0
|
||||||
|
|
||||||
|
mgr.grammar_queue = [ready_req, pending_req]
|
||||||
|
|
||||||
|
result = mgr.get_ready_grammar_requests()
|
||||||
|
self.assertEqual(len(result), 1)
|
||||||
|
self.assertIs(result[0], ready_req)
|
||||||
|
self.assertEqual(len(mgr.grammar_queue), 1)
|
||||||
|
self.assertIs(mgr.grammar_queue[0], pending_req)
|
||||||
|
|
||||||
|
def test_empty_queue(self):
|
||||||
|
"""get_ready_grammar_requests on empty queue should return empty list."""
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
result = mgr.get_ready_grammar_requests()
|
||||||
|
self.assertEqual(len(result), 0)
|
||||||
|
self.assertEqual(len(mgr.grammar_queue), 0)
|
||||||
|
|
||||||
|
def test_progressive_timeout(self):
|
||||||
|
"""Request with partial wait_ct should timeout after remaining iterations."""
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
mgr.SGLANG_GRAMMAR_MAX_POLL_ITERATIONS = 3
|
||||||
|
|
||||||
|
future = Future() # Never completes
|
||||||
|
req = _make_req(json_schema="slow")
|
||||||
|
req.grammar = future
|
||||||
|
req.grammar_key = ("json", "slow")
|
||||||
|
req.grammar_wait_ct = 2 # Already waited 2 iterations
|
||||||
|
mgr.grammar_queue.append(req)
|
||||||
|
|
||||||
|
# wait_ct increments to 3 (== max), should timeout
|
||||||
|
result = mgr.get_ready_grammar_requests()
|
||||||
|
self.assertEqual(len(result), 1)
|
||||||
|
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_propagates(self):
|
||||||
|
"""A future that raised an exception should propagate on .result()."""
|
||||||
|
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)
|
||||||
|
|
||||||
|
with self.assertRaises(RuntimeError):
|
||||||
|
mgr.get_ready_grammar_requests()
|
||||||
|
|
||||||
|
@patch("sglang.srt.constrained.grammar_manager.torch.distributed.all_gather_object")
|
||||||
|
def test_multi_rank_sync_intersects_ready_unions_failed(self, mock_all_gather):
|
||||||
|
"""With multiple ranks, ready = intersection, failed = union."""
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
mgr.grammar_sync_size = 2 # Enable multi-rank path
|
||||||
|
|
||||||
|
# Two requests: idx 0 ready on both ranks, idx 1 ready only on rank 0
|
||||||
|
grammar_obj = MagicMock(spec=BaseGrammarObject)
|
||||||
|
grammar_obj.copy.return_value = grammar_obj
|
||||||
|
done_future = Future()
|
||||||
|
done_future.set_result(grammar_obj)
|
||||||
|
|
||||||
|
req0 = _make_req(json_schema="s0", rid="r0")
|
||||||
|
req0.grammar = done_future
|
||||||
|
req0.grammar_key = ("json", "s0")
|
||||||
|
|
||||||
|
pending_future = Future()
|
||||||
|
req1 = _make_req(json_schema="s1", rid="r1")
|
||||||
|
req1.grammar = pending_future
|
||||||
|
req1.grammar_key = ("json", "s1")
|
||||||
|
req1.grammar_wait_ct = 0
|
||||||
|
|
||||||
|
mgr.grammar_queue = [req0, req1]
|
||||||
|
mgr.SGLANG_GRAMMAR_MAX_POLL_ITERATIONS = 100
|
||||||
|
|
||||||
|
# Simulate all_gather: rank 0 has {0} ready, rank 1 has {0,1} ready
|
||||||
|
def fake_all_gather(output_list, _obj, group=None): # noqa: ARG001
|
||||||
|
output_list[0] = ({0}, set()) # rank 0: only idx 0 ready
|
||||||
|
output_list[1] = ({0, 1}, set()) # rank 1: both ready
|
||||||
|
|
||||||
|
mock_all_gather.side_effect = fake_all_gather
|
||||||
|
|
||||||
|
result = mgr.get_ready_grammar_requests()
|
||||||
|
# Intersection of ready: {0} ∩ {0,1} = {0}
|
||||||
|
self.assertEqual(len(result), 1)
|
||||||
|
self.assertIs(result[0], req0)
|
||||||
|
# req1 stays in queue
|
||||||
|
self.assertEqual(len(mgr.grammar_queue), 1)
|
||||||
|
self.assertIs(mgr.grammar_queue[0], req1)
|
||||||
|
|
||||||
|
@patch("sglang.srt.constrained.grammar_manager.torch.distributed.all_gather_object")
|
||||||
|
def test_multi_rank_sync_unions_failed(self, mock_all_gather):
|
||||||
|
"""Failed requests from any rank should be unioned."""
|
||||||
|
mgr = self._make_mgr()
|
||||||
|
mgr.grammar_sync_size = 2
|
||||||
|
mgr.SGLANG_GRAMMAR_MAX_POLL_ITERATIONS = 1
|
||||||
|
|
||||||
|
pending_future = Future() # Never completes
|
||||||
|
req = _make_req(json_schema="slow", rid="r0")
|
||||||
|
req.grammar = pending_future
|
||||||
|
req.grammar_key = ("json", "slow")
|
||||||
|
req.grammar_wait_ct = 0
|
||||||
|
|
||||||
|
mgr.grammar_queue = [req]
|
||||||
|
|
||||||
|
# Simulate: rank 0 has no ready and idx 0 failed, rank 1 has no ready/failed
|
||||||
|
def fake_all_gather(output_list, _obj, group=None): # noqa: ARG001
|
||||||
|
output_list[0] = (set(), {0}) # rank 0: idx 0 timed out
|
||||||
|
output_list[1] = (set(), set()) # rank 1: nothing
|
||||||
|
|
||||||
|
mock_all_gather.side_effect = fake_all_gather
|
||||||
|
|
||||||
|
result = mgr.get_ready_grammar_requests()
|
||||||
|
# Union of failed: {} ∪ {0} = {0}
|
||||||
|
self.assertEqual(len(result), 1)
|
||||||
|
req.set_finish_with_abort.assert_called_once()
|
||||||
|
self.assertIn("timed out", req.set_finish_with_abort.call_args[0][0])
|
||||||
|
self.assertEqual(len(mgr.grammar_queue), 0)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,414 @@
|
|||||||
|
"""
|
||||||
|
Unit tests for sglang.srt.constrained.reasoner_grammar_backend.
|
||||||
|
|
||||||
|
Test Coverage:
|
||||||
|
- ReasonerGrammarObject: state transitions, accept_token during thinking
|
||||||
|
vs post-thinking, rollback across think boundary, fill_vocab_mask gating,
|
||||||
|
copy semantics, finished delegation, delegation of jump methods
|
||||||
|
- ReasonerGrammarBackend: dispatch wrapping, invalid grammar passthrough,
|
||||||
|
None grammar passthrough, reasoning init on wrapped object
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python -m pytest test_reasoner_grammar_backend.py -v
|
||||||
|
"""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from unittest.mock import MagicMock, call
|
||||||
|
|
||||||
|
from sglang.srt.constrained.base_grammar_backend import (
|
||||||
|
BaseGrammarBackend,
|
||||||
|
BaseGrammarObject,
|
||||||
|
InvalidGrammarObject,
|
||||||
|
)
|
||||||
|
from sglang.srt.constrained.reasoner_grammar_backend import (
|
||||||
|
ReasonerGrammarBackend,
|
||||||
|
ReasonerGrammarObject,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(2.0, "stage-a-cpu-only")
|
||||||
|
|
||||||
|
THINK_END_ID = 99
|
||||||
|
|
||||||
|
|
||||||
|
class TestReasonerGrammarObjectStateTransitions(unittest.TestCase):
|
||||||
|
"""Test thinking state machine in ReasonerGrammarObject."""
|
||||||
|
|
||||||
|
def _make(self):
|
||||||
|
grammar = MagicMock(spec=BaseGrammarObject)
|
||||||
|
return ReasonerGrammarObject(grammar, THINK_END_ID), grammar
|
||||||
|
|
||||||
|
def test_initial_state_thinking(self):
|
||||||
|
obj, _ = self._make()
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, -1)
|
||||||
|
|
||||||
|
def test_transfer_state_during_thinking(self):
|
||||||
|
"""Regular tokens during thinking don't change state."""
|
||||||
|
obj, _ = self._make()
|
||||||
|
obj.transfer_state(10)
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, -1)
|
||||||
|
|
||||||
|
def test_transfer_state_think_end_token(self):
|
||||||
|
"""Think end token transitions from -1 to 0."""
|
||||||
|
obj, _ = self._make()
|
||||||
|
obj.transfer_state(THINK_END_ID)
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, 0)
|
||||||
|
|
||||||
|
def test_transfer_state_increments_after_thinking(self):
|
||||||
|
"""After thinking ends, each token increments counter."""
|
||||||
|
obj, _ = self._make()
|
||||||
|
obj.tokens_after_think_end = 0
|
||||||
|
obj.transfer_state(10)
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, 1)
|
||||||
|
obj.transfer_state(20)
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, 2)
|
||||||
|
|
||||||
|
def test_think_end_after_thinking_already_ended(self):
|
||||||
|
"""Second think_end_id after thinking ended just increments."""
|
||||||
|
obj, _ = self._make()
|
||||||
|
obj.tokens_after_think_end = 3
|
||||||
|
obj.transfer_state(THINK_END_ID)
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, 4)
|
||||||
|
|
||||||
|
def test_rollback_state_from_post_thinking(self):
|
||||||
|
obj, _ = self._make()
|
||||||
|
obj.tokens_after_think_end = 3
|
||||||
|
obj.rollback_state()
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, 2)
|
||||||
|
|
||||||
|
def test_rollback_state_at_boundary(self):
|
||||||
|
"""Rollback from 0 goes back to -1 (thinking)."""
|
||||||
|
obj, _ = self._make()
|
||||||
|
obj.tokens_after_think_end = 0
|
||||||
|
obj.rollback_state()
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, -1)
|
||||||
|
|
||||||
|
def test_rollback_state_during_thinking(self):
|
||||||
|
"""Rollback during thinking stays at -1."""
|
||||||
|
obj, _ = self._make()
|
||||||
|
obj.rollback_state()
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, -1)
|
||||||
|
|
||||||
|
|
||||||
|
class TestReasonerGrammarObjectAcceptToken(unittest.TestCase):
|
||||||
|
"""Test accept_token behavior with thinking/post-thinking states."""
|
||||||
|
|
||||||
|
def _make(self):
|
||||||
|
grammar = MagicMock(spec=BaseGrammarObject)
|
||||||
|
return ReasonerGrammarObject(grammar, THINK_END_ID), grammar
|
||||||
|
|
||||||
|
def test_accept_during_thinking_skips_grammar(self):
|
||||||
|
"""During thinking phase, inner grammar should NOT receive tokens."""
|
||||||
|
obj, grammar = self._make()
|
||||||
|
obj.accept_token(10)
|
||||||
|
grammar.accept_token.assert_not_called()
|
||||||
|
# State should still be -1
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, -1)
|
||||||
|
|
||||||
|
def test_accept_think_end_token(self):
|
||||||
|
"""Think end token transitions state but doesn't call inner grammar (state was -1 before transfer)."""
|
||||||
|
obj, grammar = self._make()
|
||||||
|
# tokens_after_think_end is -1, so grammar.accept_token is not called
|
||||||
|
# But wait: accept_token checks `>= 0` BEFORE transfer_state
|
||||||
|
# At call time tokens_after_think_end == -1, so grammar.accept_token skipped
|
||||||
|
obj.accept_token(THINK_END_ID)
|
||||||
|
grammar.accept_token.assert_not_called()
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, 0)
|
||||||
|
|
||||||
|
def test_accept_after_thinking_calls_grammar(self):
|
||||||
|
"""After thinking ends, tokens go to inner grammar."""
|
||||||
|
obj, grammar = self._make()
|
||||||
|
obj.tokens_after_think_end = 0
|
||||||
|
obj.accept_token(42)
|
||||||
|
grammar.accept_token.assert_called_once_with(42)
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, 1)
|
||||||
|
|
||||||
|
def test_accept_sequence_through_thinking_and_generation(self):
|
||||||
|
"""Full sequence: think tokens -> think_end -> generation tokens."""
|
||||||
|
obj, grammar = self._make()
|
||||||
|
|
||||||
|
# Thinking phase
|
||||||
|
obj.accept_token(1)
|
||||||
|
obj.accept_token(2)
|
||||||
|
self.assertEqual(grammar.accept_token.call_count, 0)
|
||||||
|
|
||||||
|
# Think end
|
||||||
|
obj.accept_token(THINK_END_ID)
|
||||||
|
self.assertEqual(grammar.accept_token.call_count, 0)
|
||||||
|
|
||||||
|
# Generation phase
|
||||||
|
obj.accept_token(10)
|
||||||
|
obj.accept_token(20)
|
||||||
|
self.assertEqual(grammar.accept_token.call_count, 2)
|
||||||
|
grammar.accept_token.assert_has_calls([call(10), call(20)])
|
||||||
|
|
||||||
|
|
||||||
|
class TestReasonerGrammarObjectRollback(unittest.TestCase):
|
||||||
|
"""Test rollback across thinking boundary."""
|
||||||
|
|
||||||
|
def _make(self):
|
||||||
|
grammar = MagicMock(spec=BaseGrammarObject)
|
||||||
|
return ReasonerGrammarObject(grammar, THINK_END_ID), grammar
|
||||||
|
|
||||||
|
def test_rollback_within_generation(self):
|
||||||
|
"""Rollback entirely within generation phase."""
|
||||||
|
obj, grammar = self._make()
|
||||||
|
obj.tokens_after_think_end = 5
|
||||||
|
obj.rollback(3)
|
||||||
|
grammar.rollback.assert_called_once_with(3)
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, 2)
|
||||||
|
|
||||||
|
def test_rollback_across_boundary(self):
|
||||||
|
"""Rollback that crosses from generation back into thinking."""
|
||||||
|
obj, grammar = self._make()
|
||||||
|
obj.tokens_after_think_end = 2
|
||||||
|
obj.rollback(4)
|
||||||
|
# Only 2 tokens were post-thinking, so inner grammar rolls back 2
|
||||||
|
grammar.rollback.assert_called_once_with(2)
|
||||||
|
# After 4 rollback_state calls from 2: 2->1->0->-1->-1
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, -1)
|
||||||
|
|
||||||
|
def test_rollback_during_thinking(self):
|
||||||
|
"""Rollback during thinking phase doesn't touch inner grammar."""
|
||||||
|
obj, grammar = self._make()
|
||||||
|
obj.rollback(3)
|
||||||
|
grammar.rollback.assert_not_called()
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, -1)
|
||||||
|
|
||||||
|
def test_rollback_zero(self):
|
||||||
|
obj, grammar = self._make()
|
||||||
|
obj.tokens_after_think_end = 2
|
||||||
|
obj.rollback(0)
|
||||||
|
grammar.rollback.assert_not_called()
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, 2)
|
||||||
|
|
||||||
|
def test_rollback_exactly_to_boundary(self):
|
||||||
|
"""Rollback exactly the number of post-thinking tokens."""
|
||||||
|
obj, grammar = self._make()
|
||||||
|
obj.tokens_after_think_end = 3
|
||||||
|
obj.rollback(3)
|
||||||
|
grammar.rollback.assert_called_once_with(3)
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, 0)
|
||||||
|
|
||||||
|
def test_rollback_far_beyond_all_tokens(self):
|
||||||
|
"""Rollback k much larger than tokens_after_think_end clamps grammar rollback."""
|
||||||
|
obj, grammar = self._make()
|
||||||
|
obj.tokens_after_think_end = 2
|
||||||
|
obj.rollback(100)
|
||||||
|
# Inner grammar only rolls back the 2 post-thinking tokens
|
||||||
|
grammar.rollback.assert_called_once_with(2)
|
||||||
|
# State bottoms out at -1
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, -1)
|
||||||
|
|
||||||
|
def test_accept_then_rollback_roundtrip(self):
|
||||||
|
"""Accept tokens then rollback should restore original state."""
|
||||||
|
obj, grammar = self._make()
|
||||||
|
obj.tokens_after_think_end = 0 # Just finished thinking
|
||||||
|
|
||||||
|
# Accept 3 generation tokens
|
||||||
|
obj.accept_token(10)
|
||||||
|
obj.accept_token(20)
|
||||||
|
obj.accept_token(30)
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, 3)
|
||||||
|
self.assertEqual(grammar.accept_token.call_count, 3)
|
||||||
|
|
||||||
|
# Rollback all 3
|
||||||
|
obj.rollback(3)
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, 0)
|
||||||
|
grammar.rollback.assert_called_once_with(3)
|
||||||
|
|
||||||
|
|
||||||
|
class TestReasonerGrammarObjectVocabMask(unittest.TestCase):
|
||||||
|
"""Test vocab mask gating based on thinking state."""
|
||||||
|
|
||||||
|
def _make(self):
|
||||||
|
grammar = MagicMock(spec=BaseGrammarObject)
|
||||||
|
return ReasonerGrammarObject(grammar, THINK_END_ID), grammar
|
||||||
|
|
||||||
|
def test_fill_during_thinking_skips(self):
|
||||||
|
obj, grammar = self._make()
|
||||||
|
obj.fill_vocab_mask("mask", 0)
|
||||||
|
grammar.fill_vocab_mask.assert_not_called()
|
||||||
|
|
||||||
|
def test_fill_after_thinking_delegates(self):
|
||||||
|
obj, grammar = self._make()
|
||||||
|
obj.tokens_after_think_end = 0
|
||||||
|
obj.fill_vocab_mask("mask", 0)
|
||||||
|
grammar.fill_vocab_mask.assert_called_once_with("mask", 0)
|
||||||
|
|
||||||
|
def test_fill_well_into_generation(self):
|
||||||
|
obj, grammar = self._make()
|
||||||
|
obj.tokens_after_think_end = 5
|
||||||
|
obj.fill_vocab_mask("mask", 2)
|
||||||
|
grammar.fill_vocab_mask.assert_called_once_with("mask", 2)
|
||||||
|
|
||||||
|
def test_fill_at_think_end_boundary(self):
|
||||||
|
"""After accepting think_end token, fill_vocab_mask should delegate."""
|
||||||
|
obj, grammar = self._make()
|
||||||
|
# Simulate: accept think_end, state goes from -1 to 0
|
||||||
|
obj.accept_token(THINK_END_ID)
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, 0)
|
||||||
|
obj.fill_vocab_mask("mask", 0)
|
||||||
|
grammar.fill_vocab_mask.assert_called_once_with("mask", 0)
|
||||||
|
|
||||||
|
def test_allocate_delegates(self):
|
||||||
|
obj, grammar = self._make()
|
||||||
|
obj.allocate_vocab_mask(32000, 4, "cpu")
|
||||||
|
grammar.allocate_vocab_mask.assert_called_once_with(32000, 4, "cpu")
|
||||||
|
|
||||||
|
def test_move_delegates(self):
|
||||||
|
obj, grammar = self._make()
|
||||||
|
obj.move_vocab_mask("mask", "cuda")
|
||||||
|
grammar.move_vocab_mask.assert_called_once_with("mask", "cuda")
|
||||||
|
|
||||||
|
|
||||||
|
class TestReasonerGrammarObjectDelegation(unittest.TestCase):
|
||||||
|
"""Test that non-state methods delegate to inner grammar."""
|
||||||
|
|
||||||
|
def _make(self):
|
||||||
|
grammar = MagicMock(spec=BaseGrammarObject)
|
||||||
|
return ReasonerGrammarObject(grammar, THINK_END_ID), grammar
|
||||||
|
|
||||||
|
def test_is_terminated_delegates(self):
|
||||||
|
obj, grammar = self._make()
|
||||||
|
grammar.is_terminated.return_value = True
|
||||||
|
self.assertTrue(obj.is_terminated())
|
||||||
|
|
||||||
|
def test_finished_getter_delegates(self):
|
||||||
|
obj, grammar = self._make()
|
||||||
|
grammar.finished = True
|
||||||
|
self.assertTrue(obj.finished)
|
||||||
|
|
||||||
|
def test_finished_setter_delegates(self):
|
||||||
|
obj, grammar = self._make()
|
||||||
|
obj.finished = True
|
||||||
|
self.assertTrue(grammar.finished)
|
||||||
|
|
||||||
|
def test_try_jump_forward_delegates(self):
|
||||||
|
obj, grammar = self._make()
|
||||||
|
grammar.try_jump_forward.return_value = ([1, 2], "ab")
|
||||||
|
result = obj.try_jump_forward("tokenizer")
|
||||||
|
grammar.try_jump_forward.assert_called_once_with("tokenizer")
|
||||||
|
self.assertEqual(result, ([1, 2], "ab"))
|
||||||
|
|
||||||
|
def test_jump_forward_str_state_delegates(self):
|
||||||
|
obj, grammar = self._make()
|
||||||
|
grammar.jump_forward_str_state.return_value = ("str", 5)
|
||||||
|
result = obj.jump_forward_str_state("helper")
|
||||||
|
self.assertEqual(result, ("str", 5))
|
||||||
|
|
||||||
|
def test_jump_and_retokenize_delegates(self):
|
||||||
|
obj, grammar = self._make()
|
||||||
|
obj.jump_and_retokenize([1], [2], 3)
|
||||||
|
grammar.jump_and_retokenize.assert_called_once_with([1], [2], 3)
|
||||||
|
|
||||||
|
def test_apply_vocab_mask_property(self):
|
||||||
|
obj, grammar = self._make()
|
||||||
|
grammar.apply_vocab_mask = "mask_fn"
|
||||||
|
self.assertEqual(obj.apply_vocab_mask, "mask_fn")
|
||||||
|
|
||||||
|
def test_copy_creates_new_wrapper(self):
|
||||||
|
obj, grammar = self._make()
|
||||||
|
grammar_copy = MagicMock(spec=BaseGrammarObject)
|
||||||
|
grammar.copy.return_value = grammar_copy
|
||||||
|
|
||||||
|
copied = obj.copy()
|
||||||
|
self.assertIsInstance(copied, ReasonerGrammarObject)
|
||||||
|
self.assertIsNot(copied, obj)
|
||||||
|
self.assertIs(copied.grammar, grammar_copy)
|
||||||
|
self.assertEqual(copied.think_end_id, THINK_END_ID)
|
||||||
|
|
||||||
|
def test_copy_does_not_share_state(self):
|
||||||
|
"""Modifying copy's state should not affect the original."""
|
||||||
|
obj, grammar = self._make()
|
||||||
|
grammar_copy = MagicMock(spec=BaseGrammarObject)
|
||||||
|
grammar.copy.return_value = grammar_copy
|
||||||
|
|
||||||
|
copied = obj.copy()
|
||||||
|
copied.tokens_after_think_end = 5
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, -1)
|
||||||
|
|
||||||
|
|
||||||
|
class TestReasonerGrammarObjectMaybeInitReasoning(unittest.TestCase):
|
||||||
|
"""Test maybe_init_reasoning state initialization."""
|
||||||
|
|
||||||
|
def test_reasoning_true_sets_thinking(self):
|
||||||
|
grammar = MagicMock(spec=BaseGrammarObject)
|
||||||
|
obj = ReasonerGrammarObject(grammar, THINK_END_ID)
|
||||||
|
obj.maybe_init_reasoning(True)
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, -1)
|
||||||
|
|
||||||
|
def test_reasoning_false_skips_thinking(self):
|
||||||
|
grammar = MagicMock(spec=BaseGrammarObject)
|
||||||
|
obj = ReasonerGrammarObject(grammar, THINK_END_ID)
|
||||||
|
obj.maybe_init_reasoning(False)
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, 0)
|
||||||
|
|
||||||
|
def test_reasoning_toggle(self):
|
||||||
|
"""Toggling reasoning resets state regardless of current position."""
|
||||||
|
grammar = MagicMock(spec=BaseGrammarObject)
|
||||||
|
obj = ReasonerGrammarObject(grammar, THINK_END_ID)
|
||||||
|
obj.tokens_after_think_end = 5 # Deep into generation
|
||||||
|
|
||||||
|
obj.maybe_init_reasoning(True)
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, -1)
|
||||||
|
|
||||||
|
obj.maybe_init_reasoning(False)
|
||||||
|
self.assertEqual(obj.tokens_after_think_end, 0)
|
||||||
|
|
||||||
|
|
||||||
|
class TestReasonerGrammarBackend(unittest.TestCase):
|
||||||
|
"""Test ReasonerGrammarBackend dispatch wrapping."""
|
||||||
|
|
||||||
|
def _make(self):
|
||||||
|
inner = MagicMock(spec=BaseGrammarBackend)
|
||||||
|
backend = ReasonerGrammarBackend(inner, THINK_END_ID)
|
||||||
|
return backend, inner
|
||||||
|
|
||||||
|
def test_wraps_valid_grammar(self):
|
||||||
|
backend, inner = self._make()
|
||||||
|
mock_grammar = MagicMock(spec=BaseGrammarObject)
|
||||||
|
inner._init_value_dispatch.return_value = mock_grammar
|
||||||
|
|
||||||
|
result = backend._init_value_dispatch(("json", "schema"), True)
|
||||||
|
self.assertIsInstance(result, ReasonerGrammarObject)
|
||||||
|
self.assertIs(result.grammar, mock_grammar)
|
||||||
|
self.assertEqual(result.think_end_id, THINK_END_ID)
|
||||||
|
|
||||||
|
def test_passes_through_invalid_grammar(self):
|
||||||
|
backend, inner = self._make()
|
||||||
|
invalid = InvalidGrammarObject("bad grammar")
|
||||||
|
inner._init_value_dispatch.return_value = invalid
|
||||||
|
|
||||||
|
result = backend._init_value_dispatch(("json", "schema"), False)
|
||||||
|
self.assertIs(result, invalid)
|
||||||
|
self.assertIsInstance(result, InvalidGrammarObject)
|
||||||
|
|
||||||
|
def test_passes_through_none(self):
|
||||||
|
backend, inner = self._make()
|
||||||
|
inner._init_value_dispatch.return_value = None
|
||||||
|
|
||||||
|
result = backend._init_value_dispatch(("json", "schema"), False)
|
||||||
|
self.assertIsNone(result)
|
||||||
|
|
||||||
|
def test_inits_reasoning_on_wrapped(self):
|
||||||
|
backend, inner = self._make()
|
||||||
|
mock_grammar = MagicMock(spec=BaseGrammarObject)
|
||||||
|
inner._init_value_dispatch.return_value = mock_grammar
|
||||||
|
|
||||||
|
result = backend._init_value_dispatch(("json", "schema"), True)
|
||||||
|
# reasoning=True → tokens_after_think_end should be -1
|
||||||
|
self.assertEqual(result.tokens_after_think_end, -1)
|
||||||
|
|
||||||
|
def test_inits_no_reasoning_on_wrapped(self):
|
||||||
|
backend, inner = self._make()
|
||||||
|
mock_grammar = MagicMock(spec=BaseGrammarObject)
|
||||||
|
inner._init_value_dispatch.return_value = mock_grammar
|
||||||
|
|
||||||
|
result = backend._init_value_dispatch(("json", "schema"), False)
|
||||||
|
# reasoning=False → tokens_after_think_end should be 0
|
||||||
|
self.assertEqual(result.tokens_after_think_end, 0)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,74 @@
|
|||||||
|
"""
|
||||||
|
Unit tests for sglang.srt.constrained.utils.
|
||||||
|
|
||||||
|
Test Coverage:
|
||||||
|
- is_legacy_structural_tag: legacy format detection, new format detection,
|
||||||
|
missing fields, edge cases with assertion errors.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python -m pytest test_utils.py -v
|
||||||
|
"""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.srt.constrained.utils import is_legacy_structural_tag
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(1.0, "stage-a-cpu-only")
|
||||||
|
|
||||||
|
|
||||||
|
class TestIsLegacyStructuralTag(unittest.TestCase):
|
||||||
|
"""Test is_legacy_structural_tag function."""
|
||||||
|
|
||||||
|
def test_legacy_format_returns_true(self):
|
||||||
|
obj = {
|
||||||
|
"structures": [{"begin": "<tool>", "end": "</tool>"}],
|
||||||
|
"triggers": ["<tool>"],
|
||||||
|
}
|
||||||
|
self.assertTrue(is_legacy_structural_tag(obj))
|
||||||
|
|
||||||
|
def test_legacy_format_empty_lists(self):
|
||||||
|
obj = {"structures": [], "triggers": []}
|
||||||
|
self.assertTrue(is_legacy_structural_tag(obj))
|
||||||
|
|
||||||
|
def test_new_format_returns_false(self):
|
||||||
|
obj = {"format": {"type": "json_schema", "schema": {}}}
|
||||||
|
self.assertFalse(is_legacy_structural_tag(obj))
|
||||||
|
|
||||||
|
def test_new_format_empty_format(self):
|
||||||
|
obj = {"format": {}}
|
||||||
|
self.assertFalse(is_legacy_structural_tag(obj))
|
||||||
|
|
||||||
|
def test_legacy_missing_triggers_raises(self):
|
||||||
|
"""Legacy format requires both 'structures' and 'triggers'."""
|
||||||
|
obj = {"structures": [{"begin": "<tool>", "end": "</tool>"}]}
|
||||||
|
with self.assertRaises(AssertionError):
|
||||||
|
is_legacy_structural_tag(obj)
|
||||||
|
|
||||||
|
def test_new_format_missing_format_raises(self):
|
||||||
|
"""New format (no 'structures') requires 'format' key."""
|
||||||
|
obj = {"other_key": "value"}
|
||||||
|
with self.assertRaises(AssertionError):
|
||||||
|
is_legacy_structural_tag(obj)
|
||||||
|
|
||||||
|
def test_empty_dict_raises(self):
|
||||||
|
with self.assertRaises(AssertionError):
|
||||||
|
is_legacy_structural_tag({})
|
||||||
|
|
||||||
|
def test_structures_none_uses_new_format_path(self):
|
||||||
|
"""Explicitly None 'structures' should fall to new format check."""
|
||||||
|
obj = {"structures": None, "format": {"type": "json_schema"}}
|
||||||
|
self.assertFalse(is_legacy_structural_tag(obj))
|
||||||
|
|
||||||
|
def test_both_keys_present_legacy_wins(self):
|
||||||
|
"""When both 'structures' and 'format' present, 'structures' takes priority."""
|
||||||
|
obj = {
|
||||||
|
"structures": [{"begin": "<tool>"}],
|
||||||
|
"triggers": ["<tool>"],
|
||||||
|
"format": {"type": "json_schema"},
|
||||||
|
}
|
||||||
|
self.assertTrue(is_legacy_structural_tag(obj))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user