Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 48 additions & 0 deletions tests/v1/structured_output/test_backend_xgrammar.py
Original file line number Diff line number Diff line change
@@ -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()
11 changes: 10 additions & 1 deletion vllm/v1/structured_output/backend_xgrammar.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand All @@ -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:
Expand Down
Loading