diff --git a/test/registered/unit/constrained/test_base_grammar_backend.py b/test/registered/unit/constrained/test_base_grammar_backend.py new file mode 100644 index 000000000..7f2db0ece --- /dev/null +++ b/test/registered/unit/constrained/test_base_grammar_backend.py @@ -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() diff --git a/test/registered/unit/constrained/test_grammar_manager.py b/test/registered/unit/constrained/test_grammar_manager.py new file mode 100644 index 000000000..5ac92a5bb --- /dev/null +++ b/test/registered/unit/constrained/test_grammar_manager.py @@ -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() diff --git a/test/registered/unit/constrained/test_reasoner_grammar_backend.py b/test/registered/unit/constrained/test_reasoner_grammar_backend.py new file mode 100644 index 000000000..ba02e11e0 --- /dev/null +++ b/test/registered/unit/constrained/test_reasoner_grammar_backend.py @@ -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() diff --git a/test/registered/unit/constrained/test_utils.py b/test/registered/unit/constrained/test_utils.py new file mode 100644 index 000000000..5279c232c --- /dev/null +++ b/test/registered/unit/constrained/test_utils.py @@ -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": "", "end": ""}], + "triggers": [""], + } + 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": "", "end": ""}]} + 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": ""}], + "triggers": [""], + "format": {"type": "json_schema"}, + } + self.assertTrue(is_legacy_structural_tag(obj)) + + +if __name__ == "__main__": + unittest.main()