From 89bfeb69f9ef819ba3a7da5d942390534d873d47 Mon Sep 17 00:00:00 2001 From: Alex Bilichenko Date: Fri, 12 Jun 2026 09:09:36 +0000 Subject: [PATCH] Stop xgrammar advancement after grammar termination --- .../test_backend_xgrammar.py | 48 +++++++++++++++++++ vllm/v1/structured_output/backend_xgrammar.py | 11 ++++- 2 files changed, 58 insertions(+), 1 deletion(-) create mode 100644 tests/v1/structured_output/test_backend_xgrammar.py diff --git a/tests/v1/structured_output/test_backend_xgrammar.py b/tests/v1/structured_output/test_backend_xgrammar.py new file mode 100644 index 000000000000..405e9ee79539 --- /dev/null +++ b/tests/v1/structured_output/test_backend_xgrammar.py @@ -0,0 +1,48 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from vllm.v1.structured_output.backend_xgrammar import XgrammarGrammar + + +class _TerminatingMatcher: + def __init__(self) -> None: + self.accepted_tokens: list[int] = [] + self.terminated = False + self.rollback_count = 0 + + def accept_token(self, token: int) -> bool: + assert not self.terminated, f"accepted token after termination: {token}" + self.accepted_tokens.append(token) + if token == 1: + self.terminated = True + return True + + def is_terminated(self) -> bool: + return self.terminated + + def rollback(self, num_tokens: int) -> None: + self.rollback_count += num_tokens + del self.accepted_tokens[-num_tokens:] + self.terminated = False + + +def test_accept_tokens_stops_after_grammar_termination(): + matcher = _TerminatingMatcher() + grammar = XgrammarGrammar(vocab_size=10, matcher=matcher, ctx=None) # type: ignore[arg-type] + + assert grammar.accept_tokens("req", [1, 198]) + + assert matcher.accepted_tokens == [1] + assert grammar.num_processed_tokens == 1 + assert grammar.is_terminated() + + +def test_validate_tokens_stops_after_grammar_termination_and_rolls_back(): + matcher = _TerminatingMatcher() + grammar = XgrammarGrammar(vocab_size=10, matcher=matcher, ctx=None) # type: ignore[arg-type] + + assert grammar.validate_tokens([1, 198]) == [1] + + assert matcher.accepted_tokens == [] + assert matcher.rollback_count == 1 + assert not grammar.is_terminated() diff --git a/vllm/v1/structured_output/backend_xgrammar.py b/vllm/v1/structured_output/backend_xgrammar.py index a92be3d44320..ef4f6c9027ca 100644 --- a/vllm/v1/structured_output/backend_xgrammar.py +++ b/vllm/v1/structured_output/backend_xgrammar.py @@ -163,7 +163,11 @@ def accept_tokens(self, request_id: str, tokens: list[int]) -> bool: ) return False self.num_processed_tokens += 1 - self._is_terminated = self.matcher.is_terminated() + if self.matcher.is_terminated(): + self._is_terminated = True + break + else: + self._is_terminated = self.matcher.is_terminated() return True def validate_tokens(self, tokens: list[int]) -> list[int]: @@ -173,9 +177,14 @@ def validate_tokens(self, tokens: list[int]) -> list[int]: Returns the prefix list of tokens that are accepted by the FSM. """ accepted_tokens = [] + if self._is_terminated: + return accepted_tokens + for token in tokens: if self.matcher.accept_token(token): accepted_tokens.append(token) + if self.matcher.is_terminated(): + break else: break if len(accepted_tokens) > 0: