chore: import upstream snapshot with attribution
This commit is contained in:
@@ -0,0 +1,16 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
from vllm.parser.abstract_parser import (
|
||||
DelegatingParser,
|
||||
Parser,
|
||||
)
|
||||
from vllm.parser.harmony import HarmonyParser
|
||||
from vllm.parser.parser_manager import ParserManager
|
||||
|
||||
__all__ = [
|
||||
"Parser",
|
||||
"DelegatingParser",
|
||||
"HarmonyParser",
|
||||
"ParserManager",
|
||||
]
|
||||
@@ -0,0 +1,975 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
import contextlib
|
||||
import json
|
||||
from abc import abstractmethod
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from functools import cached_property
|
||||
|
||||
from openai.types.responses import ToolChoiceFunction
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from vllm.entrypoints.chat_utils import (
|
||||
get_tool_call_id_type,
|
||||
make_tool_call_id,
|
||||
)
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import (
|
||||
ChatCompletionNamedToolChoiceParam,
|
||||
ChatCompletionRequest,
|
||||
)
|
||||
from vllm.entrypoints.openai.engine.protocol import (
|
||||
DeltaMessage,
|
||||
ExtractedToolCallInformation,
|
||||
FunctionCall,
|
||||
FunctionDefinition,
|
||||
)
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
from vllm.logger import init_logger
|
||||
from vllm.parser.metrics import record_tool_parser_invocation
|
||||
from vllm.parser.utils import count_history_tool_calls
|
||||
from vllm.reasoning.abs_reasoning_parsers import ReasoningParser
|
||||
from vllm.sampling_params import StructuredOutputsParams
|
||||
from vllm.tokenizers import TokenizerLike
|
||||
from vllm.tool_parsers.abstract_tool_parser import Tool, ToolParser
|
||||
from vllm.tool_parsers.streaming import (
|
||||
extract_named_tool_call_streaming,
|
||||
extract_required_tool_call_streaming,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class StreamState:
|
||||
"""Mutable state for ``Parser.parse_delta()``. One per stream."""
|
||||
|
||||
reasoning_ended: bool = False
|
||||
tool_call_text_started: bool = False
|
||||
prompt_reasoning_checked: bool = False
|
||||
previous_text: str = ""
|
||||
previous_token_ids: list[int] = field(default_factory=list)
|
||||
history_tool_call_cnt: int = 0
|
||||
history_tool_call_cnt_initialized: bool = False
|
||||
tool_call_id_type: str = "random"
|
||||
# only used for "required" and "named tool" choices,
|
||||
# tracks whether function name has been fully returned in the stream yet
|
||||
function_name_returned: bool = False
|
||||
engine_based: bool = False
|
||||
|
||||
def advance(
|
||||
self,
|
||||
delta_text: str,
|
||||
delta_token_ids: list[int],
|
||||
) -> tuple[str, list[int]]:
|
||||
if self.engine_based:
|
||||
return delta_text, delta_token_ids
|
||||
return (
|
||||
self.previous_text + delta_text,
|
||||
self.previous_token_ids + delta_token_ids,
|
||||
)
|
||||
|
||||
def commit(
|
||||
self,
|
||||
current_text: str,
|
||||
current_token_ids: list[int],
|
||||
) -> None:
|
||||
if self.engine_based:
|
||||
self.previous_text = ""
|
||||
self.previous_token_ids = []
|
||||
else:
|
||||
self.previous_text = current_text
|
||||
self.previous_token_ids = current_token_ids
|
||||
|
||||
|
||||
class Parser:
|
||||
"""
|
||||
Abstract Parser class that unifies ReasoningParser and ToolParser into
|
||||
a single interface for parsing model output.
|
||||
|
||||
This class provides a unified way to handle both reasoning extraction
|
||||
(e.g., chain-of-thought content in <think> tags) and tool call extraction
|
||||
(e.g., function calls in XML/JSON format) from model outputs.
|
||||
|
||||
Subclasses can either:
|
||||
1. Override the abstract methods directly for custom parsing logic
|
||||
2. Set `reasoning_parser` and `tool_parser` properties to delegate to
|
||||
existing parser implementations
|
||||
|
||||
Class Attributes:
|
||||
reasoning_parser_cls: The ReasoningParser class to use (for compatibility
|
||||
with code that needs the class, not instance).
|
||||
tool_parser_cls: The ToolParser class to use (for compatibility with
|
||||
code that needs the class, not instance).
|
||||
"""
|
||||
|
||||
# Class-level parser classes for compatibility with existing patterns
|
||||
# Subclasses should override these if they use specific parser classes
|
||||
reasoning_parser_cls: type[ReasoningParser] | None = None
|
||||
tool_parser_cls: type[ToolParser] | None = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: TokenizerLike,
|
||||
tools: list[Tool] | None = None,
|
||||
*args,
|
||||
model_config=None,
|
||||
**kwargs,
|
||||
):
|
||||
self.model_tokenizer = tokenizer
|
||||
self._reasoning_parser: ReasoningParser | None = None
|
||||
self._tool_parser: ToolParser | None = None
|
||||
if self.__class__.reasoning_parser_cls is not None:
|
||||
self._reasoning_parser = self.__class__.reasoning_parser_cls(
|
||||
tokenizer, *args, **kwargs
|
||||
)
|
||||
if self.__class__.tool_parser_cls is not None:
|
||||
self._tool_parser = self.__class__.tool_parser_cls(tokenizer, tools)
|
||||
|
||||
self._engine_based = (
|
||||
self._reasoning_parser is None
|
||||
or self._reasoning_parser.engine_based_streaming
|
||||
) and (self._tool_parser is None or self._tool_parser.engine_based_streaming)
|
||||
self._stream_state = StreamState(
|
||||
tool_call_id_type=(
|
||||
get_tool_call_id_type(model_config)
|
||||
if model_config is not None
|
||||
else "random"
|
||||
),
|
||||
engine_based=self._engine_based,
|
||||
)
|
||||
|
||||
@cached_property
|
||||
def vocab(self) -> dict[str, int]:
|
||||
"""Get the vocabulary mapping from tokens to IDs."""
|
||||
return self.model_tokenizer.get_vocab()
|
||||
|
||||
@property
|
||||
def reasoning_parser(self) -> ReasoningParser | None:
|
||||
"""The underlying reasoning parser, if any."""
|
||||
return self._reasoning_parser
|
||||
|
||||
@reasoning_parser.setter
|
||||
def reasoning_parser(self, parser: ReasoningParser | None) -> None:
|
||||
self._reasoning_parser = parser
|
||||
|
||||
@property
|
||||
def tool_parser(self) -> ToolParser | None:
|
||||
"""The underlying tool parser, if any."""
|
||||
return self._tool_parser
|
||||
|
||||
@tool_parser.setter
|
||||
def tool_parser(self, parser: ToolParser | None) -> None:
|
||||
self._tool_parser = parser
|
||||
|
||||
def _initialize_history_tool_call_cnt(
|
||||
self,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> None:
|
||||
state = self._stream_state
|
||||
if state.history_tool_call_cnt_initialized:
|
||||
return
|
||||
if state.tool_call_id_type != "kimi_k2":
|
||||
state.history_tool_call_cnt_initialized = True
|
||||
return
|
||||
state.history_tool_call_cnt = count_history_tool_calls(request)
|
||||
state.history_tool_call_cnt_initialized = True
|
||||
|
||||
# ========== Reasoning Parser Methods ==========
|
||||
|
||||
@abstractmethod
|
||||
def is_reasoning_end(self, input_ids: list[int]) -> bool:
|
||||
"""
|
||||
Check if the reasoning content ends in the input_ids.
|
||||
|
||||
Used by structured engines like `xgrammar` to check if the
|
||||
reasoning content ends in the model output.
|
||||
|
||||
Args:
|
||||
input_ids: The token IDs of the model output.
|
||||
|
||||
Returns:
|
||||
True if the reasoning content ends in the input_ids.
|
||||
"""
|
||||
|
||||
def is_reasoning_end_streaming(
|
||||
self, input_ids: list[int], delta_ids: list[int]
|
||||
) -> bool:
|
||||
"""
|
||||
Check if the reasoning content ends during a decode step.
|
||||
|
||||
Args:
|
||||
input_ids: The entire model output token IDs.
|
||||
delta_ids: The last few computed tokens at the current decode step.
|
||||
|
||||
Returns:
|
||||
True if the reasoning content ends in the delta_ids.
|
||||
"""
|
||||
return self.is_reasoning_end(input_ids)
|
||||
|
||||
@abstractmethod
|
||||
def extract_content_ids(self, input_ids: list[int]) -> list[int]:
|
||||
"""
|
||||
Extract content token IDs from the input_ids.
|
||||
|
||||
This extracts the non-reasoning content (e.g., everything after
|
||||
the </think> tag).
|
||||
|
||||
Args:
|
||||
input_ids: The token IDs of the model output.
|
||||
|
||||
Returns:
|
||||
The extracted content token IDs.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def extract_reasoning(
|
||||
self,
|
||||
model_output: str,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> tuple[str | None, str | None]:
|
||||
"""
|
||||
Extract reasoning content from a complete model-generated string.
|
||||
|
||||
Used for non-streaming responses where we have the entire model
|
||||
response available before sending to the client.
|
||||
|
||||
Args:
|
||||
model_output: The complete model-generated string.
|
||||
request: The request object used to generate the output.
|
||||
|
||||
Returns:
|
||||
A tuple of (reasoning, response_content).
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def extract_reasoning_streaming(
|
||||
self,
|
||||
previous_text: str,
|
||||
current_text: str,
|
||||
delta_text: str,
|
||||
previous_token_ids: Sequence[int],
|
||||
current_token_ids: Sequence[int],
|
||||
delta_token_ids: Sequence[int],
|
||||
) -> DeltaMessage | None:
|
||||
"""
|
||||
Extract reasoning content from a streaming delta message.
|
||||
|
||||
Args:
|
||||
previous_text: Text from all previous tokens.
|
||||
current_text: Text including the current delta.
|
||||
delta_text: The new text in this delta.
|
||||
previous_token_ids: Token IDs from previous generation.
|
||||
current_token_ids: All token IDs including current.
|
||||
delta_token_ids: The new token IDs in this delta.
|
||||
|
||||
Returns:
|
||||
A DeltaMessage with reasoning and/or content fields, or None.
|
||||
"""
|
||||
|
||||
# ========== Tool Parser Methods ==========
|
||||
|
||||
def adjust_request(
|
||||
self, request: ChatCompletionRequest | ResponsesRequest
|
||||
) -> ChatCompletionRequest | ResponsesRequest:
|
||||
"""
|
||||
Adjust the request parameters for tool calling.
|
||||
|
||||
Can be overridden by subclasses to modify request parameters
|
||||
(e.g., setting structured output schemas for tool calling).
|
||||
|
||||
Args:
|
||||
request: The original request.
|
||||
|
||||
Returns:
|
||||
The adjusted request.
|
||||
"""
|
||||
return request
|
||||
|
||||
@abstractmethod
|
||||
def extract_tool_calls(
|
||||
self,
|
||||
model_output: str,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> ExtractedToolCallInformation:
|
||||
"""
|
||||
Extract tool calls from a complete model-generated string.
|
||||
|
||||
Used for non-streaming responses.
|
||||
|
||||
Args:
|
||||
model_output: The complete model-generated string.
|
||||
request: The request object used to generate the output.
|
||||
|
||||
Returns:
|
||||
ExtractedToolCallInformation containing the tool calls.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def extract_tool_calls_streaming(
|
||||
self,
|
||||
previous_text: str,
|
||||
current_text: str,
|
||||
delta_text: str,
|
||||
previous_token_ids: Sequence[int],
|
||||
current_token_ids: Sequence[int],
|
||||
delta_token_ids: Sequence[int],
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> DeltaMessage | None:
|
||||
"""
|
||||
Extract tool calls from a streaming delta message.
|
||||
|
||||
Args:
|
||||
previous_text: Text from all previous tokens.
|
||||
current_text: Text including the current delta.
|
||||
delta_text: The new text in this delta.
|
||||
previous_token_ids: Token IDs from previous generation.
|
||||
current_token_ids: All token IDs including current.
|
||||
delta_token_ids: The new token IDs in this delta.
|
||||
request: The request object.
|
||||
|
||||
Returns:
|
||||
A DeltaMessage with tool_calls field, or None.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def parse(
|
||||
self,
|
||||
model_output: str,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
enable_auto_tools: bool = False,
|
||||
model_output_token_ids: Sequence[int] = (),
|
||||
) -> tuple[str | None, str | None, list[FunctionCall] | None]:
|
||||
"""Parse a complete model output, extracting reasoning and tool calls.
|
||||
|
||||
Args:
|
||||
model_output: The complete model-generated string.
|
||||
request: The request object used to generate the output.
|
||||
enable_auto_tools: Whether to enable automatic tool call parsing.
|
||||
model_output_token_ids: The generated raw output token IDs.
|
||||
|
||||
Returns:
|
||||
A tuple of (reasoning, content, tool_calls).
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def parse_delta(
|
||||
self,
|
||||
delta_text: str,
|
||||
delta_token_ids: list[int],
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
prompt_token_ids: list[int] | None = None,
|
||||
*,
|
||||
finished: bool,
|
||||
) -> DeltaMessage | None:
|
||||
"""Parse a single streaming delta, orchestrating reasoning then
|
||||
tool call extraction via internal stream state.
|
||||
"""
|
||||
|
||||
|
||||
class DelegatingParser(Parser):
|
||||
"""
|
||||
A Parser implementation that delegates to separate ReasoningParser and
|
||||
ToolParser instances.
|
||||
|
||||
This is the recommended base class for creating model-specific parsers
|
||||
that combine existing reasoning and tool parser implementations.
|
||||
Subclasses should set `self._reasoning_parser` and `self._tool_parser`
|
||||
in their `__init__` method.
|
||||
|
||||
If either parser is None, the corresponding methods will return default
|
||||
values (no reasoning extraction, no tool calls).
|
||||
"""
|
||||
|
||||
def extract_reasoning(
|
||||
self,
|
||||
model_output: str,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> tuple[str | None, str | None]:
|
||||
if self._reasoning_parser is None:
|
||||
return None, model_output
|
||||
return self._reasoning_parser.extract_reasoning(model_output, request)
|
||||
|
||||
def _get_function_name(
|
||||
self, request: ChatCompletionRequest | ResponsesRequest
|
||||
) -> str:
|
||||
if request.tool_choice and isinstance(request.tool_choice, ToolChoiceFunction):
|
||||
return request.tool_choice.name
|
||||
if request.tool_choice and isinstance(
|
||||
request.tool_choice, ChatCompletionNamedToolChoiceParam
|
||||
):
|
||||
return request.tool_choice.function.name
|
||||
raise ValueError("Invalid tool_choice for function name extraction.")
|
||||
|
||||
def _make_tool_call_id(self, function_name: str) -> str | None:
|
||||
state = self._stream_state
|
||||
if state.tool_call_id_type != "kimi_k2":
|
||||
return None
|
||||
tool_call_id = make_tool_call_id(
|
||||
id_type=state.tool_call_id_type,
|
||||
func_name=function_name,
|
||||
idx=state.history_tool_call_cnt,
|
||||
)
|
||||
state.history_tool_call_cnt += 1
|
||||
return tool_call_id
|
||||
|
||||
def _extract_tool_calls(
|
||||
self,
|
||||
content: str | None,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
enable_auto_tools: bool = False,
|
||||
) -> tuple[list[FunctionCall] | None, str | None]:
|
||||
tool_parser = self._tool_parser
|
||||
if tool_parser is None:
|
||||
return [], content
|
||||
|
||||
if request.tool_choice == "none":
|
||||
if self._engine_based:
|
||||
result = self.extract_tool_calls(content or "", request=request)
|
||||
return [], result.content
|
||||
return [], content
|
||||
|
||||
supports_required_and_named = tool_parser.supports_required_and_named
|
||||
is_named_tool_choice = request.tool_choice and isinstance(
|
||||
request.tool_choice,
|
||||
(ToolChoiceFunction, ChatCompletionNamedToolChoiceParam),
|
||||
)
|
||||
is_required_tool_choice = request.tool_choice == "required"
|
||||
is_auto_tool_choice = enable_auto_tools and (
|
||||
request.tool_choice == "auto"
|
||||
or request.tool_choice is None
|
||||
or (
|
||||
not supports_required_and_named
|
||||
and (is_named_tool_choice or is_required_tool_choice)
|
||||
)
|
||||
)
|
||||
|
||||
tool_calls = list[FunctionCall]()
|
||||
if is_named_tool_choice and supports_required_and_named:
|
||||
if content is None or (isinstance(content, str) and not content.strip()):
|
||||
return [], None
|
||||
function_name = self._get_function_name(request)
|
||||
tool_calls.append(
|
||||
FunctionCall(
|
||||
id=self._make_tool_call_id(function_name),
|
||||
name=function_name,
|
||||
arguments=content,
|
||||
)
|
||||
)
|
||||
content = None
|
||||
elif is_required_tool_choice and supports_required_and_named:
|
||||
# "required" with standard JSON-based parsing
|
||||
parsed_calls = []
|
||||
with contextlib.suppress(ValidationError):
|
||||
content = content or ""
|
||||
parsed_calls = TypeAdapter(list[FunctionDefinition]).validate_json(
|
||||
content
|
||||
)
|
||||
for tc in parsed_calls:
|
||||
tool_calls.append(
|
||||
FunctionCall(
|
||||
id=self._make_tool_call_id(tc.name),
|
||||
name=tc.name,
|
||||
arguments=json.dumps(tc.parameters, ensure_ascii=False),
|
||||
)
|
||||
)
|
||||
content = None
|
||||
elif is_auto_tool_choice:
|
||||
# Automatic Tool Call Parsing (also used as fallback for
|
||||
# required/named when supports_required_and_named=False)
|
||||
tool_call_info = self.extract_tool_calls(
|
||||
content if content is not None else "",
|
||||
request=request,
|
||||
)
|
||||
if tool_call_info is not None and tool_call_info.tools_called:
|
||||
tool_calls.extend(
|
||||
FunctionCall(
|
||||
id=tc.id,
|
||||
name=tc.function.name,
|
||||
arguments=tc.function.arguments,
|
||||
)
|
||||
for tc in tool_call_info.tool_calls
|
||||
)
|
||||
content = tool_call_info.content
|
||||
if content and content.strip() == "":
|
||||
content = None
|
||||
else:
|
||||
# No tool calls.
|
||||
# For required/named tool choice (when falling back to auto
|
||||
# parsing), if content is empty or whitespace-only, return
|
||||
# empty list with None content.
|
||||
if (is_required_tool_choice or is_named_tool_choice) and (
|
||||
content is None
|
||||
or (isinstance(content, str) and not content.strip())
|
||||
):
|
||||
return [], None
|
||||
return None, content
|
||||
|
||||
return tool_calls, content
|
||||
|
||||
def adjust_request(
|
||||
self, request: ChatCompletionRequest | ResponsesRequest
|
||||
) -> ChatCompletionRequest | ResponsesRequest:
|
||||
if self._reasoning_parser is not None:
|
||||
request = self._reasoning_parser.adjust_request(request)
|
||||
if self._tool_parser is not None:
|
||||
request = self._apply_structural_tag(request)
|
||||
if self._tool_parser is not None:
|
||||
request = self._tool_parser.adjust_request(request)
|
||||
return request
|
||||
|
||||
def _apply_structural_tag(
|
||||
self, request: ChatCompletionRequest | ResponsesRequest
|
||||
) -> ChatCompletionRequest | ResponsesRequest:
|
||||
if (
|
||||
self._tool_parser is None
|
||||
or self._tool_parser.structural_tag_model is None
|
||||
or not request.tools
|
||||
):
|
||||
return request
|
||||
|
||||
need_tool_calling = (
|
||||
request.tool_choice == "auto"
|
||||
or request.tool_choice == "required"
|
||||
or isinstance(
|
||||
request.tool_choice,
|
||||
(ChatCompletionNamedToolChoiceParam, ToolChoiceFunction),
|
||||
)
|
||||
)
|
||||
if not need_tool_calling:
|
||||
return request
|
||||
|
||||
structure_tag = self._tool_parser.get_structural_tag(
|
||||
request,
|
||||
reasoning=False,
|
||||
)
|
||||
if structure_tag is None:
|
||||
return request
|
||||
|
||||
structural_tag = json.dumps(structure_tag.model_dump())
|
||||
request.structured_outputs = StructuredOutputsParams(
|
||||
structural_tag=structural_tag,
|
||||
)
|
||||
if isinstance(request, ResponsesRequest):
|
||||
request.text = None
|
||||
else:
|
||||
request.response_format = None
|
||||
return request
|
||||
|
||||
def extract_reasoning_streaming(
|
||||
self,
|
||||
previous_text: str,
|
||||
current_text: str,
|
||||
delta_text: str,
|
||||
previous_token_ids: Sequence[int],
|
||||
current_token_ids: Sequence[int],
|
||||
delta_token_ids: Sequence[int],
|
||||
) -> DeltaMessage | None:
|
||||
if self._reasoning_parser is None:
|
||||
return DeltaMessage(content=delta_text)
|
||||
return self._reasoning_parser.extract_reasoning_streaming(
|
||||
previous_text,
|
||||
current_text,
|
||||
delta_text,
|
||||
previous_token_ids,
|
||||
current_token_ids,
|
||||
delta_token_ids,
|
||||
)
|
||||
|
||||
def extract_tool_calls(
|
||||
self,
|
||||
model_output: str,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> ExtractedToolCallInformation:
|
||||
if self._tool_parser is None:
|
||||
return ExtractedToolCallInformation(
|
||||
tools_called=False, tool_calls=[], content=model_output
|
||||
)
|
||||
result = None
|
||||
is_tool_called: bool | Exception = False
|
||||
try:
|
||||
result = self._tool_parser.extract_tool_calls(
|
||||
model_output,
|
||||
request=request, # type: ignore[arg-type]
|
||||
)
|
||||
is_tool_called = bool(result.tools_called)
|
||||
except Exception as e:
|
||||
is_tool_called = e
|
||||
raise
|
||||
finally:
|
||||
record_tool_parser_invocation(
|
||||
is_tool_called=is_tool_called,
|
||||
is_streaming=False,
|
||||
request=request,
|
||||
)
|
||||
return result
|
||||
|
||||
def extract_tool_calls_streaming(
|
||||
self,
|
||||
previous_text: str,
|
||||
current_text: str,
|
||||
delta_text: str,
|
||||
previous_token_ids: Sequence[int],
|
||||
current_token_ids: Sequence[int],
|
||||
delta_token_ids: Sequence[int],
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> DeltaMessage | None:
|
||||
if self._tool_parser is None:
|
||||
return None
|
||||
result = None
|
||||
is_tool_called: bool | Exception = False
|
||||
try:
|
||||
result = self._tool_parser.extract_tool_calls_streaming(
|
||||
previous_text,
|
||||
current_text,
|
||||
delta_text,
|
||||
previous_token_ids,
|
||||
current_token_ids,
|
||||
delta_token_ids,
|
||||
request, # type: ignore[arg-type]
|
||||
)
|
||||
is_tool_called = bool(result and result.tool_calls)
|
||||
except Exception as e:
|
||||
is_tool_called = e
|
||||
raise
|
||||
finally:
|
||||
record_tool_parser_invocation(
|
||||
is_tool_called=is_tool_called,
|
||||
is_streaming=True,
|
||||
request=request,
|
||||
)
|
||||
return result
|
||||
|
||||
def _extract_tool_calls_streaming(
|
||||
self,
|
||||
previous_text: str,
|
||||
current_text: str,
|
||||
delta_text: str,
|
||||
previous_token_ids: Sequence[int],
|
||||
current_token_ids: Sequence[int],
|
||||
delta_token_ids: Sequence[int],
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
# The following parameters are used for "required" tool choice parsing and are
|
||||
# tracked in StreamState for streaming parsing.
|
||||
tool_call_idx: int | None = None,
|
||||
tool_call_id_type: str = "random",
|
||||
function_name_returned: bool = False,
|
||||
) -> tuple[DeltaMessage | None, bool]:
|
||||
assert self._tool_parser is not None
|
||||
supports_required_and_named = self._tool_parser.supports_required_and_named
|
||||
|
||||
if request.tool_choice == "none":
|
||||
if self._engine_based:
|
||||
# Engine-backed parsers route content extraction through
|
||||
# extract_tool_calls_streaming, so run the full pipeline
|
||||
# and strip tool_calls after.
|
||||
delta_message = self.extract_tool_calls_streaming(
|
||||
previous_text,
|
||||
current_text,
|
||||
delta_text,
|
||||
previous_token_ids,
|
||||
current_token_ids,
|
||||
delta_token_ids,
|
||||
request, # type: ignore[arg-type]
|
||||
)
|
||||
if delta_message:
|
||||
delta_message.tool_calls = []
|
||||
return delta_message, False
|
||||
return (DeltaMessage(content=delta_text) if delta_text else None), False
|
||||
|
||||
if (
|
||||
supports_required_and_named
|
||||
and request.tool_choice
|
||||
and isinstance(
|
||||
request.tool_choice,
|
||||
(ToolChoiceFunction, ChatCompletionNamedToolChoiceParam),
|
||||
)
|
||||
):
|
||||
delta_message, function_name_returned = extract_named_tool_call_streaming(
|
||||
delta_text=delta_text,
|
||||
function_name=self._get_function_name(request),
|
||||
function_name_returned=function_name_returned,
|
||||
tool_call_idx=tool_call_idx,
|
||||
tool_call_id_type=tool_call_id_type,
|
||||
tokenizer=self.model_tokenizer,
|
||||
)
|
||||
return delta_message, function_name_returned
|
||||
|
||||
if supports_required_and_named and request.tool_choice == "required":
|
||||
delta_message, function_name_returned = (
|
||||
extract_required_tool_call_streaming(
|
||||
previous_text=previous_text,
|
||||
current_text=current_text,
|
||||
delta_text=delta_text,
|
||||
function_name_returned=function_name_returned,
|
||||
tool_call_idx=tool_call_idx,
|
||||
tool_call_id_type=tool_call_id_type,
|
||||
)
|
||||
)
|
||||
return delta_message, function_name_returned
|
||||
return self.extract_tool_calls_streaming(
|
||||
previous_text,
|
||||
current_text,
|
||||
delta_text,
|
||||
previous_token_ids,
|
||||
current_token_ids,
|
||||
delta_token_ids,
|
||||
request,
|
||||
), False
|
||||
|
||||
def is_reasoning_end(self, input_ids: list[int]) -> bool:
|
||||
if self._reasoning_parser is None:
|
||||
return False
|
||||
return self._reasoning_parser.is_reasoning_end(input_ids)
|
||||
|
||||
def is_reasoning_end_streaming(
|
||||
self, input_ids: list[int], delta_ids: list[int]
|
||||
) -> bool:
|
||||
if self._reasoning_parser is None:
|
||||
return False
|
||||
return self._reasoning_parser.is_reasoning_end_streaming(input_ids, delta_ids)
|
||||
|
||||
def extract_content_ids(self, input_ids: list[int]) -> list[int]:
|
||||
if self._reasoning_parser is None:
|
||||
return input_ids
|
||||
return self._reasoning_parser.extract_content_ids(input_ids)
|
||||
|
||||
def _in_reasoning_phase(self, state: StreamState) -> bool:
|
||||
if self._reasoning_parser is None:
|
||||
return False
|
||||
return not state.reasoning_ended
|
||||
|
||||
def _in_tool_call_phase(self, state: StreamState) -> bool:
|
||||
if self._tool_parser is None:
|
||||
return False
|
||||
return state.reasoning_ended
|
||||
|
||||
def _append_unstreamed_tool_args(
|
||||
self,
|
||||
delta_message: DeltaMessage | None,
|
||||
) -> None:
|
||||
"""Append parsed-but-unstreamed tool-call arguments to *delta_message*."""
|
||||
if (
|
||||
self._tool_parser is not None
|
||||
and delta_message
|
||||
and delta_message.tool_calls
|
||||
and (last_tc := delta_message.tool_calls[-1]).function
|
||||
):
|
||||
last_tc.function.arguments = (
|
||||
last_tc.function.arguments or ""
|
||||
) + self._tool_parser.get_remaining_unstreamed_args()
|
||||
|
||||
def finalize_generation(
|
||||
self,
|
||||
delta_message: DeltaMessage | None,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
state: StreamState,
|
||||
) -> DeltaMessage | None:
|
||||
"""Finalize generation for cases where generation was incomplete.
|
||||
For example, if streaming terminated before reasoning ended
|
||||
"""
|
||||
fallback_fn = getattr(
|
||||
self._reasoning_parser, "get_streaming_fallback_content", None
|
||||
)
|
||||
if fallback_fn is not None and not state.reasoning_ended:
|
||||
promoted = fallback_fn(state.previous_text, request)
|
||||
if promoted:
|
||||
if delta_message is None:
|
||||
delta_message = DeltaMessage()
|
||||
delta_message.content = (delta_message.content or "") + promoted
|
||||
|
||||
self._append_unstreamed_tool_args(delta_message)
|
||||
return delta_message
|
||||
|
||||
def parse(
|
||||
self,
|
||||
model_output: str,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
enable_auto_tools: bool = False,
|
||||
model_output_token_ids: Sequence[int] = (),
|
||||
) -> tuple[str | None, str | None, list[FunctionCall] | None]:
|
||||
self._initialize_history_tool_call_cnt(request)
|
||||
reasoning, content = self.extract_reasoning(model_output, request)
|
||||
tool_calls, content = self._extract_tool_calls(
|
||||
content=content,
|
||||
request=request,
|
||||
enable_auto_tools=enable_auto_tools,
|
||||
)
|
||||
return reasoning, content, tool_calls
|
||||
|
||||
def parse_delta(
|
||||
self,
|
||||
delta_text: str,
|
||||
delta_token_ids: list[int],
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
prompt_token_ids: list[int] | None = None,
|
||||
*,
|
||||
finished: bool,
|
||||
) -> DeltaMessage | None:
|
||||
self._initialize_history_tool_call_cnt(request)
|
||||
state = self._stream_state
|
||||
|
||||
if not state.prompt_reasoning_checked and prompt_token_ids is not None:
|
||||
state.prompt_reasoning_checked = True
|
||||
if self._reasoning_parser is None or self.is_reasoning_end(
|
||||
prompt_token_ids
|
||||
):
|
||||
state.reasoning_ended = True
|
||||
else:
|
||||
# Reasoning is still open at the end of the prompt; let the
|
||||
# reasoning parser adjust its initial parsing state so the
|
||||
# first generated tokens are classified correctly.
|
||||
self._reasoning_parser.adjust_initial_state_from_prompt(
|
||||
prompt_token_ids
|
||||
)
|
||||
|
||||
current_text, current_token_ids = state.advance(delta_text, delta_token_ids)
|
||||
delta_message: DeltaMessage | None = None
|
||||
reasoning_transitioned = False
|
||||
|
||||
# Reasoning extraction
|
||||
if self._in_reasoning_phase(state):
|
||||
delta_message = self.extract_reasoning_streaming(
|
||||
previous_text=state.previous_text,
|
||||
current_text=current_text,
|
||||
delta_text=delta_text,
|
||||
previous_token_ids=state.previous_token_ids,
|
||||
current_token_ids=current_token_ids,
|
||||
delta_token_ids=delta_token_ids,
|
||||
)
|
||||
reasoning_parser = self._reasoning_parser
|
||||
if reasoning_parser is not None and reasoning_parser.engine_based_streaming:
|
||||
should_transition = (
|
||||
reasoning_parser.has_engine_confirmed_reasoning_end()
|
||||
)
|
||||
else:
|
||||
should_transition = self.is_reasoning_end_streaming(
|
||||
current_token_ids, delta_token_ids
|
||||
)
|
||||
if should_transition:
|
||||
state.reasoning_ended = True
|
||||
reasoning_transitioned = True
|
||||
current_token_ids = self.extract_content_ids(delta_token_ids)
|
||||
if self._engine_based:
|
||||
flush_delta = reasoning_parser.finish_streaming() # type: ignore[union-attr, attr-defined]
|
||||
current_text = (
|
||||
(delta_message.content if delta_message else None) or ""
|
||||
) + ((flush_delta.content if flush_delta else None) or "")
|
||||
if delta_message and self._tool_parser is not None:
|
||||
delta_message.content = None
|
||||
else:
|
||||
current_text = (
|
||||
delta_message.content
|
||||
if delta_message and delta_message.content
|
||||
else ""
|
||||
)
|
||||
delta_text = current_text
|
||||
|
||||
# Tool call extraction
|
||||
if self._in_tool_call_phase(state):
|
||||
if not state.tool_call_text_started:
|
||||
state.tool_call_text_started = True
|
||||
state.previous_text = ""
|
||||
state.previous_token_ids = []
|
||||
delta_text = current_text
|
||||
delta_token_ids = current_token_ids
|
||||
|
||||
reasoning_from_this_batch = (
|
||||
delta_message.reasoning if delta_message else None
|
||||
)
|
||||
|
||||
delta_message, state.function_name_returned = (
|
||||
self._extract_tool_calls_streaming(
|
||||
previous_text=state.previous_text,
|
||||
current_text=current_text,
|
||||
delta_text=delta_text,
|
||||
previous_token_ids=state.previous_token_ids,
|
||||
current_token_ids=current_token_ids,
|
||||
delta_token_ids=delta_token_ids,
|
||||
request=request, # type: ignore[arg-type]
|
||||
tool_call_idx=state.history_tool_call_cnt,
|
||||
tool_call_id_type=state.tool_call_id_type,
|
||||
function_name_returned=state.function_name_returned,
|
||||
)
|
||||
)
|
||||
|
||||
if reasoning_from_this_batch:
|
||||
if delta_message is None:
|
||||
delta_message = DeltaMessage(reasoning=reasoning_from_this_batch)
|
||||
elif not delta_message.reasoning:
|
||||
delta_message.reasoning = reasoning_from_this_batch
|
||||
|
||||
if (
|
||||
delta_message
|
||||
and delta_message.tool_calls
|
||||
and delta_message.tool_calls[0].id is not None
|
||||
):
|
||||
state.history_tool_call_cnt += 1
|
||||
|
||||
# No phase active: pass through as content.
|
||||
# Skip when reasoning just ended in this delta — the engine already
|
||||
# consumed the end-of-reasoning marker (e.g. </think>) and
|
||||
# delta_text still contains the raw marker text.
|
||||
if (
|
||||
delta_message is None
|
||||
and not reasoning_transitioned
|
||||
and not self._in_reasoning_phase(state)
|
||||
and not self._in_tool_call_phase(state)
|
||||
):
|
||||
delta_message = DeltaMessage(content=delta_text)
|
||||
|
||||
state.commit(current_text, current_token_ids)
|
||||
|
||||
if finished:
|
||||
delta_message = self.finalize_generation(delta_message, request, state)
|
||||
delta_message = self._flush_engine_parsers(delta_message)
|
||||
|
||||
# Suppress reasoning deltas if not requested
|
||||
if delta_message and not request.include_reasoning:
|
||||
delta_message.reasoning = None
|
||||
|
||||
# If only reasoning was in the message (no content, no tool_calls)
|
||||
# skip emitting entirely
|
||||
if not delta_message.content and not delta_message.tool_calls:
|
||||
delta_message = None
|
||||
|
||||
return delta_message
|
||||
|
||||
def _flush_engine_parsers(
|
||||
self, delta_message: DeltaMessage | None
|
||||
) -> DeltaMessage | None:
|
||||
"""Flush buffered state from engine-based parsers at stream end."""
|
||||
reasoning_ended = self._stream_state.reasoning_ended
|
||||
for parser in (self._reasoning_parser, self._tool_parser):
|
||||
if not getattr(parser, "engine_based_streaming", False):
|
||||
continue
|
||||
# When reasoning has ended and we transitioned to the tool
|
||||
# phase, the reasoning parser's engine may still have buffered
|
||||
# characters from tool-call markup it saw with
|
||||
# skip_tool_parsing=True. Flushing that would leak spurious
|
||||
# content (e.g. a stray '"'), so skip it.
|
||||
if parser is self._reasoning_parser and reasoning_ended:
|
||||
continue
|
||||
finish = getattr(parser, "finish_streaming", None)
|
||||
if finish is None:
|
||||
continue
|
||||
flush_delta = finish()
|
||||
if flush_delta is None:
|
||||
continue
|
||||
if delta_message is None:
|
||||
delta_message = flush_delta
|
||||
else:
|
||||
if flush_delta.content:
|
||||
delta_message.content = (
|
||||
delta_message.content or ""
|
||||
) + flush_delta.content
|
||||
if flush_delta.reasoning:
|
||||
delta_message.reasoning = (
|
||||
delta_message.reasoning or ""
|
||||
) + flush_delta.reasoning
|
||||
if flush_delta.tool_calls:
|
||||
delta_message.tool_calls = (
|
||||
delta_message.tool_calls or []
|
||||
) + flush_delta.tool_calls
|
||||
return delta_message
|
||||
@@ -0,0 +1,131 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""DeepSeek V3.2 parser: DSML tool calls with ``function_calls`` wrapper.
|
||||
|
||||
DeepSeek V3.2 output format::
|
||||
|
||||
<|DSML|function_calls>
|
||||
<|DSML|invoke name="func_name">
|
||||
<|DSML|parameter name="location" string="true">杭州</|DSML|parameter>
|
||||
<|DSML|parameter name="count" string="false">5</|DSML|parameter>
|
||||
</|DSML|invoke>
|
||||
</|DSML|function_calls>
|
||||
|
||||
This is identical to DeepSeek V4 except for the outer wrapper
|
||||
(``function_calls`` instead of ``tool_calls``) and the absence of
|
||||
``<think>``/``</think>`` reasoning tags.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from vllm.parser.deepseek_v4 import (
|
||||
DSML_INVOKE_END,
|
||||
DSML_INVOKE_NAME_END,
|
||||
DSML_INVOKE_PREFIX,
|
||||
DSML_PARAM_CLOSE,
|
||||
_dsml_arg_converter,
|
||||
_unwrap_wrapper_args,
|
||||
)
|
||||
from vllm.parser.engine.events import EventType
|
||||
from vllm.parser.engine.parser_engine import ParserEngine
|
||||
from vllm.parser.engine.parser_engine_config import (
|
||||
ParserEngineConfig,
|
||||
ParserState,
|
||||
Transition,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.tokenizers import TokenizerLike
|
||||
from vllm.tool_parsers.abstract_tool_parser import Tool
|
||||
|
||||
_DSML = "|DSML|"
|
||||
|
||||
DSML_FUNC_START = f"<{_DSML}function_calls>"
|
||||
DSML_FUNC_END = f"</{_DSML}function_calls>"
|
||||
|
||||
|
||||
@functools.cache
|
||||
def deepseek_v32_config() -> ParserEngineConfig:
|
||||
return ParserEngineConfig(
|
||||
name="deepseek_v32",
|
||||
initial_state=ParserState.CONTENT,
|
||||
terminals={
|
||||
"TOOL_START": DSML_FUNC_START,
|
||||
"TOOL_END": DSML_FUNC_END,
|
||||
"INVOKE_PREFIX": DSML_INVOKE_PREFIX,
|
||||
"INVOKE_NAME_END": DSML_INVOKE_NAME_END,
|
||||
"INVOKE_END": DSML_INVOKE_END,
|
||||
"PARAM_CLOSE": DSML_PARAM_CLOSE,
|
||||
},
|
||||
token_id_terminals={
|
||||
"TOOL_START": DSML_FUNC_START,
|
||||
"TOOL_END": DSML_FUNC_END,
|
||||
},
|
||||
transitions={
|
||||
(ParserState.CONTENT, "TOOL_START"): Transition(
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
(),
|
||||
),
|
||||
(ParserState.TOOL_PREAMBLE, "INVOKE_PREFIX"): Transition(
|
||||
ParserState.TOOL_NAME,
|
||||
(EventType.TOOL_CALL_START,),
|
||||
),
|
||||
(ParserState.TOOL_NAME, "INVOKE_NAME_END"): Transition(
|
||||
ParserState.TOOL_ARGS,
|
||||
(),
|
||||
),
|
||||
(ParserState.TOOL_ARGS, "INVOKE_END"): Transition(
|
||||
ParserState.TOOL_BETWEEN,
|
||||
(EventType.TOOL_CALL_END,),
|
||||
),
|
||||
(ParserState.TOOL_ARGS, "TOOL_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(EventType.TOOL_CALL_END,),
|
||||
),
|
||||
# Parallel tool calls
|
||||
(ParserState.TOOL_BETWEEN, "INVOKE_PREFIX"): Transition(
|
||||
ParserState.TOOL_NAME,
|
||||
(EventType.TOOL_CALL_START,),
|
||||
),
|
||||
(ParserState.TOOL_BETWEEN, "TOOL_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(),
|
||||
),
|
||||
},
|
||||
content_events={
|
||||
ParserState.CONTENT: EventType.TEXT_CHUNK,
|
||||
ParserState.TOOL_NAME: EventType.TOOL_NAME,
|
||||
ParserState.TOOL_ARGS: EventType.ARG_VALUE_CHUNK,
|
||||
},
|
||||
arg_converter=_dsml_arg_converter,
|
||||
arg_structural_chars=frozenset(">"),
|
||||
strip_content_whitespace_with_tools=False,
|
||||
tool_args_json=False,
|
||||
)
|
||||
|
||||
|
||||
class DeepSeekV32Parser(ParserEngine):
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: TokenizerLike,
|
||||
tools: list[Tool] | None = None,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
kwargs.pop("chat_template_kwargs", None)
|
||||
super().__init__(
|
||||
tokenizer,
|
||||
tools,
|
||||
parser_engine_config=deepseek_v32_config(),
|
||||
**kwargs,
|
||||
)
|
||||
self._arg_converter = self._convert_args
|
||||
|
||||
def _convert_args(self, raw_args: str, partial: bool) -> str:
|
||||
result = _dsml_arg_converter(raw_args, partial)
|
||||
if not self._tools:
|
||||
return result
|
||||
func_name = next((s.name for s in self._tool_slots if s.args == raw_args), None)
|
||||
return _unwrap_wrapper_args(result, self._tools, func_name)
|
||||
@@ -0,0 +1,237 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""DeepSeek V4 parser: ``<think>``/``</think>``
|
||||
reasoning plus DSML tool calls in a single state machine.
|
||||
|
||||
DeepSeek V4 output format::
|
||||
|
||||
<think>
|
||||
...reasoning...
|
||||
</think>
|
||||
<|DSML|tool_calls>
|
||||
<|DSML|invoke name="func_name">
|
||||
<|DSML|parameter name="location" string="true">杭州</|DSML|parameter>
|
||||
<|DSML|parameter name="count" string="false">5</|DSML|parameter>
|
||||
</|DSML|invoke>
|
||||
</|DSML|tool_calls>
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import functools
|
||||
import json
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import regex as re
|
||||
|
||||
from vllm.parser.engine.events import EventType
|
||||
from vllm.parser.engine.parser_engine import ParserEngine
|
||||
from vllm.parser.engine.parser_engine_config import (
|
||||
ParserEngineConfig,
|
||||
ParserState,
|
||||
Transition,
|
||||
)
|
||||
from vllm.tool_parsers.utils import find_tool_properties
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.tokenizers import TokenizerLike
|
||||
from vllm.tool_parsers.abstract_tool_parser import Tool
|
||||
|
||||
_DSML = "|DSML|"
|
||||
|
||||
DSML_THINK_START = "<think>"
|
||||
DSML_THINK_END = "</think>"
|
||||
DSML_TOOL_START = f"<{_DSML}tool_calls>"
|
||||
DSML_TOOL_END = f"</{_DSML}tool_calls>"
|
||||
DSML_INVOKE_PREFIX = f'<{_DSML}invoke name="'
|
||||
DSML_INVOKE_NAME_END = '">'
|
||||
DSML_INVOKE_END = f"</{_DSML}invoke>"
|
||||
DSML_PARAM_CLOSE = f"</{_DSML}parameter>"
|
||||
|
||||
_ESCAPED_DSML = re.escape(_DSML)
|
||||
_PARAM_RE = re.compile(
|
||||
rf'<{_ESCAPED_DSML}parameter\s+name="([^"]+)"\s+string="(true|false)">'
|
||||
rf"(.*?)</{_ESCAPED_DSML}parameter>",
|
||||
re.DOTALL,
|
||||
)
|
||||
_PARTIAL_PARAM_RE = re.compile(
|
||||
rf'<{_ESCAPED_DSML}parameter\s+name="([^"]+)"\s+string="(true|false)">'
|
||||
rf"(.*)$",
|
||||
re.DOTALL,
|
||||
)
|
||||
|
||||
|
||||
def _dsml_arg_converter(raw_args: str, partial: bool) -> str:
|
||||
params: dict[str, object] = {}
|
||||
|
||||
last_end = 0
|
||||
for m in _PARAM_RE.finditer(raw_args):
|
||||
name, is_str, value = m.group(1), m.group(2), m.group(3)
|
||||
if is_str == "true":
|
||||
params[name] = value
|
||||
else:
|
||||
try:
|
||||
params[name] = json.loads(value)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
params[name] = value
|
||||
last_end = m.end()
|
||||
|
||||
if partial:
|
||||
pm = _PARTIAL_PARAM_RE.search(raw_args, last_end)
|
||||
if pm:
|
||||
name, is_str, value = pm.group(1), pm.group(2), pm.group(3)
|
||||
if is_str == "true":
|
||||
params[name] = value
|
||||
else:
|
||||
with contextlib.suppress(json.JSONDecodeError, ValueError):
|
||||
params[name] = json.loads(value)
|
||||
|
||||
return json.dumps(params, ensure_ascii=False)
|
||||
|
||||
|
||||
def _unwrap_wrapper_args(
|
||||
args_json: str,
|
||||
tools: list[Tool] | None,
|
||||
func_name: str | None,
|
||||
) -> str:
|
||||
if not tools or not func_name:
|
||||
return args_json
|
||||
try:
|
||||
args = json.loads(args_json)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return args_json
|
||||
if not isinstance(args, dict):
|
||||
return args_json
|
||||
properties = find_tool_properties(tools, func_name)
|
||||
if not properties:
|
||||
return args_json
|
||||
allowed = set(properties.keys())
|
||||
for wrapper in ("arguments", "input"):
|
||||
if set(args.keys()) != {wrapper} or wrapper in allowed:
|
||||
continue
|
||||
inner = args[wrapper]
|
||||
if isinstance(inner, str):
|
||||
try:
|
||||
inner = json.loads(inner)
|
||||
except json.JSONDecodeError:
|
||||
return args_json
|
||||
if isinstance(inner, dict) and set(inner.keys()).issubset(allowed):
|
||||
return json.dumps(inner, ensure_ascii=False)
|
||||
return args_json
|
||||
|
||||
|
||||
@functools.cache
|
||||
def deepseek_v4_config(thinking: bool = False) -> ParserEngineConfig:
|
||||
return ParserEngineConfig(
|
||||
name="deepseek_v4",
|
||||
initial_state=ParserState.REASONING if thinking else ParserState.CONTENT,
|
||||
terminals={
|
||||
"THINK_START": DSML_THINK_START,
|
||||
"THINK_END": DSML_THINK_END,
|
||||
"TOOL_START": DSML_TOOL_START,
|
||||
"TOOL_END": DSML_TOOL_END,
|
||||
"INVOKE_PREFIX": DSML_INVOKE_PREFIX,
|
||||
"INVOKE_NAME_END": DSML_INVOKE_NAME_END,
|
||||
"INVOKE_END": DSML_INVOKE_END,
|
||||
"PARAM_CLOSE": DSML_PARAM_CLOSE,
|
||||
},
|
||||
token_id_terminals={
|
||||
"THINK_START": DSML_THINK_START,
|
||||
"THINK_END": DSML_THINK_END,
|
||||
"TOOL_START": DSML_TOOL_START,
|
||||
"TOOL_END": DSML_TOOL_END,
|
||||
},
|
||||
transitions={
|
||||
(ParserState.CONTENT, "THINK_START"): Transition(
|
||||
ParserState.REASONING,
|
||||
(EventType.REASONING_START,),
|
||||
),
|
||||
# Absorb a bare </think> with no prior <think>
|
||||
(ParserState.CONTENT, "THINK_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(),
|
||||
),
|
||||
# Absorb a duplicate <think> while already reasoning
|
||||
(ParserState.REASONING, "THINK_START"): Transition(
|
||||
ParserState.REASONING,
|
||||
(),
|
||||
),
|
||||
(ParserState.REASONING, "THINK_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(EventType.REASONING_END,),
|
||||
),
|
||||
# Tool call beginning while still inside <think>
|
||||
(ParserState.REASONING, "TOOL_START"): Transition(
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
(EventType.REASONING_END,),
|
||||
),
|
||||
(ParserState.CONTENT, "TOOL_START"): Transition(
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
(),
|
||||
),
|
||||
(ParserState.TOOL_PREAMBLE, "INVOKE_PREFIX"): Transition(
|
||||
ParserState.TOOL_NAME,
|
||||
(EventType.TOOL_CALL_START,),
|
||||
),
|
||||
(ParserState.TOOL_NAME, "INVOKE_NAME_END"): Transition(
|
||||
ParserState.TOOL_ARGS,
|
||||
(),
|
||||
),
|
||||
(ParserState.TOOL_ARGS, "INVOKE_END"): Transition(
|
||||
ParserState.TOOL_BETWEEN,
|
||||
(EventType.TOOL_CALL_END,),
|
||||
),
|
||||
(ParserState.TOOL_ARGS, "TOOL_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(EventType.TOOL_CALL_END,),
|
||||
),
|
||||
# Parallel tool calls
|
||||
(ParserState.TOOL_BETWEEN, "INVOKE_PREFIX"): Transition(
|
||||
ParserState.TOOL_NAME,
|
||||
(EventType.TOOL_CALL_START,),
|
||||
),
|
||||
(ParserState.TOOL_BETWEEN, "TOOL_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(),
|
||||
),
|
||||
},
|
||||
content_events={
|
||||
ParserState.CONTENT: EventType.TEXT_CHUNK,
|
||||
ParserState.REASONING: EventType.REASONING_CHUNK,
|
||||
ParserState.TOOL_NAME: EventType.TOOL_NAME,
|
||||
ParserState.TOOL_ARGS: EventType.ARG_VALUE_CHUNK,
|
||||
},
|
||||
arg_converter=_dsml_arg_converter,
|
||||
arg_structural_chars=frozenset(">"),
|
||||
strip_content_whitespace_with_tools=False,
|
||||
tool_args_json=False,
|
||||
)
|
||||
|
||||
|
||||
class DeepSeekV4Parser(ParserEngine):
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: TokenizerLike,
|
||||
tools: list[Tool] | None = None,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
chat_kwargs = kwargs.pop("chat_template_kwargs", None) or {}
|
||||
thinking = (
|
||||
bool(chat_kwargs.get("thinking") or chat_kwargs.get("enable_thinking"))
|
||||
and chat_kwargs.get("reasoning_effort") != "none"
|
||||
)
|
||||
super().__init__(
|
||||
tokenizer,
|
||||
tools,
|
||||
parser_engine_config=deepseek_v4_config(thinking=thinking),
|
||||
**kwargs,
|
||||
)
|
||||
self._arg_converter = self._convert_args
|
||||
|
||||
def _convert_args(self, raw_args: str, partial: bool) -> str:
|
||||
result = _dsml_arg_converter(raw_args, partial)
|
||||
if not self._tools:
|
||||
return result
|
||||
func_name = next((s.name for s in self._tool_slots if s.args == raw_args), None)
|
||||
return _unwrap_wrapper_args(result, self._tools, func_name)
|
||||
@@ -0,0 +1,17 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
Streaming parser engine framework for tool call and reasoning extraction.
|
||||
|
||||
Instead of hand-rolling a parser for every model's tool-call / reasoning
|
||||
format, each format is declared as a ParserEngineConfig (terminals,
|
||||
states, and transitions) and a shared incremental engine handles
|
||||
streaming, ambiguity buffering, token-ID mapping, and delta computation.
|
||||
"""
|
||||
|
||||
from vllm.parser.engine.events import EventType, SemanticEvent
|
||||
|
||||
__all__ = [
|
||||
"EventType",
|
||||
"SemanticEvent",
|
||||
]
|
||||
@@ -0,0 +1,210 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Adapters that expose :class:`ParserEngine` through the legacy
|
||||
:class:`ReasoningParser` and :class:`ToolParser` interfaces.
|
||||
|
||||
This lets parser engines flow through the existing serving-layer code
|
||||
paths that expect separate reasoning and tool parser instances, without
|
||||
any changes to the serving layer itself.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterator, Sequence
|
||||
from contextlib import contextmanager
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from vllm.parser.engine.parser_engine_config import ParserState
|
||||
from vllm.reasoning.abs_reasoning_parsers import ReasoningParser
|
||||
from vllm.tool_parsers.abstract_tool_parser import ToolParser
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import (
|
||||
ChatCompletionRequest,
|
||||
)
|
||||
from vllm.entrypoints.openai.engine.protocol import (
|
||||
DeltaMessage,
|
||||
ExtractedToolCallInformation,
|
||||
)
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
from vllm.parser.engine.parser_engine import ParserEngine
|
||||
from vllm.tokenizers import TokenizerLike
|
||||
from vllm.tool_parsers.utils import Tool
|
||||
|
||||
|
||||
class ParserEngineReasoningAdapter(ReasoningParser):
|
||||
"""Adapts a :class:`ParserEngine` to the :class:`ReasoningParser`
|
||||
interface so parser engines can be used as reasoning parsers in the
|
||||
existing serving code.
|
||||
|
||||
Subclasses set :attr:`_parser_engine_cls` to the concrete
|
||||
:class:`ParserEngine` class.
|
||||
"""
|
||||
|
||||
_parser_engine_cls: type[ParserEngine]
|
||||
engine_based_streaming: bool = True
|
||||
|
||||
def __init__(self, tokenizer: TokenizerLike, *args, **kwargs) -> None:
|
||||
super().__init__(tokenizer, *args, **kwargs)
|
||||
self._parser_engine = self._parser_engine_cls(tokenizer, **kwargs) # type: ignore[call-arg]
|
||||
|
||||
@contextmanager
|
||||
def _skip_tool_parsing(self) -> Iterator[None]:
|
||||
saved = self._parser_engine.skip_tool_parsing
|
||||
self._parser_engine.skip_tool_parsing = True
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
self._parser_engine.skip_tool_parsing = saved
|
||||
|
||||
def is_reasoning_end(self, input_ids: Sequence[int]) -> bool:
|
||||
return self._parser_engine.is_reasoning_end(list(input_ids))
|
||||
|
||||
def adjust_initial_state_from_prompt(self, prompt_token_ids: Sequence[int]) -> None:
|
||||
self._parser_engine.adjust_initial_state_from_prompt(prompt_token_ids)
|
||||
|
||||
def extract_content_ids(self, input_ids: list[int]) -> list[int]:
|
||||
return self._parser_engine.extract_content_ids(input_ids)
|
||||
|
||||
def extract_reasoning(
|
||||
self,
|
||||
model_output: str,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> tuple[str | None, str | None]:
|
||||
with self._skip_tool_parsing():
|
||||
return self._parser_engine.extract_reasoning(model_output, request)
|
||||
|
||||
def extract_reasoning_streaming(
|
||||
self,
|
||||
previous_text: str,
|
||||
current_text: str,
|
||||
delta_text: str,
|
||||
previous_token_ids: Sequence[int],
|
||||
current_token_ids: Sequence[int],
|
||||
delta_token_ids: Sequence[int],
|
||||
) -> DeltaMessage | None:
|
||||
with self._skip_tool_parsing():
|
||||
return self._parser_engine.extract_reasoning_streaming(
|
||||
previous_text,
|
||||
current_text,
|
||||
delta_text,
|
||||
previous_token_ids,
|
||||
current_token_ids,
|
||||
delta_token_ids,
|
||||
)
|
||||
|
||||
@property
|
||||
def reasoning_start_str(self) -> str | None:
|
||||
return self._parser_engine.reasoning_start_str
|
||||
|
||||
@property
|
||||
def reasoning_end_str(self) -> str | None:
|
||||
return self._parser_engine.reasoning_end_str
|
||||
|
||||
def adjust_request(
|
||||
self,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> ChatCompletionRequest | ResponsesRequest:
|
||||
return self._parser_engine.adjust_request(request)
|
||||
|
||||
def has_engine_confirmed_reasoning_end(self) -> bool:
|
||||
return self._parser_engine.reasoning_ended
|
||||
|
||||
def finish_streaming(self) -> DeltaMessage | None:
|
||||
with self._skip_tool_parsing():
|
||||
return self._parser_engine.finish_streaming()
|
||||
|
||||
def get_streaming_fallback_content(
|
||||
self,
|
||||
text: str,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> str | None:
|
||||
return self._parser_engine.get_streaming_fallback_content(text, request)
|
||||
|
||||
def count_reasoning_tokens(self, token_ids: Sequence[int]) -> int:
|
||||
return self._parser_engine.count_reasoning_tokens(token_ids)
|
||||
|
||||
|
||||
class ParserEngineToolAdapter(ToolParser):
|
||||
"""Adapts a :class:`ParserEngine` to the :class:`ToolParser` interface.
|
||||
|
||||
:meth:`extract_tool_calls` starts the parser engine in ``CONTENT``
|
||||
state so it can parse reasoning-stripped content (i.e. the output of
|
||||
:meth:`ReasoningParser.extract_reasoning`).
|
||||
|
||||
Subclasses set :attr:`_parser_engine_cls` to the concrete
|
||||
:class:`ParserEngine` class.
|
||||
"""
|
||||
|
||||
_parser_engine_cls: type[ParserEngine]
|
||||
engine_based_streaming: bool = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: TokenizerLike,
|
||||
tools: list[Tool] | None = None,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
super().__init__(tokenizer, tools)
|
||||
self._parser_engine = self._parser_engine_cls(tokenizer, tools, **kwargs) # type: ignore[call-arg]
|
||||
|
||||
def adjust_request(
|
||||
self,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> ChatCompletionRequest | ResponsesRequest:
|
||||
request = super().adjust_request(request)
|
||||
return self._parser_engine.adjust_request(request)
|
||||
|
||||
def extract_tool_calls(
|
||||
self,
|
||||
model_output: str,
|
||||
request: ChatCompletionRequest,
|
||||
) -> ExtractedToolCallInformation:
|
||||
return self._parser_engine.extract_tool_calls_from_content(
|
||||
model_output, request
|
||||
)
|
||||
|
||||
def extract_tool_calls_streaming(
|
||||
self,
|
||||
previous_text: str,
|
||||
current_text: str,
|
||||
delta_text: str,
|
||||
previous_token_ids: Sequence[int],
|
||||
current_token_ids: Sequence[int],
|
||||
delta_token_ids: Sequence[int],
|
||||
request: ChatCompletionRequest,
|
||||
) -> DeltaMessage | None:
|
||||
engine = self._parser_engine
|
||||
engine.initialize_streaming(initial_state=ParserState.CONTENT)
|
||||
return engine.extract_tool_calls_streaming(
|
||||
previous_text,
|
||||
current_text,
|
||||
delta_text,
|
||||
previous_token_ids,
|
||||
current_token_ids,
|
||||
delta_token_ids,
|
||||
request,
|
||||
)
|
||||
|
||||
def finish_streaming(self) -> DeltaMessage | None:
|
||||
return self._parser_engine.finish_streaming()
|
||||
|
||||
|
||||
def make_adapters(
|
||||
parser_engine_cls: type[ParserEngine],
|
||||
) -> tuple[type[ParserEngineReasoningAdapter], type[ParserEngineToolAdapter]]:
|
||||
reasoning_adapter = type(
|
||||
f"{parser_engine_cls.__name__}ReasoningAdapter",
|
||||
(ParserEngineReasoningAdapter,),
|
||||
{"_parser_engine_cls": parser_engine_cls},
|
||||
)
|
||||
tool_adapter = type(
|
||||
f"{parser_engine_cls.__name__}ToolAdapter",
|
||||
(ParserEngineToolAdapter,),
|
||||
{"_parser_engine_cls": parser_engine_cls},
|
||||
)
|
||||
# Let the serving layer find the adapters and call adjust_request(),
|
||||
# which sets skip_special_tokens=False for the detokenizer.
|
||||
parser_engine_cls.reasoning_parser_cls = reasoning_adapter # type: ignore[attr-defined]
|
||||
parser_engine_cls.tool_parser_cls = tool_adapter # type: ignore[attr-defined]
|
||||
return reasoning_adapter, tool_adapter
|
||||
@@ -0,0 +1,26 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Semantic event types emitted by the streaming parser engine."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum, auto
|
||||
|
||||
|
||||
class EventType(Enum):
|
||||
TEXT_CHUNK = auto()
|
||||
REASONING_START = auto()
|
||||
REASONING_CHUNK = auto()
|
||||
REASONING_END = auto()
|
||||
TOOL_CALL_START = auto()
|
||||
TOOL_NAME = auto()
|
||||
ARG_VALUE_CHUNK = auto()
|
||||
TOOL_CALL_END = auto()
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class SemanticEvent:
|
||||
type: EventType
|
||||
value: str = ""
|
||||
tool_index: int = -1
|
||||
@@ -0,0 +1,223 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Incremental text lexer that converts text chunks into terminal
|
||||
tokens, with prefix-match buffering for ambiguous boundaries."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import regex as re
|
||||
|
||||
CONTENT_TERMINAL = "__CONTENT__"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class TerminalDef:
|
||||
name: str
|
||||
pattern: re.Pattern[str]
|
||||
is_literal: bool = False
|
||||
literal: str = ""
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class LexToken:
|
||||
terminal: str
|
||||
value: str
|
||||
|
||||
|
||||
class LexerShape:
|
||||
"""Immutable pre-computed data derived from terminal definitions.
|
||||
|
||||
Created once per :class:`ParserEngineConfig` and shared across all
|
||||
:class:`IncrementalLexer` instances that use the same config.
|
||||
"""
|
||||
|
||||
__slots__ = (
|
||||
"terminals",
|
||||
"literal_strings",
|
||||
"max_literal_len",
|
||||
"literal_first_chars",
|
||||
"has_only_literals",
|
||||
"prefix_set",
|
||||
"literals_by_first",
|
||||
)
|
||||
|
||||
def __init__(self, terminals: list[TerminalDef]) -> None:
|
||||
self.terminals = sorted(
|
||||
terminals,
|
||||
key=lambda t: (not t.is_literal, -len(t.pattern.pattern)),
|
||||
)
|
||||
literal_strings: list[tuple[str, str]] = []
|
||||
for t in self.terminals:
|
||||
if t.is_literal:
|
||||
literal_strings.append((t.literal, t.name))
|
||||
|
||||
self.literal_strings = literal_strings
|
||||
max_len = 0
|
||||
for lit, _ in literal_strings:
|
||||
if len(lit) > max_len:
|
||||
max_len = len(lit)
|
||||
self.max_literal_len = max_len
|
||||
self.literal_first_chars = frozenset(
|
||||
lit[0] for lit, _ in literal_strings if lit
|
||||
)
|
||||
self.has_only_literals = all(t.is_literal for t in terminals)
|
||||
|
||||
prefix_set: set[str] = set()
|
||||
for lit, _ in literal_strings:
|
||||
for i in range(1, len(lit)):
|
||||
prefix_set.add(lit[:i])
|
||||
self.prefix_set = frozenset(prefix_set)
|
||||
|
||||
by_first: dict[str, list[tuple[str, str]]] = {}
|
||||
for lit, name in literal_strings:
|
||||
if lit:
|
||||
by_first.setdefault(lit[0], []).append((lit, name))
|
||||
self.literals_by_first = by_first
|
||||
|
||||
|
||||
class IncrementalLexer:
|
||||
"""Converts streaming text into terminal tokens.
|
||||
|
||||
The key feature is **prefix-match buffering**: when the text in the
|
||||
buffer could be the start of a multi-character terminal (e.g.
|
||||
``"<tool_"`` that could become ``"<tool_call>"``), the lexer holds
|
||||
the text rather than emitting it. When the next chunk arrives, it
|
||||
either completes the terminal or flushes the buffered text as
|
||||
content.
|
||||
|
||||
Terminals are tried in priority order (literals first, then by
|
||||
descending priority, then by pattern length).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
terminals: list[TerminalDef] | LexerShape,
|
||||
content_terminal: str = CONTENT_TERMINAL,
|
||||
) -> None:
|
||||
if isinstance(terminals, LexerShape):
|
||||
shape = terminals
|
||||
else:
|
||||
shape = LexerShape(terminals)
|
||||
self._shape = shape
|
||||
self.terminals = shape.terminals
|
||||
self.content_terminal = content_terminal
|
||||
self.buffer = ""
|
||||
|
||||
self._literal_strings = shape.literal_strings
|
||||
self._max_literal_len = shape.max_literal_len
|
||||
self._literal_first_chars = shape.literal_first_chars
|
||||
self._has_only_literals = shape.has_only_literals
|
||||
self._prefix_set = shape.prefix_set
|
||||
self._literals_by_first = shape.literals_by_first
|
||||
|
||||
def reset(self) -> None:
|
||||
self.buffer = ""
|
||||
|
||||
def feed(self, text: str) -> list[LexToken]:
|
||||
if not self.buffer and self._has_only_literals and self._literal_first_chars:
|
||||
for ch in text:
|
||||
if ch in self._literal_first_chars:
|
||||
break
|
||||
else:
|
||||
return [LexToken(self.content_terminal, text)]
|
||||
self.buffer += text
|
||||
return self._drain()
|
||||
|
||||
def flush(self) -> list[LexToken]:
|
||||
tokens: list[LexToken] = []
|
||||
if self.buffer:
|
||||
tokens.extend(self._drain(final=True))
|
||||
if self.buffer:
|
||||
tokens.append(LexToken(self.content_terminal, self.buffer))
|
||||
self.buffer = ""
|
||||
return tokens
|
||||
|
||||
def _drain(self, *, final: bool = False) -> list[LexToken]:
|
||||
tokens: list[LexToken] = []
|
||||
first_chars = self._literal_first_chars
|
||||
content_terminal = self.content_terminal
|
||||
has_only_literals = self._has_only_literals
|
||||
literals_by_first = self._literals_by_first
|
||||
prefix_set = self._prefix_set
|
||||
|
||||
while self.buffer:
|
||||
if has_only_literals and first_chars:
|
||||
has_potential = False
|
||||
for ch in self.buffer:
|
||||
if ch in first_chars:
|
||||
has_potential = True
|
||||
break
|
||||
if not has_potential:
|
||||
tokens.append(LexToken(content_terminal, self.buffer))
|
||||
self.buffer = ""
|
||||
break
|
||||
|
||||
best_match: tuple[str, str, int] | None = None
|
||||
|
||||
first = self.buffer[0]
|
||||
for lit, name in literals_by_first.get(first, ()):
|
||||
if self.buffer.startswith(lit) and (
|
||||
best_match is None or len(lit) > best_match[2]
|
||||
):
|
||||
best_match = (name, lit, len(lit))
|
||||
|
||||
# If the current buffer is both a complete literal and the prefix
|
||||
# of a longer literal, wait for the next chunk. For example,
|
||||
# "<invoke name=" should not be emitted before the next chunk
|
||||
# proves whether this is the quoted form '<invoke name="'.
|
||||
if self.buffer in prefix_set and not final:
|
||||
if best_match is not None:
|
||||
longer_match = False
|
||||
for lit, _ in literals_by_first.get(first, ()):
|
||||
if len(lit) > best_match[2] and lit.startswith(self.buffer):
|
||||
longer_match = True
|
||||
break
|
||||
if not longer_match:
|
||||
tokens.append(LexToken(best_match[0], best_match[1]))
|
||||
self.buffer = self.buffer[best_match[2] :]
|
||||
continue
|
||||
break
|
||||
else:
|
||||
break
|
||||
|
||||
if best_match is not None:
|
||||
tokens.append(LexToken(best_match[0], best_match[1]))
|
||||
self.buffer = self.buffer[best_match[2] :]
|
||||
else:
|
||||
content_end = self._find_content_boundary()
|
||||
if content_end > 0:
|
||||
tokens.append(LexToken(content_terminal, self.buffer[:content_end]))
|
||||
self.buffer = self.buffer[content_end:]
|
||||
else:
|
||||
tokens.append(LexToken(content_terminal, self.buffer[0]))
|
||||
self.buffer = self.buffer[1:]
|
||||
|
||||
return tokens
|
||||
|
||||
def _find_content_boundary(self) -> int:
|
||||
buf = self.buffer
|
||||
n = len(buf)
|
||||
first_chars = self._literal_first_chars
|
||||
for i in range(1, n):
|
||||
if buf[i] not in first_chars:
|
||||
continue
|
||||
remaining = n - i
|
||||
for lit, _ in self._literal_strings:
|
||||
check_len = min(remaining, len(lit))
|
||||
if buf[i : i + check_len] == lit[:check_len]:
|
||||
return i
|
||||
return n
|
||||
|
||||
|
||||
def terminals_from_literals(literals: dict[str, str]) -> list[TerminalDef]:
|
||||
return [
|
||||
TerminalDef(
|
||||
name=name,
|
||||
pattern=re.compile(re.escape(lit)),
|
||||
is_literal=True,
|
||||
literal=lit,
|
||||
)
|
||||
for name, lit in literals.items()
|
||||
]
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,108 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Declarative configuration for model tool-call and reasoning formats.
|
||||
|
||||
Each model format is described by a :class:`ParserEngineConfig` that specifies:
|
||||
|
||||
* **terminals** – literal strings or regex patterns that delimit the format
|
||||
(e.g. ``<tool_call>``, ``</think>``).
|
||||
* **token_id_terminals** – terminals that should be matched by token ID
|
||||
rather than (or in addition to) text.
|
||||
* **transitions** – a state machine mapping
|
||||
``(state, terminal) → (new_state, events_to_emit)`` that drives semantic
|
||||
event generation during streaming.
|
||||
* **content_events** – what :class:`EventType` to emit for plain content
|
||||
(non-terminal text) in each state.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum, auto
|
||||
from functools import cached_property
|
||||
|
||||
from vllm.parser.engine.events import EventType
|
||||
|
||||
|
||||
class ParserState(Enum):
|
||||
CONTENT = auto()
|
||||
REASONING = auto()
|
||||
TOOL_PREAMBLE = auto()
|
||||
TOOL_NAME = auto()
|
||||
TOOL_ARGS = auto()
|
||||
TOOL_BETWEEN = auto()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Transition:
|
||||
next_state: ParserState
|
||||
events: tuple[EventType, ...] = field(default_factory=tuple)
|
||||
skip_in_token_id_mode: bool = False
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ParserEngineConfig:
|
||||
"""Declarative description of a model's tool-call / reasoning format.
|
||||
|
||||
The engine feeds terminals from the incremental lexer into the
|
||||
transition table and emits the corresponding semantic events.
|
||||
Content tokens (text between terminals) are classified by the
|
||||
current state via ``content_events``.
|
||||
"""
|
||||
|
||||
name: str
|
||||
|
||||
terminals: dict[str, str] = field(default_factory=dict)
|
||||
|
||||
token_id_terminals: dict[str, str] = field(default_factory=dict)
|
||||
|
||||
transitions: dict[tuple[ParserState, str], Transition] = field(
|
||||
default_factory=dict,
|
||||
)
|
||||
|
||||
content_events: dict[ParserState, EventType] = field(
|
||||
default_factory=lambda: {
|
||||
ParserState.CONTENT: EventType.TEXT_CHUNK,
|
||||
ParserState.REASONING: EventType.REASONING_CHUNK,
|
||||
ParserState.TOOL_NAME: EventType.TOOL_NAME,
|
||||
ParserState.TOOL_ARGS: EventType.ARG_VALUE_CHUNK,
|
||||
},
|
||||
)
|
||||
|
||||
initial_state: ParserState = ParserState.CONTENT
|
||||
|
||||
arg_converter: Callable[[str, bool], str] | None = None
|
||||
|
||||
stream_arg_deltas: bool = True
|
||||
|
||||
tool_args_json: bool = True
|
||||
|
||||
arg_structural_chars: frozenset[str] | None = None
|
||||
|
||||
# Special tokens exempt from auto-drop but not state-machine terminals.
|
||||
preserve_tokens: frozenset[str] = field(default_factory=frozenset)
|
||||
|
||||
# Prevents trailing-whitespace accumulation across multi-turn conversations.
|
||||
strip_trailing_reasoning_whitespace: bool = True
|
||||
|
||||
# Drop content that is entirely whitespace when tool calls follow.
|
||||
drop_whitespace_only_content_before_tools: bool = True
|
||||
|
||||
# .strip() content text when tool calls are present.
|
||||
strip_content_whitespace_with_tools: bool = True
|
||||
|
||||
# Reject tool calls whose names are absent from the request tools.
|
||||
validate_tool_names: bool = False
|
||||
|
||||
@cached_property
|
||||
def terminal_defs(self):
|
||||
from vllm.parser.engine.incremental_lexer import terminals_from_literals
|
||||
|
||||
return terminals_from_literals(self.terminals)
|
||||
|
||||
@cached_property
|
||||
def lexer_shape(self):
|
||||
from vllm.parser.engine.incremental_lexer import LexerShape
|
||||
|
||||
return LexerShape(self.terminal_defs)
|
||||
@@ -0,0 +1,64 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Concrete adapter classes for each registered parser engine.
|
||||
|
||||
These are created via :func:`make_adapters` and exposed as module-level
|
||||
names so that :class:`ReasoningParserManager` and
|
||||
:class:`ToolParserManager` can load them lazily.
|
||||
"""
|
||||
|
||||
from vllm.parser.deepseek_v4 import DeepSeekV4Parser
|
||||
from vllm.parser.deepseek_v32 import DeepSeekV32Parser
|
||||
from vllm.parser.engine.adapters import make_adapters
|
||||
from vllm.parser.gemma4 import Gemma4Parser
|
||||
from vllm.parser.glm47_moe import Glm47MoeParser
|
||||
from vllm.parser.kimi_k2 import KimiK2Parser
|
||||
from vllm.parser.minimax_m2 import MinimaxM2Parser
|
||||
from vllm.parser.nemotron_v3 import NemotronV3Parser
|
||||
from vllm.parser.qwen3 import Qwen3Parser
|
||||
from vllm.parser.seed_oss import SeedOssParser
|
||||
|
||||
(
|
||||
DeepSeekV32ParserReasoningAdapter,
|
||||
DeepSeekV32ParserToolAdapter,
|
||||
) = make_adapters(DeepSeekV32Parser)
|
||||
|
||||
(
|
||||
DeepSeekV4ParserReasoningAdapter,
|
||||
DeepSeekV4ParserToolAdapter,
|
||||
) = make_adapters(DeepSeekV4Parser)
|
||||
|
||||
(
|
||||
MinimaxM2ParserReasoningAdapter,
|
||||
MinimaxM2ParserToolAdapter,
|
||||
) = make_adapters(MinimaxM2Parser)
|
||||
|
||||
(
|
||||
Gemma4ParserReasoningAdapter,
|
||||
Gemma4ParserToolAdapter,
|
||||
) = make_adapters(Gemma4Parser)
|
||||
|
||||
(
|
||||
NemotronV3ParserReasoningAdapter,
|
||||
NemotronV3ParserToolAdapter,
|
||||
) = make_adapters(NemotronV3Parser)
|
||||
|
||||
(
|
||||
Qwen3ParserReasoningAdapter,
|
||||
Qwen3ParserToolAdapter,
|
||||
) = make_adapters(Qwen3Parser)
|
||||
|
||||
(
|
||||
SeedOssParserReasoningAdapter,
|
||||
SeedOssParserToolAdapter,
|
||||
) = make_adapters(SeedOssParser)
|
||||
|
||||
(
|
||||
Glm47MoeParserReasoningAdapter,
|
||||
Glm47MoeParserToolAdapter,
|
||||
) = make_adapters(Glm47MoeParser)
|
||||
|
||||
(
|
||||
KimiK2ParserReasoningAdapter,
|
||||
KimiK2ParserToolAdapter,
|
||||
) = make_adapters(KimiK2Parser)
|
||||
@@ -0,0 +1,472 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Streaming parser engine that orchestrates token ID scanning,
|
||||
incremental lexing, and state-machine-driven semantic event emission."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass
|
||||
|
||||
from vllm.parser.engine.events import EventType, SemanticEvent
|
||||
from vllm.parser.engine.incremental_lexer import (
|
||||
CONTENT_TERMINAL,
|
||||
IncrementalLexer,
|
||||
LexerShape,
|
||||
LexToken,
|
||||
TerminalDef,
|
||||
)
|
||||
from vllm.parser.engine.parser_engine_config import (
|
||||
ParserEngineConfig,
|
||||
ParserState,
|
||||
Transition,
|
||||
)
|
||||
from vllm.parser.engine.token_id_scanner import (
|
||||
DROP_TERMINAL,
|
||||
LexerInput,
|
||||
PreLexedTerminal,
|
||||
TextChunk,
|
||||
TokenIDScanner,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _DropInfo:
|
||||
lexer_shape: LexerShape
|
||||
extra_token_ids: dict[int, str]
|
||||
|
||||
|
||||
def _build_drop_info(
|
||||
config: ParserEngineConfig,
|
||||
tokenizer,
|
||||
) -> _DropInfo | None:
|
||||
try:
|
||||
special_tokens: list[str] = list(tokenizer.all_special_tokens)
|
||||
special_ids: list[int] = list(tokenizer.all_special_ids)
|
||||
except (AttributeError, NotImplementedError):
|
||||
return None
|
||||
|
||||
if not special_tokens:
|
||||
return None
|
||||
|
||||
configured_texts = (
|
||||
set(config.token_id_terminals.values())
|
||||
| set(config.terminals.values())
|
||||
| config.preserve_tokens
|
||||
)
|
||||
|
||||
extra_token_ids: dict[int, str] = {}
|
||||
drop_texts: set[str] = set()
|
||||
for text, tid in zip(special_tokens, special_ids):
|
||||
if text not in configured_texts:
|
||||
extra_token_ids[tid] = DROP_TERMINAL
|
||||
drop_texts.add(text)
|
||||
|
||||
if not drop_texts:
|
||||
return None
|
||||
|
||||
import regex as re
|
||||
|
||||
drop_terminal_defs = [
|
||||
TerminalDef(
|
||||
name=DROP_TERMINAL,
|
||||
pattern=re.compile(re.escape(text)),
|
||||
is_literal=True,
|
||||
literal=text,
|
||||
)
|
||||
for text in drop_texts
|
||||
]
|
||||
|
||||
all_terminal_defs = list(config.terminal_defs) + drop_terminal_defs
|
||||
lexer_shape = LexerShape(all_terminal_defs)
|
||||
|
||||
return _DropInfo(
|
||||
lexer_shape=lexer_shape,
|
||||
extra_token_ids=extra_token_ids,
|
||||
)
|
||||
|
||||
|
||||
class StreamingParserEngine:
|
||||
"""Consumes ``(delta_text, delta_token_ids)`` pairs and produces a
|
||||
stream of :class:`SemanticEvent` instances.
|
||||
|
||||
This is the main entry point for streaming parsing.
|
||||
Create one per request (it is stateful).
|
||||
|
||||
The pipeline is::
|
||||
|
||||
delta_text + delta_token_ids
|
||||
→ TokenIDScanner (special token pre-lexing)
|
||||
→ IncrementalLexer (text → terminal tokens with prefix buffering)
|
||||
→ State Machine (terminal → semantic events)
|
||||
→ list[SemanticEvent]
|
||||
|
||||
Usage::
|
||||
|
||||
engine = StreamingParserEngine(config, tokenizer)
|
||||
for each streaming delta:
|
||||
events = engine.feed(delta_text, delta_token_ids)
|
||||
# convert events to DeltaMessage
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: ParserEngineConfig,
|
||||
tokenizer,
|
||||
initial_state: ParserState | None = None,
|
||||
vocab: dict[str, int] | None = None,
|
||||
) -> None:
|
||||
self.config = config
|
||||
|
||||
resolved_token_ids: dict[int, str] = {}
|
||||
if tokenizer is not None:
|
||||
if vocab is None:
|
||||
vocab = tokenizer.get_vocab()
|
||||
if config.token_id_terminals:
|
||||
for terminal_name, token_text in config.token_id_terminals.items():
|
||||
tid = vocab.get(token_text)
|
||||
if tid is not None:
|
||||
resolved_token_ids[tid] = terminal_name
|
||||
|
||||
drop_info: _DropInfo | None = None
|
||||
if tokenizer is not None:
|
||||
drop_info = _build_drop_info(config, tokenizer)
|
||||
|
||||
lexer_shape = config.lexer_shape
|
||||
if drop_info is not None:
|
||||
resolved_token_ids.update(drop_info.extra_token_ids)
|
||||
lexer_shape = drop_info.lexer_shape
|
||||
|
||||
self._resolved_token_ids = resolved_token_ids
|
||||
self._has_drops = drop_info is not None
|
||||
|
||||
self._scanner = TokenIDScanner(
|
||||
resolved_token_ids,
|
||||
tokenizer,
|
||||
)
|
||||
|
||||
self._token_id_terminal_names: frozenset[str] = frozenset(
|
||||
resolved_token_ids.values()
|
||||
)
|
||||
|
||||
self._lexer = IncrementalLexer(lexer_shape, content_terminal=CONTENT_TERMINAL)
|
||||
|
||||
self._tool_terminals: frozenset[str] = frozenset(
|
||||
terminal
|
||||
for (state, terminal), tr in config.transitions.items()
|
||||
if tr.next_state in self._TOOL_STATES or state in self._TOOL_STATES
|
||||
)
|
||||
|
||||
self.skip_tool_parsing = False
|
||||
self.reset(initial_state=initial_state)
|
||||
|
||||
def _reset_args_state(self) -> None:
|
||||
self._args_buffer: str = ""
|
||||
self._args_safe_end: int = 0
|
||||
self._args_brace_depth: int = 0
|
||||
self._args_in_string: bool = False
|
||||
self._args_escape_next: bool = False
|
||||
|
||||
def reset(self, initial_state: ParserState | None = None) -> None:
|
||||
"""Reset mutable state for reuse across requests.
|
||||
|
||||
Preserves cached immutable structures (compiled terminals,
|
||||
resolved token IDs, lexer shape, token text cache) to avoid
|
||||
redundant initialization work.
|
||||
"""
|
||||
self.state = (
|
||||
initial_state if initial_state is not None else self.config.initial_state
|
||||
)
|
||||
self.tool_index = -1
|
||||
self._ever_had_token_ids = False
|
||||
# DO NOT reset skip_tool_parsing here — callers set it before
|
||||
# calling methods that trigger reset() (e.g. extract_reasoning),
|
||||
# and clearing it silently breaks non-streaming tool-call-as-
|
||||
# implicit-reasoning-end (content returns None).
|
||||
self._scanner.reset()
|
||||
self._lexer.reset()
|
||||
self._reset_args_state()
|
||||
|
||||
def feed(
|
||||
self,
|
||||
delta_text: str,
|
||||
delta_token_ids: Sequence[int],
|
||||
) -> list[SemanticEvent]:
|
||||
if delta_token_ids:
|
||||
self._ever_had_token_ids = True
|
||||
|
||||
# Fast path: skip scanner and lexer when the delta is plain
|
||||
# content with no special tokens and no terminal-starting chars.
|
||||
if (
|
||||
delta_text
|
||||
and not self._lexer.buffer
|
||||
and not self._scanner._deferred_terminals
|
||||
and self._lexer._literal_first_chars.isdisjoint(delta_text)
|
||||
):
|
||||
has_special = False
|
||||
for tid in delta_token_ids:
|
||||
if tid in self._resolved_token_ids:
|
||||
has_special = True
|
||||
break
|
||||
if not has_special:
|
||||
return self._emit_for_state(delta_text)
|
||||
|
||||
scanner_items = self._scanner.scan(delta_text, delta_token_ids)
|
||||
|
||||
if len(scanner_items) == 1 and isinstance(scanner_items[0], TextChunk):
|
||||
lex_tokens = self._lexer.feed(scanner_items[0].text)
|
||||
if len(lex_tokens) == 1 and lex_tokens[0].terminal == CONTENT_TERMINAL:
|
||||
text = lex_tokens[0].value
|
||||
return self._emit_for_state(text)
|
||||
return self._process_lex_tokens(lex_tokens)
|
||||
|
||||
return self._process_scanner_items(scanner_items)
|
||||
|
||||
def _process_scanner_items(
|
||||
self, items: Sequence[LexerInput]
|
||||
) -> list[SemanticEvent]:
|
||||
events: list[SemanticEvent] = []
|
||||
for item in items:
|
||||
if isinstance(item, PreLexedTerminal):
|
||||
events.extend(self._process_lex_tokens(self._lexer.flush()))
|
||||
events.extend(self._on_terminal(item.terminal, item.text))
|
||||
elif isinstance(item, TextChunk):
|
||||
events.extend(self._process_lex_tokens(self._lexer.feed(item.text)))
|
||||
return events
|
||||
|
||||
def finish(self) -> list[SemanticEvent]:
|
||||
events = self._process_scanner_items(self._scanner.flush_pending())
|
||||
|
||||
events.extend(self._process_lex_tokens(self._lexer.flush()))
|
||||
|
||||
if self._args_buffer:
|
||||
events.append(
|
||||
SemanticEvent(
|
||||
EventType.ARG_VALUE_CHUNK,
|
||||
value=self._args_buffer,
|
||||
tool_index=self.tool_index,
|
||||
)
|
||||
)
|
||||
self._args_buffer = ""
|
||||
self._args_safe_end = 0
|
||||
|
||||
if self.state in (
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
ParserState.TOOL_ARGS,
|
||||
ParserState.TOOL_NAME,
|
||||
ParserState.TOOL_BETWEEN,
|
||||
):
|
||||
if self.tool_index >= 0:
|
||||
events.append(
|
||||
SemanticEvent(
|
||||
EventType.TOOL_CALL_END,
|
||||
tool_index=self.tool_index,
|
||||
)
|
||||
)
|
||||
self.state = ParserState.CONTENT
|
||||
elif self.state == ParserState.REASONING:
|
||||
events.append(
|
||||
SemanticEvent(EventType.REASONING_END, tool_index=self.tool_index)
|
||||
)
|
||||
self.state = ParserState.CONTENT
|
||||
|
||||
return events
|
||||
|
||||
def parse_complete(self, text: str) -> list[SemanticEvent]:
|
||||
token_ids: list[int] = []
|
||||
events = self.feed(text, token_ids)
|
||||
events.extend(self.finish())
|
||||
return events
|
||||
|
||||
def _process_lex_tokens(self, tokens: list[LexToken]) -> list[SemanticEvent]:
|
||||
events: list[SemanticEvent] = []
|
||||
strict = self._token_id_terminal_names if self._ever_had_token_ids else None
|
||||
for tok in tokens:
|
||||
if tok.terminal == CONTENT_TERMINAL or (strict and tok.terminal in strict):
|
||||
events.extend(self._on_content(tok.value))
|
||||
else:
|
||||
events.extend(self._on_terminal(tok.terminal, tok.value))
|
||||
return events
|
||||
|
||||
_TOOL_STATES = frozenset(
|
||||
{
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
ParserState.TOOL_NAME,
|
||||
ParserState.TOOL_ARGS,
|
||||
ParserState.TOOL_BETWEEN,
|
||||
}
|
||||
)
|
||||
|
||||
def _on_terminal(self, terminal: str, value: str) -> list[SemanticEvent]:
|
||||
key = (self.state, terminal)
|
||||
transition = self.config.transitions.get(key)
|
||||
|
||||
if transition is None:
|
||||
if (
|
||||
self._has_drops
|
||||
and terminal == DROP_TERMINAL
|
||||
# Preserve drop tokens when skip_tool_parsing is active so
|
||||
# the reasoning pass doesn't silently remove tokens that a
|
||||
# later tool-call pass might need to see.
|
||||
and not self.skip_tool_parsing
|
||||
):
|
||||
return []
|
||||
return self._emit_for_state(value)
|
||||
|
||||
if self.skip_tool_parsing and terminal in self._tool_terminals:
|
||||
if EventType.REASONING_END in transition.events:
|
||||
self.state = ParserState.CONTENT
|
||||
return [
|
||||
SemanticEvent(
|
||||
EventType.REASONING_END,
|
||||
value=value,
|
||||
tool_index=self.tool_index,
|
||||
),
|
||||
SemanticEvent(
|
||||
EventType.TEXT_CHUNK,
|
||||
value=value,
|
||||
tool_index=self.tool_index,
|
||||
),
|
||||
]
|
||||
content_type = self.config.content_events.get(self.state)
|
||||
if content_type is not None:
|
||||
return [
|
||||
SemanticEvent(content_type, value=value, tool_index=self.tool_index)
|
||||
]
|
||||
return []
|
||||
|
||||
if transition.skip_in_token_id_mode and self._ever_had_token_ids:
|
||||
return self._emit_for_state(value)
|
||||
|
||||
return self._apply_transition(transition, value)
|
||||
|
||||
def _emit_for_state(self, text: str) -> list[SemanticEvent]:
|
||||
if self.state == ParserState.TOOL_ARGS:
|
||||
if self.config.tool_args_json:
|
||||
return self._feed_args_text(text)
|
||||
return [
|
||||
SemanticEvent(
|
||||
EventType.ARG_VALUE_CHUNK,
|
||||
value=text,
|
||||
tool_index=self.tool_index,
|
||||
)
|
||||
]
|
||||
content_type = self.config.content_events.get(self.state)
|
||||
if content_type is not None:
|
||||
return [SemanticEvent(content_type, value=text, tool_index=self.tool_index)]
|
||||
return []
|
||||
|
||||
def _on_content(self, text: str) -> list[SemanticEvent]:
|
||||
if not text:
|
||||
return []
|
||||
return self._emit_for_state(text)
|
||||
|
||||
def _apply_transition(
|
||||
self,
|
||||
transition: Transition,
|
||||
value: str,
|
||||
) -> list[SemanticEvent]:
|
||||
events: list[SemanticEvent] = []
|
||||
|
||||
if (
|
||||
self.state == ParserState.TOOL_ARGS
|
||||
and transition.next_state != ParserState.TOOL_ARGS
|
||||
and self._args_buffer
|
||||
):
|
||||
events.append(
|
||||
SemanticEvent(
|
||||
EventType.ARG_VALUE_CHUNK,
|
||||
value=self._args_buffer,
|
||||
tool_index=self.tool_index,
|
||||
)
|
||||
)
|
||||
self._args_buffer = ""
|
||||
|
||||
self.state = transition.next_state
|
||||
|
||||
for event_type in transition.events:
|
||||
if event_type == EventType.TOOL_CALL_START:
|
||||
self.tool_index += 1
|
||||
events.append(
|
||||
SemanticEvent(
|
||||
event_type,
|
||||
value=value,
|
||||
tool_index=self.tool_index,
|
||||
)
|
||||
)
|
||||
|
||||
if self.state == ParserState.TOOL_ARGS:
|
||||
self._args_brace_depth = 0
|
||||
self._args_in_string = False
|
||||
self._args_escape_next = False
|
||||
self._args_safe_end = 0
|
||||
|
||||
return events
|
||||
|
||||
def _feed_args_text(self, text: str) -> list[SemanticEvent]:
|
||||
"""Feed text into the JSON argument streaming buffer.
|
||||
|
||||
Streams argument characters incrementally while holding back
|
||||
closing braces/brackets that might change as more input arrives.
|
||||
"""
|
||||
events: list[SemanticEvent] = []
|
||||
for ch in text:
|
||||
result = self._feed_args_char(ch)
|
||||
events.extend(result)
|
||||
return events
|
||||
|
||||
def _feed_args_char(self, ch: str) -> list[SemanticEvent]:
|
||||
self._args_buffer += ch
|
||||
|
||||
if self._args_escape_next:
|
||||
self._args_escape_next = False
|
||||
self._args_safe_end = len(self._args_buffer)
|
||||
return self._flush_safe_args()
|
||||
|
||||
if self._args_in_string:
|
||||
if ch == "\\":
|
||||
self._args_escape_next = True
|
||||
elif ch == '"':
|
||||
self._args_in_string = False
|
||||
self._args_safe_end = len(self._args_buffer)
|
||||
return self._flush_safe_args()
|
||||
|
||||
if ch == '"':
|
||||
self._args_in_string = True
|
||||
self._args_safe_end = len(self._args_buffer)
|
||||
return self._flush_safe_args()
|
||||
|
||||
if ch in ("{", "["):
|
||||
self._args_brace_depth += 1
|
||||
self._args_safe_end = len(self._args_buffer)
|
||||
return self._flush_safe_args()
|
||||
|
||||
if ch in ("}", "]"):
|
||||
if self._args_brace_depth > 0:
|
||||
self._args_brace_depth -= 1
|
||||
if self._args_brace_depth == 0:
|
||||
return []
|
||||
self._args_safe_end = len(self._args_buffer)
|
||||
return self._flush_safe_args()
|
||||
|
||||
self._args_safe_end = len(self._args_buffer)
|
||||
return self._flush_safe_args()
|
||||
|
||||
def _flush_safe_args(self) -> list[SemanticEvent]:
|
||||
"""Emit buffered argument characters up to the safe-end watermark.
|
||||
|
||||
Top-level closing braces are held back (safe_end not advanced)
|
||||
until confirmed safe by a subsequent character or finish().
|
||||
"""
|
||||
if self._args_safe_end == 0:
|
||||
return []
|
||||
to_emit = self._args_buffer[: self._args_safe_end]
|
||||
self._args_buffer = self._args_buffer[self._args_safe_end :]
|
||||
self._args_safe_end = 0
|
||||
return [
|
||||
SemanticEvent(
|
||||
EventType.ARG_VALUE_CHUNK,
|
||||
value=to_emit,
|
||||
tool_index=self.tool_index,
|
||||
)
|
||||
]
|
||||
@@ -0,0 +1,296 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Scan delta token IDs for special tokens and split the stream into
|
||||
pre-lexed terminals and plain text chunks."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass
|
||||
|
||||
DROP_TERMINAL = "__DROP__"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class TextChunk:
|
||||
text: str
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class PreLexedTerminal:
|
||||
terminal: str
|
||||
token_id: int
|
||||
text: str
|
||||
|
||||
|
||||
LexerInput = TextChunk | PreLexedTerminal
|
||||
|
||||
|
||||
class TokenIDScanner:
|
||||
"""Maps special token IDs in the delta to terminals.
|
||||
|
||||
Before text-based lexing happens, the scanner checks each token ID
|
||||
in the delta against a mapping of ``{token_id: terminal_name}``.
|
||||
Matched tokens are emitted as :class:`PreLexedTerminal` items;
|
||||
everything else is grouped into :class:`TextChunk` items for the
|
||||
incremental lexer to process.
|
||||
|
||||
When a terminal's text is not yet in ``delta_text`` (held back by
|
||||
the detokenizer), the terminal is deferred until the text arrives
|
||||
in a subsequent delta.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
token_id_to_terminal: dict[int, str],
|
||||
tokenizer,
|
||||
) -> None:
|
||||
self.token_id_to_terminal = token_id_to_terminal
|
||||
self.tokenizer = tokenizer
|
||||
self._token_text_cache: dict[int, str] = {}
|
||||
self._deferred_terminals: list[PreLexedTerminal] = []
|
||||
self._deferred_post_text: str = ""
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Clear mutable state for reuse. Preserves the token text cache."""
|
||||
self._deferred_terminals.clear()
|
||||
self._deferred_post_text = ""
|
||||
|
||||
def _decode_token(self, token_id: int) -> str:
|
||||
if token_id not in self._token_text_cache:
|
||||
self._token_text_cache[token_id] = self.tokenizer.decode([token_id])
|
||||
return self._token_text_cache[token_id]
|
||||
|
||||
_EMPTY: tuple[LexerInput, ...] = ()
|
||||
|
||||
def scan(
|
||||
self,
|
||||
delta_text: str,
|
||||
delta_token_ids: Sequence[int],
|
||||
) -> Sequence[LexerInput]:
|
||||
prefix_items: list[LexerInput] = []
|
||||
effective_text = delta_text
|
||||
|
||||
if self._deferred_terminals:
|
||||
prefix_items, effective_text = self._resolve_deferred(delta_text)
|
||||
|
||||
if not self.token_id_to_terminal:
|
||||
if effective_text:
|
||||
prefix_items.append(TextChunk(effective_text))
|
||||
return prefix_items
|
||||
|
||||
has_special = False
|
||||
token_id_to_terminal = self.token_id_to_terminal
|
||||
for tid in delta_token_ids:
|
||||
if tid in token_id_to_terminal:
|
||||
has_special = True
|
||||
break
|
||||
|
||||
if not has_special:
|
||||
if effective_text:
|
||||
if not prefix_items:
|
||||
return [TextChunk(effective_text)]
|
||||
prefix_items.append(TextChunk(effective_text))
|
||||
return prefix_items or self._EMPTY
|
||||
|
||||
token_texts = [self._decode_token(tid) for tid in delta_token_ids]
|
||||
|
||||
results: list[LexerInput] = []
|
||||
text_accum: list[str] = []
|
||||
|
||||
for idx, tid in enumerate(delta_token_ids):
|
||||
terminal = self.token_id_to_terminal.get(tid)
|
||||
if terminal is not None:
|
||||
if text_accum:
|
||||
joined = "".join(text_accum)
|
||||
if joined:
|
||||
results.append(TextChunk(joined))
|
||||
text_accum.clear()
|
||||
results.append(PreLexedTerminal(terminal, tid, token_texts[idx]))
|
||||
else:
|
||||
text_accum.append(token_texts[idx])
|
||||
|
||||
if text_accum:
|
||||
joined = "".join(text_accum)
|
||||
if joined:
|
||||
results.append(TextChunk(joined))
|
||||
|
||||
if effective_text:
|
||||
results = self._recover_holdback_text(effective_text, results)
|
||||
else:
|
||||
# No detokenizer text to validate against — individually-decoded
|
||||
# TextChunks are unreliable (context-dependent decoding).
|
||||
# Defer PreLexedTerminals so the state machine doesn't
|
||||
# transition before the preceding text has arrived. The
|
||||
# deferred terminals will be resolved against the actual
|
||||
# delta_text in a subsequent scan() or flushed by finish().
|
||||
for r in results:
|
||||
if isinstance(r, PreLexedTerminal):
|
||||
self._deferred_terminals.append(r)
|
||||
results = []
|
||||
|
||||
return prefix_items + results
|
||||
|
||||
def flush_pending(self) -> list[LexerInput]:
|
||||
if not self._deferred_terminals and not self._deferred_post_text:
|
||||
return []
|
||||
results: list[LexerInput] = []
|
||||
if self._deferred_post_text:
|
||||
results.append(TextChunk(self._deferred_post_text))
|
||||
self._deferred_post_text = ""
|
||||
results.extend(self._deferred_terminals)
|
||||
self._deferred_terminals.clear()
|
||||
return results
|
||||
|
||||
def _resolve_deferred(
|
||||
self,
|
||||
delta_text: str,
|
||||
) -> tuple[list[LexerInput], str]:
|
||||
"""Resolve deferred terminals against new delta_text.
|
||||
|
||||
When a previous ``scan()`` deferred a terminal (its text hadn't
|
||||
arrived yet), the next delta's text should contain that terminal's
|
||||
text. Split delta_text at the terminal boundary: text before
|
||||
belongs to the previous parser state, the terminal triggers the
|
||||
state transition, and text after belongs to the new state.
|
||||
|
||||
Returns ``(prefix_items, remaining_text)`` where prefix_items
|
||||
are the resolved deferred terminals (with any preceding text)
|
||||
and remaining_text is the unconsumed portion of delta_text that
|
||||
should be scanned with the current delta's token IDs.
|
||||
"""
|
||||
deferred = self._deferred_terminals
|
||||
self._deferred_terminals = []
|
||||
|
||||
results: list[LexerInput] = []
|
||||
remaining = delta_text
|
||||
|
||||
if self._deferred_post_text:
|
||||
remaining = self._deferred_post_text + remaining
|
||||
self._deferred_post_text = ""
|
||||
|
||||
# Duplicate-text deferred terminals resolve left-to-right via
|
||||
# find(); correct when each terminal text appears once in sequence.
|
||||
for terminal in deferred:
|
||||
pos = remaining.find(terminal.text)
|
||||
if pos > 0:
|
||||
results.append(TextChunk(remaining[:pos]))
|
||||
results.append(terminal)
|
||||
remaining = remaining[pos + len(terminal.text) :]
|
||||
elif pos == 0:
|
||||
results.append(terminal)
|
||||
remaining = remaining[len(terminal.text) :]
|
||||
else:
|
||||
# Accumulate text until terminal text arrives —
|
||||
# only the terminal provides a reliable split point.
|
||||
if remaining:
|
||||
self._deferred_post_text += remaining
|
||||
remaining = ""
|
||||
self._deferred_terminals.append(terminal)
|
||||
|
||||
return results, remaining
|
||||
|
||||
def _recover_holdback_text(
|
||||
self,
|
||||
delta_text: str,
|
||||
results: list[LexerInput],
|
||||
) -> list[LexerInput]:
|
||||
"""Recover detokenizer hold-back text not in delta_token_ids.
|
||||
|
||||
The detokenizer may flush previously held-back text in
|
||||
``delta_text`` that has no corresponding token ID in
|
||||
``delta_token_ids``. This hold-back text always appears as a
|
||||
prefix of ``delta_text``.
|
||||
"""
|
||||
if not results:
|
||||
return [TextChunk(delta_text)]
|
||||
|
||||
reconstructed = self._join_decoded_text(results)
|
||||
|
||||
if not reconstructed:
|
||||
return [TextChunk(delta_text)] + results
|
||||
|
||||
pos = delta_text.find(reconstructed)
|
||||
if pos > 0:
|
||||
return [TextChunk(delta_text[:pos])] + results
|
||||
if pos == 0:
|
||||
return results
|
||||
|
||||
# Fallback: SentencePiece context-dependent decoding mismatch.
|
||||
# Rebuild from delta_text using PreLexedTerminals as split anchors.
|
||||
return self._rebuild_from_anchors(delta_text, results)
|
||||
|
||||
def _join_decoded_text(self, results: list[LexerInput]) -> str:
|
||||
"""Join TextChunk and PreLexedTerminal text into one string."""
|
||||
parts: list[str] = []
|
||||
for item in results:
|
||||
if isinstance(item, (TextChunk, PreLexedTerminal)):
|
||||
parts.append(item.text)
|
||||
return "".join(parts)
|
||||
|
||||
def _rebuild_from_anchors(
|
||||
self,
|
||||
delta_text: str,
|
||||
results: list[LexerInput],
|
||||
) -> list[LexerInput]:
|
||||
"""Rebuild results from delta_text using terminals as anchors.
|
||||
|
||||
When context-dependent decoding creates a mismatch between
|
||||
individually-decoded tokens and delta_text, use
|
||||
PreLexedTerminals as split points and reallocate text from
|
||||
delta_text. If a terminal's text is not found in delta_text,
|
||||
it is deferred to the next scan() call.
|
||||
|
||||
Anchors are resolved right-to-left with ``rfind`` so that each
|
||||
anchor binds to the *rightmost* available occurrence of its
|
||||
text. This prevents earlier literal lookalikes (e.g. a user
|
||||
mentioning ``<tool_call>`` in prose) from stealing the position
|
||||
of a real special-token anchor that appears later.
|
||||
|
||||
If the same anchor text appears multiple times as real special
|
||||
tokens (not prose), the rightmost-first binding could misalign.
|
||||
In practice this doesn't happen: each special token ID maps to
|
||||
a distinct PreLexedTerminal, and duplicates in prose are resolved
|
||||
by the token-ID filtering layer above.
|
||||
"""
|
||||
anchors = [item for item in results if isinstance(item, PreLexedTerminal)]
|
||||
if not anchors:
|
||||
return [TextChunk(delta_text)]
|
||||
|
||||
# Resolve positions right-to-left: each anchor gets the
|
||||
# rightmost occurrence that is still before the next anchor.
|
||||
positions: list[int] = [-1] * len(anchors)
|
||||
search_end = len(delta_text)
|
||||
for i in range(len(anchors) - 1, -1, -1):
|
||||
pos = delta_text.rfind(anchors[i].text, 0, search_end)
|
||||
if pos >= 0:
|
||||
positions[i] = pos
|
||||
search_end = pos
|
||||
|
||||
# Build results left-to-right using the resolved positions.
|
||||
new_results: list[LexerInput] = []
|
||||
consumed = 0
|
||||
for i, anchor in enumerate(anchors):
|
||||
pos = positions[i]
|
||||
if pos >= consumed:
|
||||
if pos > consumed:
|
||||
new_results.append(TextChunk(delta_text[consumed:pos]))
|
||||
new_results.append(anchor)
|
||||
consumed = pos + len(anchor.text)
|
||||
else:
|
||||
has_later_valid = any(p >= 0 for p in positions[i + 1 :])
|
||||
# DROP anchors (EOS, etc.) may have text that never
|
||||
# arrives in delta_text (stripped by detokenizer).
|
||||
# Don't defer remaining content waiting for text
|
||||
# that will never come.
|
||||
if (
|
||||
not has_later_valid
|
||||
and consumed < len(delta_text)
|
||||
and anchor.terminal != DROP_TERMINAL
|
||||
):
|
||||
self._deferred_post_text += delta_text[consumed:]
|
||||
consumed = len(delta_text)
|
||||
self._deferred_terminals.append(anchor)
|
||||
if consumed < len(delta_text):
|
||||
new_results.append(TextChunk(delta_text[consumed:]))
|
||||
return new_results
|
||||
@@ -0,0 +1,555 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Gemma4 parser.
|
||||
|
||||
Handles channel-based reasoning plus custom tool call format in a single
|
||||
state machine::
|
||||
|
||||
<|channel>thought
|
||||
...reasoning...<channel|>
|
||||
<|tool_call>call:func_name{key:<|"|>value<|"|>,num:42}<tool_call|>
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from vllm.entrypoints.openai.engine.protocol import DeltaMessage
|
||||
from vllm.logger import init_logger
|
||||
from vllm.parser.engine.events import EventType, SemanticEvent
|
||||
from vllm.parser.engine.parser_engine import ParserEngine
|
||||
from vllm.parser.engine.parser_engine_config import (
|
||||
ParserEngineConfig,
|
||||
ParserState,
|
||||
Transition,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import (
|
||||
ChatCompletionRequest,
|
||||
)
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
from vllm.tokenizers import TokenizerLike
|
||||
from vllm.tool_parsers.abstract_tool_parser import Tool
|
||||
|
||||
CHANNEL_START = "<|channel>"
|
||||
CHANNEL_END = "<channel|>"
|
||||
TOOL_CALL_START = "<|tool_call>"
|
||||
TOOL_CALL_END = "<tool_call|>"
|
||||
STRING_DELIM = '<|"|>'
|
||||
_DELIM_LEN = len(STRING_DELIM)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Gemma4 argument parser
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_PARTIAL_DELIM_SUFFIXES = tuple(
|
||||
STRING_DELIM[:k] for k in range(len(STRING_DELIM), 0, -1)
|
||||
)
|
||||
|
||||
|
||||
def _strip_partial_delim(value: str) -> str:
|
||||
"""Strip a trailing partial ``STRING_DELIM`` prefix from *value*.
|
||||
|
||||
Prevents partial delimiters from leaking into the streamed JSON diff.
|
||||
"""
|
||||
for suffix in _PARTIAL_DELIM_SUFFIXES:
|
||||
if value.endswith(suffix):
|
||||
return value[: -len(suffix)]
|
||||
return value
|
||||
|
||||
|
||||
def _parse_gemma4_args(args_str: str, *, partial: bool = False) -> dict:
|
||||
"""Parse Gemma4's custom key:value format into a Python dict.
|
||||
|
||||
Format examples::
|
||||
|
||||
location:<|"|>Tokyo<|"|>
|
||||
location:<|"|>San Francisco<|"|>,unit:<|"|>celsius<|"|>
|
||||
count:42,flag:true
|
||||
nested:{inner_key:<|"|>val<|"|>}
|
||||
items:[<|"|>a<|"|>,<|"|>b<|"|>]
|
||||
|
||||
Args:
|
||||
args_str: The raw Gemma4 argument string.
|
||||
partial: When True (streaming), bare values at end of string are
|
||||
omitted because they may be incomplete and type-unstable
|
||||
(e.g. partial boolean parsed as bare string).
|
||||
|
||||
Returns a dict ready for ``json.dumps()``.
|
||||
"""
|
||||
if not args_str or not args_str.strip():
|
||||
return {}
|
||||
|
||||
result: dict = {}
|
||||
i = 0
|
||||
n = len(args_str)
|
||||
|
||||
while i < n:
|
||||
while i < n and args_str[i] in (" ", ",", "\n", "\t"):
|
||||
i += 1
|
||||
if i >= n:
|
||||
break
|
||||
|
||||
key_start = i
|
||||
while i < n and args_str[i] != ":":
|
||||
i += 1
|
||||
if i >= n:
|
||||
break
|
||||
key = args_str[key_start:i].strip()
|
||||
if key.startswith(STRING_DELIM) and key.endswith(STRING_DELIM):
|
||||
key = key[_DELIM_LEN:-_DELIM_LEN]
|
||||
i += 1
|
||||
|
||||
if i >= n:
|
||||
if not partial:
|
||||
result[key] = ""
|
||||
break
|
||||
|
||||
while i < n and args_str[i] in (" ", "\n", "\t"):
|
||||
i += 1
|
||||
if i >= n:
|
||||
if not partial:
|
||||
result[key] = ""
|
||||
break
|
||||
|
||||
if args_str[i : i + _DELIM_LEN] == STRING_DELIM:
|
||||
i += _DELIM_LEN
|
||||
val_start = i
|
||||
end_pos = args_str.find(STRING_DELIM, i)
|
||||
if end_pos == -1:
|
||||
# Unterminated string — take rest, strip partial delimiter.
|
||||
value = args_str[val_start:]
|
||||
if partial:
|
||||
value = _strip_partial_delim(value)
|
||||
result[key] = value
|
||||
break
|
||||
result[key] = args_str[val_start:end_pos]
|
||||
i = end_pos + _DELIM_LEN
|
||||
|
||||
elif args_str[i] == "{":
|
||||
depth = 1
|
||||
obj_start = i + 1
|
||||
i += 1
|
||||
while i < n and depth > 0:
|
||||
if args_str[i : i + _DELIM_LEN] == STRING_DELIM:
|
||||
# Skip over string contents to avoid counting { inside strings
|
||||
i += _DELIM_LEN
|
||||
next_delim = args_str.find(STRING_DELIM, i)
|
||||
i = n if next_delim == -1 else next_delim + _DELIM_LEN
|
||||
continue
|
||||
if args_str[i] == "{":
|
||||
depth += 1
|
||||
elif args_str[i] == "}":
|
||||
depth -= 1
|
||||
i += 1
|
||||
if depth > 0:
|
||||
# Incomplete nested object — use i (not i-1) to avoid
|
||||
# dropping the last char, and recurse as partial.
|
||||
result[key] = _parse_gemma4_args(args_str[obj_start:i], partial=True)
|
||||
else:
|
||||
result[key] = _parse_gemma4_args(args_str[obj_start : i - 1])
|
||||
|
||||
elif args_str[i] == "[":
|
||||
depth = 1
|
||||
arr_start = i + 1
|
||||
i += 1
|
||||
while i < n and depth > 0:
|
||||
if args_str[i : i + _DELIM_LEN] == STRING_DELIM:
|
||||
i += _DELIM_LEN
|
||||
next_delim = args_str.find(STRING_DELIM, i)
|
||||
i = n if next_delim == -1 else next_delim + _DELIM_LEN
|
||||
continue
|
||||
if args_str[i] == "[":
|
||||
depth += 1
|
||||
elif args_str[i] == "]":
|
||||
depth -= 1
|
||||
i += 1
|
||||
if depth > 0:
|
||||
result[key] = _parse_gemma4_array(args_str[arr_start:i], partial=True)
|
||||
else:
|
||||
result[key] = _parse_gemma4_array(args_str[arr_start : i - 1])
|
||||
|
||||
else:
|
||||
val_start = i
|
||||
while i < n and args_str[i] not in (",", "}", "]"):
|
||||
i += 1
|
||||
if partial and i >= n:
|
||||
# Value may be incomplete (e.g. partial boolean) —
|
||||
# withhold to avoid type instability during streaming.
|
||||
break
|
||||
if i == val_start:
|
||||
logger.warning(
|
||||
"Gemma4 args parser made no progress at position %d; "
|
||||
"aborting on malformed input.",
|
||||
i,
|
||||
)
|
||||
break
|
||||
raw_val = args_str[val_start:i].strip()
|
||||
if partial and raw_val.endswith("."):
|
||||
# Digits may still arrive (e.g. "108." -> "108.2");
|
||||
# withhold to avoid corrupting the streaming diff.
|
||||
break
|
||||
result[key] = raw_val
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def _parse_gemma4_array(arr_str: str, *, partial: bool = False) -> list:
|
||||
items: list = []
|
||||
i = 0
|
||||
n = len(arr_str)
|
||||
|
||||
while i < n:
|
||||
while i < n and arr_str[i] in (" ", ",", "\n", "\t"):
|
||||
i += 1
|
||||
if i >= n:
|
||||
break
|
||||
|
||||
if arr_str[i : i + _DELIM_LEN] == STRING_DELIM:
|
||||
i += _DELIM_LEN
|
||||
end_pos = arr_str.find(STRING_DELIM, i)
|
||||
if end_pos == -1:
|
||||
items.append(arr_str[i:])
|
||||
break
|
||||
items.append(arr_str[i:end_pos])
|
||||
i = end_pos + _DELIM_LEN
|
||||
|
||||
elif arr_str[i] == "{":
|
||||
depth = 1
|
||||
obj_start = i + 1
|
||||
i += 1
|
||||
while i < n and depth > 0:
|
||||
if arr_str[i : i + _DELIM_LEN] == STRING_DELIM:
|
||||
i += _DELIM_LEN
|
||||
nd = arr_str.find(STRING_DELIM, i)
|
||||
i = nd + _DELIM_LEN if nd != -1 else n
|
||||
continue
|
||||
if arr_str[i] == "{":
|
||||
depth += 1
|
||||
elif arr_str[i] == "}":
|
||||
depth -= 1
|
||||
i += 1
|
||||
if depth > 0:
|
||||
items.append(_parse_gemma4_args(arr_str[obj_start:i], partial=True))
|
||||
else:
|
||||
items.append(_parse_gemma4_args(arr_str[obj_start : i - 1]))
|
||||
|
||||
elif arr_str[i] == "[":
|
||||
depth = 1
|
||||
sub_start = i + 1
|
||||
i += 1
|
||||
while i < n and depth > 0:
|
||||
if arr_str[i : i + _DELIM_LEN] == STRING_DELIM:
|
||||
i += _DELIM_LEN
|
||||
nd = arr_str.find(STRING_DELIM, i)
|
||||
i = nd + _DELIM_LEN if nd != -1 else n
|
||||
continue
|
||||
if arr_str[i] == "[":
|
||||
depth += 1
|
||||
elif arr_str[i] == "]":
|
||||
depth -= 1
|
||||
i += 1
|
||||
if depth > 0:
|
||||
items.append(_parse_gemma4_array(arr_str[sub_start:i], partial=True))
|
||||
else:
|
||||
items.append(_parse_gemma4_array(arr_str[sub_start : i - 1]))
|
||||
|
||||
else:
|
||||
val_start = i
|
||||
while i < n and arr_str[i] not in (",", "]"):
|
||||
i += 1
|
||||
if partial and i >= n:
|
||||
break
|
||||
if i == val_start:
|
||||
logger.warning(
|
||||
"Gemma4 array parser made no progress at position %d; "
|
||||
"aborting on malformed input.",
|
||||
i,
|
||||
)
|
||||
break
|
||||
raw_val = arr_str[val_start:i].strip()
|
||||
if partial and raw_val.endswith("."):
|
||||
break
|
||||
items.append(raw_val)
|
||||
|
||||
return items
|
||||
|
||||
|
||||
def _gemma4_arg_converter(raw_args: str, partial: bool) -> str:
|
||||
"""Convert Gemma4 custom arg format to a JSON string."""
|
||||
text = raw_args.strip()
|
||||
if text.endswith("}"):
|
||||
text = text[:-1]
|
||||
|
||||
parsed = _parse_gemma4_args(text, partial=partial)
|
||||
return json.dumps(parsed, ensure_ascii=False)
|
||||
|
||||
|
||||
@functools.cache
|
||||
def gemma4_config() -> ParserEngineConfig:
|
||||
return ParserEngineConfig(
|
||||
name="gemma4",
|
||||
initial_state=ParserState.CONTENT,
|
||||
terminals={
|
||||
"THINK_START": CHANNEL_START,
|
||||
"THINK_END": CHANNEL_END,
|
||||
"TOOL_START": TOOL_CALL_START,
|
||||
"TOOL_END": TOOL_CALL_END,
|
||||
"CALL_PREFIX": "call:",
|
||||
"OPEN_BRACE": "{",
|
||||
},
|
||||
token_id_terminals={
|
||||
"THINK_START": CHANNEL_START,
|
||||
"THINK_END": CHANNEL_END,
|
||||
"TOOL_START": TOOL_CALL_START,
|
||||
"TOOL_END": TOOL_CALL_END,
|
||||
},
|
||||
transitions={
|
||||
# -- Reasoning transitions --
|
||||
(ParserState.CONTENT, "THINK_START"): Transition(
|
||||
ParserState.REASONING,
|
||||
(EventType.REASONING_START,),
|
||||
),
|
||||
# No-op: if we pre-initialised the engine to REASONING from the
|
||||
# prompt (see ``adjust_initial_state_from_prompt``) but the model
|
||||
# still emits its own ``<|channel>`` opener, swallow it instead
|
||||
# of leaking it as TEXT_CHUNK.
|
||||
(ParserState.REASONING, "THINK_START"): Transition(
|
||||
ParserState.REASONING,
|
||||
(),
|
||||
),
|
||||
(ParserState.REASONING, "THINK_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(EventType.REASONING_END,),
|
||||
),
|
||||
# Tool call directly from reasoning (no explicit <channel|>)
|
||||
(ParserState.REASONING, "TOOL_START"): Transition(
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
(EventType.REASONING_END, EventType.TOOL_CALL_START),
|
||||
),
|
||||
# -- Tool call transitions --
|
||||
(ParserState.CONTENT, "TOOL_START"): Transition(
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
(EventType.REASONING_END, EventType.TOOL_CALL_START),
|
||||
),
|
||||
(ParserState.TOOL_PREAMBLE, "TOOL_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(EventType.TOOL_CALL_END,),
|
||||
),
|
||||
(ParserState.TOOL_PREAMBLE, "CALL_PREFIX"): Transition(
|
||||
ParserState.TOOL_NAME,
|
||||
(),
|
||||
),
|
||||
(ParserState.TOOL_NAME, "OPEN_BRACE"): Transition(
|
||||
ParserState.TOOL_ARGS,
|
||||
(),
|
||||
),
|
||||
(ParserState.TOOL_ARGS, "TOOL_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(EventType.TOOL_CALL_END,),
|
||||
),
|
||||
# Back-to-back tool calls
|
||||
(ParserState.CONTENT, "TOOL_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(),
|
||||
),
|
||||
# Absorb a bare <channel|> that arrives after we already
|
||||
# returned to CONTENT; prevents leaking it as TEXT_CHUNK.
|
||||
(ParserState.CONTENT, "THINK_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(),
|
||||
),
|
||||
},
|
||||
content_events={
|
||||
ParserState.CONTENT: EventType.TEXT_CHUNK,
|
||||
ParserState.REASONING: EventType.REASONING_CHUNK,
|
||||
ParserState.TOOL_NAME: EventType.TOOL_NAME,
|
||||
ParserState.TOOL_ARGS: EventType.ARG_VALUE_CHUNK,
|
||||
},
|
||||
arg_converter=_gemma4_arg_converter,
|
||||
tool_args_json=False,
|
||||
arg_structural_chars=frozenset(",:{}[]<"),
|
||||
preserve_tokens=frozenset({STRING_DELIM}),
|
||||
)
|
||||
|
||||
|
||||
_GEMMA4_THOUGHT_PREFIX = "thought\n"
|
||||
_GEMMA4_THOUGHT_TOKEN = "thought"
|
||||
|
||||
|
||||
class Gemma4Parser(ParserEngine):
|
||||
"""Gemma4 parser: ``<|channel>`` reasoning + ``<|tool_call>``
|
||||
tool calls in a single engine.
|
||||
|
||||
- Strips the ``thought\\n`` prefix from reasoning content
|
||||
- Sets ``skip_special_tokens=False`` so boundary tokens are visible
|
||||
- Detects ``<|tool_call>`` token as implicit reasoning end
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: TokenizerLike,
|
||||
tools: list[Tool] | None = None,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
chat_kwargs = kwargs.get("chat_template_kwargs", {}) or {}
|
||||
self._thinking_enabled = chat_kwargs.get("enable_thinking", True)
|
||||
super().__init__(
|
||||
tokenizer,
|
||||
tools,
|
||||
parser_engine_config=gemma4_config(),
|
||||
**kwargs,
|
||||
)
|
||||
vocab = self.vocab
|
||||
self._tool_call_token_id: int | None = vocab.get("<|tool_call>")
|
||||
self._new_turn_token_id: int | None = vocab.get("<|turn>")
|
||||
self._tool_response_token_id: int | None = vocab.get("<|tool_response>")
|
||||
self._reasoning_text: str = ""
|
||||
self._prefix_stripped: bool = False
|
||||
self._is_first_feed: bool = True
|
||||
|
||||
def _reset(self, initial_state=None) -> None:
|
||||
super()._reset(initial_state=initial_state)
|
||||
self._reasoning_text = ""
|
||||
self._prefix_stripped = False
|
||||
self._is_first_feed = True
|
||||
|
||||
def _preprocess_feed(
|
||||
self,
|
||||
delta_text: str,
|
||||
delta_token_ids: Sequence[int],
|
||||
) -> tuple[str, Sequence[int]]:
|
||||
if not self._is_first_feed:
|
||||
return delta_text, delta_token_ids
|
||||
self._is_first_feed = False
|
||||
|
||||
if (
|
||||
not delta_text
|
||||
or self._engine.state != ParserState.CONTENT
|
||||
or self._reasoning_start_token_id is None
|
||||
or self._reasoning_end_token_id is None
|
||||
):
|
||||
return delta_text, delta_token_ids
|
||||
|
||||
if CHANNEL_START in delta_text:
|
||||
return delta_text, delta_token_ids
|
||||
|
||||
needs_injection = (
|
||||
CHANNEL_END in delta_text
|
||||
or delta_text.startswith(_GEMMA4_THOUGHT_PREFIX)
|
||||
or delta_text == _GEMMA4_THOUGHT_TOKEN
|
||||
)
|
||||
if not needs_injection:
|
||||
return delta_text, delta_token_ids
|
||||
|
||||
delta_text = CHANNEL_START + delta_text
|
||||
if delta_token_ids:
|
||||
delta_token_ids = [self._reasoning_start_token_id, *delta_token_ids]
|
||||
|
||||
return delta_text, delta_token_ids
|
||||
|
||||
def is_reasoning_end(self, input_ids: list[int]) -> bool:
|
||||
end_id = self._reasoning_end_token_id
|
||||
start_id = self._reasoning_start_token_id
|
||||
tool_call_id = self._tool_call_token_id
|
||||
new_turn_id = self._new_turn_token_id
|
||||
tool_response_id = self._tool_response_token_id
|
||||
|
||||
if end_id is not None and not input_ids:
|
||||
return self.parser_engine_config.initial_state != ParserState.REASONING
|
||||
|
||||
for i in range(len(input_ids) - 1, -1, -1):
|
||||
tid = input_ids[i]
|
||||
if start_id is not None and tid == start_id:
|
||||
return False
|
||||
if tool_call_id is not None and tid == tool_call_id:
|
||||
return True
|
||||
if new_turn_id is not None and tid == new_turn_id:
|
||||
return not self._thinking_enabled
|
||||
if tool_response_id is not None and tid == tool_response_id:
|
||||
return not self._thinking_enabled
|
||||
if end_id is not None and tid == end_id:
|
||||
return True
|
||||
return True
|
||||
|
||||
def adjust_initial_state_from_prompt(self, prompt_token_ids: Sequence[int]) -> None:
|
||||
"""Pre-initialise the engine to ``REASONING`` when the prompt does
|
||||
not already end with reasoning concluded.
|
||||
|
||||
This covers the post-tool-response continuation case where the chat
|
||||
template leaves the prompt ending inside an open ``<|channel>``
|
||||
block (issue #45834). It is also safe in the common new-turn case
|
||||
where the model itself emits ``<|channel>`` first: the no-op
|
||||
``(REASONING, THINK_START)`` transition swallows it, and the
|
||||
``thought\n`` prefix in the first reasoning chunk is stripped by
|
||||
``_events_to_delta`` as it already is in the default flow.
|
||||
"""
|
||||
if self.is_reasoning_end(list(prompt_token_ids)):
|
||||
return
|
||||
self._engine.reset(initial_state=ParserState.REASONING)
|
||||
# Prevent a later default ``initialize_streaming()`` (e.g. from
|
||||
# ``ParserEngineReasoningAdapter.extract_reasoning_streaming``) from
|
||||
# clobbering this with ``CONTENT``.
|
||||
self._streaming_initialized = True
|
||||
|
||||
def _events_to_delta(
|
||||
self,
|
||||
events: list[SemanticEvent],
|
||||
finished: bool = False,
|
||||
) -> DeltaMessage | None:
|
||||
delta = super()._events_to_delta(events, finished=finished)
|
||||
if delta is None or delta.reasoning is None:
|
||||
return delta
|
||||
|
||||
if self._prefix_stripped:
|
||||
return delta
|
||||
self._reasoning_text += delta.reasoning
|
||||
|
||||
if self._reasoning_text.startswith(_GEMMA4_THOUGHT_PREFIX):
|
||||
prefix_len = len(_GEMMA4_THOUGHT_PREFIX)
|
||||
prev_reasoning_len = len(self._reasoning_text) - len(delta.reasoning)
|
||||
if prev_reasoning_len >= prefix_len:
|
||||
self._prefix_stripped = True
|
||||
return delta
|
||||
chars_of_prefix_in_delta = prefix_len - prev_reasoning_len
|
||||
stripped = delta.reasoning[chars_of_prefix_in_delta:]
|
||||
if stripped:
|
||||
self._prefix_stripped = True
|
||||
delta.reasoning = stripped
|
||||
return delta
|
||||
if len(self._reasoning_text) >= prefix_len:
|
||||
self._prefix_stripped = True
|
||||
delta.reasoning = None
|
||||
if delta.content is not None or delta.tool_calls:
|
||||
return delta
|
||||
return None
|
||||
return None
|
||||
|
||||
if _GEMMA4_THOUGHT_PREFIX.startswith(self._reasoning_text):
|
||||
if finished:
|
||||
self._prefix_stripped = True
|
||||
return None
|
||||
|
||||
self._prefix_stripped = True
|
||||
delta.reasoning = self._reasoning_text
|
||||
return delta
|
||||
|
||||
def extract_reasoning(
|
||||
self,
|
||||
model_output: str,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> tuple[str | None, str | None]:
|
||||
reasoning, content = super().extract_reasoning(model_output, request)
|
||||
if reasoning:
|
||||
if reasoning.startswith(_GEMMA4_THOUGHT_PREFIX):
|
||||
reasoning = reasoning[len(_GEMMA4_THOUGHT_PREFIX) :]
|
||||
elif reasoning == _GEMMA4_THOUGHT_PREFIX.rstrip():
|
||||
reasoning = None
|
||||
return reasoning or None, content
|
||||
@@ -0,0 +1,226 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""GLM-4.7 parser for reasoning and tool calls.
|
||||
|
||||
GLM-4.7 uses XML-like tool calls::
|
||||
|
||||
<tool_call>func_name<arg_key>key</arg_key><arg_value>value</arg_value></tool_call>
|
||||
|
||||
The function name can be followed directly by the first ``<arg_key>`` tag,
|
||||
and tool calls may have no arguments.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
import json
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import regex as re
|
||||
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
from vllm.parser.engine.events import EventType
|
||||
from vllm.parser.engine.parser_engine import ParserEngine
|
||||
from vllm.parser.engine.parser_engine_config import (
|
||||
ParserEngineConfig,
|
||||
ParserState,
|
||||
Transition,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.tokenizers import TokenizerLike
|
||||
from vllm.tool_parsers.abstract_tool_parser import Tool
|
||||
|
||||
THINK_START = "<think>"
|
||||
THINK_END = "</think>"
|
||||
TOOL_CALL_START = "<tool_call>"
|
||||
TOOL_CALL_END = "</tool_call>"
|
||||
ARG_KEY_START = "<arg_key>"
|
||||
ARG_KEY_END = "</arg_key>"
|
||||
ARG_VALUE_START = "<arg_value>"
|
||||
ARG_VALUE_END = "</arg_value>"
|
||||
|
||||
_ARG_RE = re.compile(
|
||||
r"<arg_key>(?P<key>.*?)</arg_key>\s*"
|
||||
r"<arg_value>(?P<value>.*?)</arg_value>",
|
||||
re.DOTALL,
|
||||
)
|
||||
_PARTIAL_ARG_RE = re.compile(
|
||||
r"<arg_key>(?P<key>.*?)</arg_key>\s*"
|
||||
r"<arg_value>(?P<value>.*)$",
|
||||
re.DOTALL,
|
||||
)
|
||||
|
||||
|
||||
def _glm47_arg_converter(raw_args: str, partial: bool) -> str:
|
||||
params: dict[str, object] = {}
|
||||
|
||||
for match in _ARG_RE.finditer(raw_args):
|
||||
params[match.group("key").strip()] = match.group("value")
|
||||
|
||||
if partial:
|
||||
remaining = _ARG_RE.sub("", raw_args)
|
||||
match = _PARTIAL_ARG_RE.search(remaining)
|
||||
if match:
|
||||
key = match.group("key").strip()
|
||||
if key:
|
||||
params[key] = match.group("value")
|
||||
|
||||
return json.dumps(params, ensure_ascii=False)
|
||||
|
||||
|
||||
@functools.cache
|
||||
def glm47_moe_config(thinking: bool = True) -> ParserEngineConfig:
|
||||
arg_tag_transitions = {
|
||||
(ParserState.TOOL_ARGS, terminal): Transition(
|
||||
ParserState.TOOL_ARGS,
|
||||
(EventType.ARG_VALUE_CHUNK,),
|
||||
)
|
||||
for terminal in (
|
||||
"ARG_KEY_START",
|
||||
"ARG_KEY_END",
|
||||
"ARG_VALUE_START",
|
||||
"ARG_VALUE_END",
|
||||
)
|
||||
}
|
||||
|
||||
reasoning_terminals = (
|
||||
{
|
||||
"THINK_START": THINK_START,
|
||||
"THINK_END": THINK_END,
|
||||
}
|
||||
if thinking
|
||||
else {}
|
||||
)
|
||||
reasoning_token_id_terminals = (
|
||||
{
|
||||
"THINK_START": THINK_START,
|
||||
"THINK_END": THINK_END,
|
||||
}
|
||||
if thinking
|
||||
else {}
|
||||
)
|
||||
reasoning_transitions = (
|
||||
{
|
||||
(ParserState.CONTENT, "THINK_START"): Transition(
|
||||
ParserState.REASONING,
|
||||
(EventType.REASONING_START,),
|
||||
),
|
||||
(ParserState.REASONING, "THINK_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(EventType.REASONING_END,),
|
||||
),
|
||||
(ParserState.CONTENT, "THINK_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(),
|
||||
),
|
||||
}
|
||||
if thinking
|
||||
else {}
|
||||
)
|
||||
|
||||
return ParserEngineConfig(
|
||||
name="glm47_moe",
|
||||
initial_state=ParserState.REASONING if thinking else ParserState.CONTENT,
|
||||
terminals={
|
||||
**reasoning_terminals,
|
||||
"TOOL_START": TOOL_CALL_START,
|
||||
"TOOL_END": TOOL_CALL_END,
|
||||
"ARG_KEY_START": ARG_KEY_START,
|
||||
"ARG_KEY_END": ARG_KEY_END,
|
||||
"ARG_VALUE_START": ARG_VALUE_START,
|
||||
"ARG_VALUE_END": ARG_VALUE_END,
|
||||
},
|
||||
token_id_terminals={
|
||||
**reasoning_token_id_terminals,
|
||||
"TOOL_START": TOOL_CALL_START,
|
||||
"TOOL_END": TOOL_CALL_END,
|
||||
},
|
||||
transitions={
|
||||
**reasoning_transitions,
|
||||
(ParserState.REASONING, "THINK_START"): Transition(
|
||||
ParserState.REASONING,
|
||||
(),
|
||||
),
|
||||
(ParserState.REASONING, "TOOL_START"): Transition(
|
||||
ParserState.TOOL_NAME,
|
||||
(EventType.REASONING_END, EventType.TOOL_CALL_START),
|
||||
),
|
||||
(ParserState.CONTENT, "TOOL_START"): Transition(
|
||||
ParserState.TOOL_NAME,
|
||||
(EventType.TOOL_CALL_START,),
|
||||
),
|
||||
(ParserState.TOOL_NAME, "ARG_KEY_START"): Transition(
|
||||
ParserState.TOOL_ARGS,
|
||||
(EventType.ARG_VALUE_CHUNK,),
|
||||
),
|
||||
(ParserState.TOOL_NAME, "TOOL_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(EventType.TOOL_CALL_END,),
|
||||
),
|
||||
(ParserState.TOOL_ARGS, "TOOL_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(EventType.TOOL_CALL_END,),
|
||||
),
|
||||
**arg_tag_transitions,
|
||||
},
|
||||
arg_converter=_glm47_arg_converter,
|
||||
stream_arg_deltas=True,
|
||||
tool_args_json=False,
|
||||
validate_tool_names=True,
|
||||
)
|
||||
|
||||
|
||||
class Glm47MoeParser(ParserEngine):
|
||||
"""GLM-4.7 parser backed by the declarative parser engine."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: TokenizerLike,
|
||||
tools: list[Tool] | None = None,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
chat_kwargs = kwargs.get("chat_template_kwargs", {}) or {}
|
||||
thinking = chat_kwargs.get("thinking", None)
|
||||
enable_thinking = chat_kwargs.get("enable_thinking", None)
|
||||
self.thinking_enabled = (
|
||||
True
|
||||
if thinking is None and enable_thinking is None
|
||||
else bool(thinking) or bool(enable_thinking)
|
||||
)
|
||||
kwargs.setdefault(
|
||||
"parser_engine_config",
|
||||
glm47_moe_config(thinking=self.thinking_enabled),
|
||||
)
|
||||
super().__init__(tokenizer, tools, **kwargs)
|
||||
|
||||
def _emit_name_delta(self, idx: int, deltas, name: str | None) -> None:
|
||||
if name is not None:
|
||||
name = name.strip()
|
||||
super()._emit_name_delta(idx, deltas, name)
|
||||
|
||||
def _handle_tool_end(self, event, deltas) -> None:
|
||||
idx = event.tool_index
|
||||
if 0 <= idx < len(self._tool_slots):
|
||||
self._tool_slots[idx].name = self._tool_slots[idx].name.strip()
|
||||
super()._handle_tool_end(event, deltas)
|
||||
|
||||
def is_reasoning_end(self, input_ids: list[int]) -> bool:
|
||||
if not self.thinking_enabled:
|
||||
return True
|
||||
return super().is_reasoning_end(input_ids)
|
||||
|
||||
def extract_content_ids(self, input_ids: list[int]) -> list[int]:
|
||||
if not self.thinking_enabled:
|
||||
return input_ids
|
||||
return super().extract_content_ids(input_ids)
|
||||
|
||||
def extract_reasoning(
|
||||
self,
|
||||
model_output: str,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> tuple[str | None, str | None]:
|
||||
if not self.thinking_enabled:
|
||||
return None, model_output
|
||||
return super().extract_reasoning(model_output, request)
|
||||
@@ -0,0 +1,358 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum, auto
|
||||
from typing import TYPE_CHECKING, NamedTuple
|
||||
|
||||
from openai_harmony import HarmonyError, Message, Role
|
||||
|
||||
from vllm.entrypoints.chat_utils import make_tool_call_id
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
|
||||
from vllm.entrypoints.openai.engine.protocol import (
|
||||
DeltaFunctionCall,
|
||||
DeltaMessage,
|
||||
DeltaToolCall,
|
||||
FunctionCall,
|
||||
)
|
||||
from vllm.entrypoints.openai.parser.harmony_utils import (
|
||||
extract_function_from_recipient,
|
||||
get_streamable_parser_for_assistant,
|
||||
is_function_recipient,
|
||||
)
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
from vllm.logger import init_logger
|
||||
from vllm.parser.abstract_parser import DelegatingParser
|
||||
from vllm.reasoning.gptoss_reasoning_parser import GptOssReasoningParser
|
||||
from vllm.tool_parsers.gptoss_tool_parser import GptOssToolParser
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from openai_harmony import Message, StreamableParser
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class _SegmentType(Enum):
|
||||
TOOL = auto()
|
||||
REASONING = auto()
|
||||
CONTENT = auto()
|
||||
IGNORE = auto()
|
||||
|
||||
@staticmethod
|
||||
def from_channel_and_recipient(
|
||||
channel: str | None, recipient: str | None
|
||||
) -> _SegmentType:
|
||||
if recipient and is_function_recipient(recipient):
|
||||
return _SegmentType.TOOL
|
||||
if channel == "analysis":
|
||||
return _SegmentType.REASONING
|
||||
if channel == "final" or (channel == "commentary" and recipient is None):
|
||||
return _SegmentType.CONTENT
|
||||
return _SegmentType.IGNORE
|
||||
|
||||
|
||||
class Segment(NamedTuple):
|
||||
channel: str | None
|
||||
recipient: str | None
|
||||
delta: str
|
||||
completed_message: Message | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChunkResult:
|
||||
segments: list[Segment]
|
||||
reasoning_token_count: int
|
||||
|
||||
|
||||
class HarmonyParser(DelegatingParser):
|
||||
def __init__(self, tokenizer, tools=None, *args, **kwargs):
|
||||
super().__init__(tokenizer, tools, *args, **kwargs)
|
||||
|
||||
if self.reasoning_parser and not isinstance(
|
||||
self.reasoning_parser, GptOssReasoningParser
|
||||
):
|
||||
raise ValueError(
|
||||
"Harmony requires GptOssReasoningParser, "
|
||||
f"got {self.reasoning_parser.__class__.__name__}."
|
||||
)
|
||||
|
||||
if self.tool_parser and not isinstance(self.tool_parser, GptOssToolParser):
|
||||
raise ValueError(
|
||||
"Harmony requires GptOssToolParser, "
|
||||
f"got {self.tool_parser.__class__.__name__}."
|
||||
)
|
||||
|
||||
self._parser: StreamableParser | None = None
|
||||
self._next_tool_call_index = 0
|
||||
self._num_processed_messages = 0
|
||||
|
||||
# For error recovery
|
||||
self._current_message_tokens: list[int] = []
|
||||
|
||||
@property
|
||||
def _harmony_parser(self) -> StreamableParser:
|
||||
"""Lazily initializes the Harmony parser."""
|
||||
if self._parser is None:
|
||||
self._parser = get_streamable_parser_for_assistant()
|
||||
return self._parser
|
||||
|
||||
def _poll_completed_message(self) -> Message | None:
|
||||
messages = self._harmony_parser.messages
|
||||
if len(messages) <= self._num_processed_messages:
|
||||
return None
|
||||
msg = messages[self._num_processed_messages]
|
||||
msg.recipient = self._normalize_recipient(msg.recipient)
|
||||
self._num_processed_messages += 1
|
||||
return msg
|
||||
|
||||
def flush(self) -> list[Segment]:
|
||||
segments: list[Segment] = []
|
||||
try:
|
||||
self._harmony_parser.process_eos()
|
||||
msg = self._poll_completed_message()
|
||||
except HarmonyError:
|
||||
logger.warning(
|
||||
"Harmony parser ended in a non-terminal state; returning the "
|
||||
"recovered raw output."
|
||||
)
|
||||
|
||||
final_channel = "final"
|
||||
text = self.model_tokenizer.decode(self._current_message_tokens)
|
||||
segments.append(
|
||||
Segment(
|
||||
channel=final_channel,
|
||||
recipient=None,
|
||||
delta=text,
|
||||
completed_message=None,
|
||||
)
|
||||
)
|
||||
msg = Message.from_role_and_content(Role.ASSISTANT, text).with_channel(
|
||||
final_channel
|
||||
)
|
||||
|
||||
# Reset to the initial assistant-parser state for the next turn.
|
||||
self._parser = None
|
||||
self._num_processed_messages = 0
|
||||
self._current_message_tokens.clear()
|
||||
|
||||
if msg is None:
|
||||
return segments
|
||||
|
||||
segments.append(
|
||||
Segment(
|
||||
channel=msg.channel,
|
||||
recipient=msg.recipient,
|
||||
delta="",
|
||||
completed_message=msg,
|
||||
)
|
||||
)
|
||||
return segments
|
||||
|
||||
def parse(
|
||||
self,
|
||||
model_output: str,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
enable_auto_tools: bool = False,
|
||||
model_output_token_ids: Sequence[int] = (),
|
||||
) -> tuple[str | None, str | None, list[FunctionCall] | None]:
|
||||
"""Parse Harmony output from token IDs.
|
||||
|
||||
Tool calls are always extracted regardless of ``enable_auto_tools``.
|
||||
Callers must decide whether to surface them.
|
||||
"""
|
||||
result = self.process_chunk(model_output_token_ids)
|
||||
flushed_segments = self.flush()
|
||||
if flushed_segments:
|
||||
result.segments.extend(flushed_segments)
|
||||
|
||||
reasoning_parts: list[str] = []
|
||||
content_parts: list[str] = []
|
||||
tool_calls: list[FunctionCall] = []
|
||||
|
||||
for segment in result.segments:
|
||||
msg = segment.completed_message
|
||||
if msg is None:
|
||||
continue
|
||||
if msg.author.role != "assistant" or not msg.content:
|
||||
continue
|
||||
text = msg.content[0].text
|
||||
segment_type = _SegmentType.from_channel_and_recipient(
|
||||
msg.channel, msg.recipient
|
||||
)
|
||||
match segment_type:
|
||||
case _SegmentType.REASONING if self.reasoning_parser and text:
|
||||
reasoning_parts.append(text)
|
||||
case _SegmentType.CONTENT if text:
|
||||
content_parts.append(text)
|
||||
case _SegmentType.TOOL if self.tool_parser:
|
||||
recipient = msg.recipient
|
||||
content_type = msg.content_type
|
||||
assert recipient is not None
|
||||
if content_type is not None and "json" not in content_type:
|
||||
arguments = text
|
||||
else:
|
||||
try:
|
||||
arguments = json.dumps(json.loads(text))
|
||||
except json.JSONDecodeError:
|
||||
arguments = text
|
||||
tool_calls.append(
|
||||
FunctionCall(
|
||||
name=extract_function_from_recipient(recipient),
|
||||
arguments=arguments,
|
||||
)
|
||||
)
|
||||
|
||||
reasoning = "\n".join(reasoning_parts) or None
|
||||
content = "\n".join(content_parts) or None
|
||||
return reasoning, content, tool_calls or None
|
||||
|
||||
def parse_delta(
|
||||
self,
|
||||
delta_text: str,
|
||||
delta_token_ids: list[int],
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
prompt_token_ids: list[int] | None = None,
|
||||
*,
|
||||
finished: bool,
|
||||
) -> DeltaMessage | None:
|
||||
prev_recipient = self._normalize_recipient(
|
||||
self._harmony_parser.current_recipient
|
||||
)
|
||||
result = self.process_chunk(delta_token_ids)
|
||||
if finished:
|
||||
flushed_segments = self.flush()
|
||||
if flushed_segments:
|
||||
result.segments.extend(flushed_segments)
|
||||
combined_content = ""
|
||||
combined_reasoning = ""
|
||||
tool_messages: list[DeltaToolCall] = []
|
||||
|
||||
for segment in result.segments:
|
||||
if segment.completed_message is not None:
|
||||
prev_recipient = None
|
||||
continue
|
||||
|
||||
segment_type = _SegmentType.from_channel_and_recipient(
|
||||
segment.channel, segment.recipient
|
||||
)
|
||||
match segment_type:
|
||||
case _SegmentType.REASONING if self.reasoning_parser:
|
||||
combined_reasoning += segment.delta
|
||||
case _SegmentType.CONTENT:
|
||||
combined_content += segment.delta
|
||||
case _SegmentType.TOOL if self.tool_parser:
|
||||
assert segment.recipient is not None
|
||||
if prev_recipient != segment.recipient:
|
||||
tool_name = extract_function_from_recipient(segment.recipient)
|
||||
tool_messages.append(
|
||||
DeltaToolCall(
|
||||
# HarmonyParser does not use _stream_state;
|
||||
# "random" tool_call_id_type is always used
|
||||
id=make_tool_call_id(),
|
||||
type="function",
|
||||
function=DeltaFunctionCall(
|
||||
name=tool_name,
|
||||
arguments=segment.delta,
|
||||
),
|
||||
index=self._next_tool_call_index,
|
||||
)
|
||||
)
|
||||
self._next_tool_call_index += 1
|
||||
prev_recipient = segment.recipient
|
||||
elif segment.delta:
|
||||
idx = self._next_tool_call_index - 1
|
||||
if tool_messages:
|
||||
tool_msg = tool_messages[-1]
|
||||
assert tool_msg.index == idx
|
||||
fn = tool_msg.function
|
||||
assert fn is not None and fn.arguments is not None
|
||||
fn.arguments += segment.delta
|
||||
else:
|
||||
tool_messages.append(
|
||||
DeltaToolCall(
|
||||
index=idx,
|
||||
function=DeltaFunctionCall(arguments=segment.delta),
|
||||
)
|
||||
)
|
||||
|
||||
if finished:
|
||||
self._next_tool_call_index = 0
|
||||
|
||||
if not combined_content and not combined_reasoning and not tool_messages:
|
||||
return None
|
||||
|
||||
delta_message = DeltaMessage()
|
||||
if combined_content:
|
||||
delta_message.content = combined_content
|
||||
if combined_reasoning:
|
||||
delta_message.reasoning = combined_reasoning
|
||||
if tool_messages:
|
||||
delta_message.tool_calls = tool_messages
|
||||
|
||||
# Suppress reasoning deltas if not requested
|
||||
if delta_message and not request.include_reasoning:
|
||||
delta_message.reasoning = None
|
||||
|
||||
# If only reasoning was in the message (no content, no tool_calls)
|
||||
# skip emitting entirely
|
||||
if not delta_message.content and not delta_message.tool_calls:
|
||||
return None
|
||||
|
||||
return delta_message
|
||||
|
||||
def process_chunk(self, token_ids: Sequence[int]) -> ChunkResult:
|
||||
if not token_ids:
|
||||
return ChunkResult(segments=[], reasoning_token_count=0)
|
||||
|
||||
segments: list[Segment] = []
|
||||
reasoning_token_count = 0
|
||||
for token_id in token_ids:
|
||||
self._harmony_parser.process(token_id)
|
||||
channel = self._harmony_parser.current_channel
|
||||
recipient = self._normalize_recipient(
|
||||
self._harmony_parser.current_recipient
|
||||
)
|
||||
delta = self._harmony_parser.last_content_delta or ""
|
||||
completed_message = self._poll_completed_message()
|
||||
|
||||
if completed_message is not None:
|
||||
self._current_message_tokens.clear()
|
||||
else:
|
||||
self._current_message_tokens.append(token_id)
|
||||
|
||||
if channel == "analysis" or (
|
||||
channel == "commentary" and recipient is not None
|
||||
):
|
||||
reasoning_token_count += 1
|
||||
|
||||
segments.append(
|
||||
Segment(
|
||||
channel=channel,
|
||||
recipient=recipient,
|
||||
delta=delta,
|
||||
completed_message=completed_message,
|
||||
)
|
||||
)
|
||||
|
||||
# TODO: Optionally merge and suppress empty Segments
|
||||
|
||||
return ChunkResult(
|
||||
segments=segments,
|
||||
reasoning_token_count=reasoning_token_count,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _normalize_recipient(recipient: str | None) -> str | None:
|
||||
"""Remove constrained formats misparsed into recipients by older Harmony."""
|
||||
if recipient is None:
|
||||
return None
|
||||
|
||||
constrain_index = recipient.find("<|constrain|>")
|
||||
if constrain_index == -1:
|
||||
return recipient
|
||||
return recipient[:constrain_index].rstrip() or None
|
||||
@@ -0,0 +1,285 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Kimi K2 parser for reasoning and tool calls.
|
||||
|
||||
Kimi K2 tool call format::
|
||||
|
||||
<|tool_calls_section_begin|>
|
||||
<|tool_call_begin|>functions.get_weather:0
|
||||
<|tool_call_argument_begin|>{"city": "Tokyo"}<|tool_call_end|>
|
||||
<|tool_calls_section_end|>
|
||||
|
||||
The header before ``<|tool_call_argument_begin|>`` is Kimi's native tool
|
||||
call id. The function name is the final component before ``:N``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
from collections.abc import Sequence
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import regex as re
|
||||
|
||||
from vllm.entrypoints.openai.engine.protocol import DeltaFunctionCall, DeltaToolCall
|
||||
from vllm.parser.engine.events import EventType
|
||||
from vllm.parser.engine.parser_engine import ParserEngine
|
||||
from vllm.parser.engine.parser_engine_config import (
|
||||
ParserEngineConfig,
|
||||
ParserState,
|
||||
Transition,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import (
|
||||
ChatCompletionRequest,
|
||||
)
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
from vllm.tokenizers import TokenizerLike
|
||||
from vllm.tool_parsers.abstract_tool_parser import Tool
|
||||
|
||||
THINK_START = "<think>"
|
||||
THINK_END = "</think>"
|
||||
TOOL_SECTION_START = "<|tool_calls_section_begin|>"
|
||||
TOOL_SECTION_END = "<|tool_calls_section_end|>"
|
||||
TOOL_CALL_START = "<|tool_call_begin|>"
|
||||
TOOL_CALL_END = "<|tool_call_end|>"
|
||||
TOOL_ARG_START = "<|tool_call_argument_begin|>"
|
||||
|
||||
_TOOL_ID_RE = re.compile(r"(?P<id>.+:\d+)")
|
||||
|
||||
|
||||
@functools.cache
|
||||
def kimi_k2_config(thinking: bool = True) -> ParserEngineConfig:
|
||||
reasoning_terminals = (
|
||||
{
|
||||
"THINK_START": THINK_START,
|
||||
"THINK_END": THINK_END,
|
||||
}
|
||||
if thinking
|
||||
else {}
|
||||
)
|
||||
reasoning_transitions = (
|
||||
{
|
||||
(ParserState.REASONING, "THINK_START"): Transition(
|
||||
ParserState.REASONING,
|
||||
(),
|
||||
),
|
||||
(ParserState.REASONING, "THINK_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(EventType.REASONING_END,),
|
||||
),
|
||||
(ParserState.CONTENT, "THINK_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(),
|
||||
),
|
||||
}
|
||||
if thinking
|
||||
else {}
|
||||
)
|
||||
|
||||
return ParserEngineConfig(
|
||||
name="kimi_k2",
|
||||
initial_state=ParserState.REASONING if thinking else ParserState.CONTENT,
|
||||
terminals={
|
||||
**reasoning_terminals,
|
||||
"TOOL_SECTION_START": TOOL_SECTION_START,
|
||||
"TOOL_SECTION_END": TOOL_SECTION_END,
|
||||
"TOOL_START": TOOL_CALL_START,
|
||||
"TOOL_END": TOOL_CALL_END,
|
||||
"ARG_START": TOOL_ARG_START,
|
||||
},
|
||||
token_id_terminals={
|
||||
**reasoning_terminals,
|
||||
"TOOL_SECTION_START": TOOL_SECTION_START,
|
||||
"TOOL_SECTION_END": TOOL_SECTION_END,
|
||||
"TOOL_START": TOOL_CALL_START,
|
||||
"TOOL_END": TOOL_CALL_END,
|
||||
"ARG_START": TOOL_ARG_START,
|
||||
},
|
||||
transitions={
|
||||
**reasoning_transitions,
|
||||
(ParserState.REASONING, "TOOL_SECTION_START"): Transition(
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
(EventType.REASONING_END,),
|
||||
),
|
||||
(ParserState.CONTENT, "TOOL_SECTION_START"): Transition(
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
(),
|
||||
),
|
||||
(ParserState.TOOL_PREAMBLE, "TOOL_START"): Transition(
|
||||
ParserState.TOOL_NAME,
|
||||
(EventType.TOOL_CALL_START,),
|
||||
),
|
||||
(ParserState.TOOL_NAME, "ARG_START"): Transition(
|
||||
ParserState.TOOL_ARGS,
|
||||
(),
|
||||
),
|
||||
(ParserState.TOOL_ARGS, "TOOL_END"): Transition(
|
||||
ParserState.TOOL_BETWEEN,
|
||||
(EventType.TOOL_CALL_END,),
|
||||
),
|
||||
(ParserState.TOOL_ARGS, "TOOL_SECTION_END"): Transition(
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
(EventType.TOOL_CALL_END,),
|
||||
),
|
||||
(ParserState.TOOL_BETWEEN, "TOOL_START"): Transition(
|
||||
ParserState.TOOL_NAME,
|
||||
(EventType.TOOL_CALL_START,),
|
||||
),
|
||||
# Keep the parser in a tool state after the section closes so
|
||||
# trailing model text after native tool calls is suppressed.
|
||||
(ParserState.TOOL_PREAMBLE, "TOOL_SECTION_END"): Transition(
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
(),
|
||||
),
|
||||
(ParserState.TOOL_BETWEEN, "TOOL_SECTION_END"): Transition(
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
(),
|
||||
),
|
||||
},
|
||||
stream_arg_deltas=True,
|
||||
tool_args_json=True,
|
||||
strip_trailing_reasoning_whitespace=True,
|
||||
drop_whitespace_only_content_before_tools=True,
|
||||
strip_content_whitespace_with_tools=False,
|
||||
validate_tool_names=False,
|
||||
)
|
||||
|
||||
|
||||
class KimiK2Parser(ParserEngine):
|
||||
"""Kimi K2 parser backed by the declarative parser engine."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: TokenizerLike,
|
||||
tools: list[Tool] | None = None,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
chat_kwargs = kwargs.get("chat_template_kwargs", {}) or {}
|
||||
thinking = chat_kwargs.get("thinking", None)
|
||||
enable_thinking = chat_kwargs.get("enable_thinking", None)
|
||||
self.thinking_enabled = (
|
||||
True
|
||||
if thinking is None and enable_thinking is None
|
||||
else bool(thinking) or bool(enable_thinking)
|
||||
)
|
||||
kwargs.setdefault(
|
||||
"parser_engine_config",
|
||||
kimi_k2_config(thinking=self.thinking_enabled),
|
||||
)
|
||||
super().__init__(tokenizer, tools, **kwargs)
|
||||
|
||||
vocab = self.vocab
|
||||
self._start_token_id = vocab.get(THINK_START)
|
||||
self._end_token_id = vocab.get(THINK_END)
|
||||
self._tool_section_start_token_id = vocab.get(TOOL_SECTION_START)
|
||||
|
||||
@staticmethod
|
||||
def _extract_tool_id_and_name(header: str | None) -> tuple[str | None, str | None]:
|
||||
if header is None:
|
||||
return None, None
|
||||
match = _TOOL_ID_RE.match(header.strip())
|
||||
if not match:
|
||||
return None, None
|
||||
|
||||
tool_id = match.group("id").strip()
|
||||
tool_name = tool_id.split(":")[0].removeprefix("functions.")
|
||||
return tool_id, tool_name
|
||||
|
||||
def _emit_name_delta(
|
||||
self,
|
||||
idx: int,
|
||||
deltas: list[DeltaToolCall],
|
||||
name: str | None,
|
||||
) -> None:
|
||||
tool_id, tool_name = self._extract_tool_id_and_name(name)
|
||||
if not tool_name:
|
||||
if 0 <= idx < len(self._tool_slots):
|
||||
self._tool_slots[idx].name = ""
|
||||
return
|
||||
|
||||
slot = self._tool_slots[idx]
|
||||
slot.id = tool_id or ""
|
||||
super()._emit_name_delta(idx, deltas, tool_name)
|
||||
|
||||
def _handle_tool_end(self, event, deltas) -> None:
|
||||
idx = event.tool_index
|
||||
if 0 <= idx < len(self._tool_slots) and not self._tool_slots[idx].name_sent:
|
||||
tool_id, tool_name = self._extract_tool_id_and_name(
|
||||
self._tool_slots[idx].name
|
||||
)
|
||||
if tool_name:
|
||||
self._tool_slots[idx].id = tool_id or ""
|
||||
self._tool_slots[idx].name = tool_name
|
||||
super()._handle_tool_end(event, deltas)
|
||||
|
||||
def _handle_arg_chunk(self, event, deltas) -> None:
|
||||
idx = event.tool_index
|
||||
name_sent_before = (
|
||||
0 <= idx < len(self._tool_slots) and self._tool_slots[idx].name_sent
|
||||
)
|
||||
super()._handle_arg_chunk(event, deltas)
|
||||
if (
|
||||
event.value
|
||||
and not name_sent_before
|
||||
and 0 <= idx < len(self._tool_slots)
|
||||
and self._tool_slots[idx].name_sent
|
||||
):
|
||||
deltas.append(
|
||||
DeltaToolCall(
|
||||
index=idx,
|
||||
function=DeltaFunctionCall(arguments=event.value),
|
||||
)
|
||||
)
|
||||
|
||||
def _extract_args_json(self, raw_args: str, func_name: str) -> str:
|
||||
return raw_args.strip() or "{}"
|
||||
|
||||
def is_reasoning_end(self, input_ids: list[int]) -> bool:
|
||||
if not self.thinking_enabled:
|
||||
return True
|
||||
|
||||
start_id = self._start_token_id
|
||||
end_id = self._end_token_id
|
||||
tool_section_id = self._tool_section_start_token_id
|
||||
|
||||
for i in range(len(input_ids) - 1, -1, -1):
|
||||
token_id = input_ids[i]
|
||||
if start_id is not None and token_id == start_id:
|
||||
return False
|
||||
if end_id is not None and token_id == end_id:
|
||||
return True
|
||||
if tool_section_id is not None and token_id == tool_section_id:
|
||||
return True
|
||||
return False
|
||||
|
||||
def extract_content_ids(self, input_ids: list[int]) -> list[int]:
|
||||
if not self.thinking_enabled:
|
||||
return input_ids
|
||||
|
||||
end_id = self._end_token_id
|
||||
if end_id is not None and end_id in input_ids:
|
||||
end_idx = len(input_ids) - 1 - input_ids[::-1].index(end_id)
|
||||
return input_ids[end_idx + 1 :]
|
||||
|
||||
tool_section_id = self._tool_section_start_token_id
|
||||
if tool_section_id is not None and tool_section_id in input_ids:
|
||||
section_idx = len(input_ids) - 1 - input_ids[::-1].index(tool_section_id)
|
||||
return input_ids[section_idx:]
|
||||
|
||||
return []
|
||||
|
||||
def extract_reasoning(
|
||||
self,
|
||||
model_output: str,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> tuple[str | None, str | None]:
|
||||
if not self.thinking_enabled:
|
||||
return None, model_output
|
||||
return super().extract_reasoning(model_output, request)
|
||||
|
||||
def count_reasoning_tokens(self, token_ids: Sequence[int]) -> int:
|
||||
if not self.thinking_enabled:
|
||||
return 0
|
||||
return super().count_reasoning_tokens(token_ids)
|
||||
@@ -0,0 +1,108 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Prometheus metrics for the parsers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import Enum
|
||||
from itertools import product
|
||||
from typing import cast
|
||||
|
||||
from prometheus_client import REGISTRY, Counter
|
||||
|
||||
_model_name: str | None = None
|
||||
|
||||
_TOOL_CALL_PARSER_INVOCATIONS_TOTAL = "vllm:tool_call_parser_invocations_total"
|
||||
_tool_call_parser_invocations: Counter | None = None
|
||||
|
||||
|
||||
class ToolCallOutcome(Enum):
|
||||
TOOL_CALL = "tool_call"
|
||||
NO_TOOL_CALL = "no_tool_call"
|
||||
|
||||
|
||||
class RequestType(Enum):
|
||||
CHAT_COMPLETIONS = "chat_completions"
|
||||
RESPONSES = "responses"
|
||||
OTHER = "other"
|
||||
|
||||
|
||||
def init_parser_metrics(*, model_name: str) -> None:
|
||||
"""Lazily register parser metrics and cache the shared model label."""
|
||||
global _model_name
|
||||
_model_name = model_name
|
||||
|
||||
global _tool_call_parser_invocations
|
||||
try:
|
||||
_tool_call_parser_invocations = Counter(
|
||||
name=_TOOL_CALL_PARSER_INVOCATIONS_TOTAL,
|
||||
documentation=(
|
||||
"Total number of ToolParser invocations. "
|
||||
"Non-streaming increments once per choice; "
|
||||
"streaming increments once per delta."
|
||||
),
|
||||
labelnames=["model_name", "mode", "outcome", "request_type"],
|
||||
)
|
||||
except ValueError:
|
||||
_tool_call_parser_invocations = cast(
|
||||
Counter,
|
||||
REGISTRY._names_to_collectors[_TOOL_CALL_PARSER_INVOCATIONS_TOTAL],
|
||||
)
|
||||
|
||||
for mode, outcome, request_type in product(
|
||||
("streaming", "non_streaming"),
|
||||
ToolCallOutcome,
|
||||
RequestType,
|
||||
):
|
||||
_tool_call_parser_invocations.labels(
|
||||
model_name=_model_name,
|
||||
mode=mode,
|
||||
outcome=outcome.value,
|
||||
request_type=request_type.value,
|
||||
)
|
||||
|
||||
|
||||
def record_tool_parser_invocation(
|
||||
*,
|
||||
is_tool_called: bool | Exception,
|
||||
is_streaming: bool,
|
||||
request: object,
|
||||
) -> None:
|
||||
"""Increment the tool-call parser invocation counter when registered.
|
||||
Currently parser failures are treated as no tool calls.
|
||||
|
||||
TODO: To accurately track parser failures, add a new ToolCallOutcome and
|
||||
more importantly, ensure exceptions are propagated out of the ToolParsers
|
||||
instead of being caught internally. This would require going through
|
||||
ToolParser implementation on a case-by-case basis.
|
||||
"""
|
||||
if _tool_call_parser_invocations is None:
|
||||
return
|
||||
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
|
||||
match request:
|
||||
case ChatCompletionRequest():
|
||||
request_type = RequestType.CHAT_COMPLETIONS
|
||||
case ResponsesRequest():
|
||||
request_type = RequestType.RESPONSES
|
||||
case _:
|
||||
request_type = RequestType.OTHER
|
||||
|
||||
match is_tool_called:
|
||||
case bool():
|
||||
outcome = (
|
||||
ToolCallOutcome.TOOL_CALL
|
||||
if is_tool_called
|
||||
else ToolCallOutcome.NO_TOOL_CALL
|
||||
)
|
||||
case _:
|
||||
outcome = ToolCallOutcome.NO_TOOL_CALL
|
||||
|
||||
_tool_call_parser_invocations.labels(
|
||||
model_name=_model_name,
|
||||
mode="streaming" if is_streaming else "non_streaming",
|
||||
outcome=outcome.value,
|
||||
request_type=request_type.value,
|
||||
).inc()
|
||||
@@ -0,0 +1,211 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""MiniMax M2 parser for XML-style tool calls.
|
||||
|
||||
MiniMax M2 tool call format::
|
||||
|
||||
<minimax:tool_call><invoke name="get_weather">
|
||||
<parameter name="city">Seattle</parameter>
|
||||
</invoke></minimax:tool_call>
|
||||
|
||||
Each ``<invoke>`` block becomes one tool call. The argument body consists
|
||||
of ``<parameter name="...">...</parameter>`` tags.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
import json
|
||||
|
||||
import regex as re
|
||||
|
||||
from vllm.parser.engine.events import EventType
|
||||
from vllm.parser.engine.parser_engine import ParserEngine
|
||||
from vllm.parser.engine.parser_engine_config import (
|
||||
ParserEngineConfig,
|
||||
ParserState,
|
||||
Transition,
|
||||
)
|
||||
|
||||
TOOL_CALL_START = "<minimax:tool_call>"
|
||||
TOOL_CALL_END = "</minimax:tool_call>"
|
||||
THINK_START = "<think>"
|
||||
THINK_END = "</think>"
|
||||
INVOKE_PREFIX_DQ = '<invoke name="'
|
||||
INVOKE_PREFIX_SQ = "<invoke name='"
|
||||
INVOKE_PREFIX_UNQUOTED = "<invoke name="
|
||||
INVOKE_END = "</invoke>"
|
||||
NAME_END_DQ = '">'
|
||||
NAME_END_SQ = "'>"
|
||||
NAME_END_UNQUOTED = ">"
|
||||
PARAM_START = "<parameter name="
|
||||
PARAM_END = "</parameter>"
|
||||
_PARAM_RE = re.compile(
|
||||
r"<\s*parameter\s+name\s*=\s*"
|
||||
r"(?:\"(?P<dq_name>[^\"]*)\"|'(?P<sq_name>[^']*)'|(?P<bare_name>[^>\s]+))"
|
||||
r"\s*>"
|
||||
r"(?P<value>.*?)"
|
||||
r"(?:<\s*/\s*parameter\s*>|(?=<\s*parameter\s+name\s*=))",
|
||||
re.DOTALL,
|
||||
)
|
||||
_PARTIAL_PARAM_RE = re.compile(
|
||||
r"<\s*parameter\s+name\s*=\s*"
|
||||
r"(?:\"(?P<dq_name>[^\"]*)\"|'(?P<sq_name>[^']*)'|(?P<bare_name>[^>\s]+))"
|
||||
r"\s*>"
|
||||
r"(?P<value>.*)$",
|
||||
re.DOTALL,
|
||||
)
|
||||
|
||||
|
||||
def _minimax_m2_arg_converter(raw_args: str, partial: bool) -> str:
|
||||
params: dict[str, object] = {}
|
||||
|
||||
for match in _PARAM_RE.finditer(raw_args):
|
||||
name = (
|
||||
match.group("dq_name")
|
||||
or match.group("sq_name")
|
||||
or match.group("bare_name")
|
||||
or ""
|
||||
).strip()
|
||||
if not name:
|
||||
continue
|
||||
params[name] = match.group("value").strip()
|
||||
|
||||
if partial:
|
||||
remaining = _PARAM_RE.sub("", raw_args)
|
||||
match = _PARTIAL_PARAM_RE.search(remaining)
|
||||
if match:
|
||||
name = (
|
||||
match.group("dq_name")
|
||||
or match.group("sq_name")
|
||||
or match.group("bare_name")
|
||||
or ""
|
||||
).strip()
|
||||
if name:
|
||||
params[name] = match.group("value").strip()
|
||||
|
||||
return json.dumps(params, ensure_ascii=False)
|
||||
|
||||
|
||||
@functools.cache
|
||||
def minimax_m2_config() -> ParserEngineConfig:
|
||||
return ParserEngineConfig(
|
||||
name="minimax_m2",
|
||||
initial_state=ParserState.REASONING,
|
||||
terminals={
|
||||
"THINK_START": THINK_START,
|
||||
"THINK_END": THINK_END,
|
||||
"TOOL_START": TOOL_CALL_START,
|
||||
"PARAM_START": PARAM_START,
|
||||
"PARAM_END": PARAM_END,
|
||||
"TOOL_END": TOOL_CALL_END,
|
||||
"INVOKE_PREFIX_DQ": INVOKE_PREFIX_DQ,
|
||||
"INVOKE_PREFIX_SQ": INVOKE_PREFIX_SQ,
|
||||
"INVOKE_PREFIX_UNQUOTED": INVOKE_PREFIX_UNQUOTED,
|
||||
"INVOKE_END": INVOKE_END,
|
||||
"NAME_END_DQ": NAME_END_DQ,
|
||||
"NAME_END_SQ": NAME_END_SQ,
|
||||
"NAME_END_UNQUOTED": NAME_END_UNQUOTED,
|
||||
},
|
||||
token_id_terminals={
|
||||
"THINK_START": THINK_START,
|
||||
"THINK_END": THINK_END,
|
||||
"TOOL_START": TOOL_CALL_START,
|
||||
"TOOL_END": TOOL_CALL_END,
|
||||
},
|
||||
transitions={
|
||||
(ParserState.REASONING, "THINK_START"): Transition(
|
||||
ParserState.REASONING,
|
||||
(),
|
||||
),
|
||||
(ParserState.REASONING, "THINK_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(EventType.REASONING_END,),
|
||||
),
|
||||
(ParserState.CONTENT, "THINK_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(),
|
||||
),
|
||||
(ParserState.REASONING, "TOOL_START"): Transition(
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
(EventType.REASONING_END,),
|
||||
),
|
||||
(ParserState.CONTENT, "TOOL_START"): Transition(
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
(),
|
||||
),
|
||||
(ParserState.TOOL_ARGS, "PARAM_START"): Transition(
|
||||
ParserState.TOOL_ARGS,
|
||||
(EventType.ARG_VALUE_CHUNK,),
|
||||
),
|
||||
(ParserState.TOOL_ARGS, "PARAM_END"): Transition(
|
||||
ParserState.TOOL_ARGS,
|
||||
(EventType.ARG_VALUE_CHUNK,),
|
||||
),
|
||||
(ParserState.TOOL_PREAMBLE, "TOOL_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(),
|
||||
),
|
||||
(ParserState.TOOL_BETWEEN, "TOOL_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(),
|
||||
),
|
||||
(ParserState.CONTENT, "TOOL_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(),
|
||||
),
|
||||
(ParserState.TOOL_ARGS, "INVOKE_END"): Transition(
|
||||
ParserState.TOOL_BETWEEN,
|
||||
(EventType.TOOL_CALL_END,),
|
||||
),
|
||||
**{
|
||||
(state, terminal): Transition(
|
||||
ParserState.TOOL_NAME,
|
||||
(EventType.TOOL_CALL_START,),
|
||||
)
|
||||
for state in (
|
||||
ParserState.CONTENT,
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
ParserState.TOOL_BETWEEN,
|
||||
)
|
||||
for terminal in (
|
||||
"INVOKE_PREFIX_DQ",
|
||||
"INVOKE_PREFIX_SQ",
|
||||
"INVOKE_PREFIX_UNQUOTED",
|
||||
)
|
||||
},
|
||||
**{
|
||||
(ParserState.TOOL_NAME, terminal): Transition(
|
||||
ParserState.TOOL_ARGS,
|
||||
(),
|
||||
)
|
||||
for terminal in (
|
||||
"NAME_END_DQ",
|
||||
"NAME_END_SQ",
|
||||
"NAME_END_UNQUOTED",
|
||||
)
|
||||
},
|
||||
},
|
||||
arg_converter=_minimax_m2_arg_converter,
|
||||
stream_arg_deltas=True,
|
||||
tool_args_json=False,
|
||||
validate_tool_names=True,
|
||||
)
|
||||
|
||||
|
||||
class MinimaxM2Parser(ParserEngine):
|
||||
"""MiniMax M2 parser backed by the declarative parser engine."""
|
||||
|
||||
def __init__(self, tokenizer, tools=None, **kwargs) -> None:
|
||||
kwargs.setdefault("parser_engine_config", minimax_m2_config())
|
||||
super().__init__(tokenizer, tools, **kwargs)
|
||||
self._think_end_token_id = self.vocab.get(THINK_END)
|
||||
|
||||
def extract_content_ids(self, input_ids: list[int]) -> list[int]:
|
||||
end_id = self._think_end_token_id
|
||||
if end_id is None:
|
||||
return []
|
||||
for i in range(len(input_ids) - 1, -1, -1):
|
||||
if input_ids[i] == end_id:
|
||||
return input_ids[i + 1 :]
|
||||
return []
|
||||
@@ -0,0 +1,82 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from vllm.entrypoints.openai.engine.protocol import DeltaMessage, FunctionCall
|
||||
from vllm.parser.abstract_parser import DelegatingParser
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import (
|
||||
ChatCompletionRequest,
|
||||
)
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
|
||||
|
||||
class MistralParser(DelegatingParser):
|
||||
def __init__(self, tokenizer, tools=None, *args, **kwargs):
|
||||
super().__init__(tokenizer, tools, *args, **kwargs)
|
||||
from vllm.tool_parsers.mistral_tool_parser import MistralToolParser
|
||||
|
||||
if not isinstance(self._tool_parser, MistralToolParser):
|
||||
raise ValueError(
|
||||
"MistralParser requires --tool-call-parser mistral, "
|
||||
f"got {self._tool_parser.__class__.__name__}."
|
||||
)
|
||||
|
||||
def _maybe_force_auto_tool_parsing(
|
||||
self, request: ChatCompletionRequest | ResponsesRequest
|
||||
) -> None:
|
||||
# When the Mistral grammar factory injected structured outputs,
|
||||
# the model emits v11+ format ([TOOL_CALLS]name{args}) that the
|
||||
# named/required parsers can't handle. Disable them so all
|
||||
# tool_choice modes fall back to auto tool parsing via
|
||||
# extract_tool_calls.
|
||||
if getattr(request, "_grammar_from_tool_parser", False):
|
||||
assert self._tool_parser is not None
|
||||
self._tool_parser.supports_required_and_named = False
|
||||
|
||||
def parse(
|
||||
self,
|
||||
model_output: str,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
enable_auto_tools: bool = False,
|
||||
model_output_token_ids: Sequence[int] = (),
|
||||
) -> tuple[str | None, str | None, list[FunctionCall] | None]:
|
||||
self._maybe_force_auto_tool_parsing(request)
|
||||
reasoning, content, tool_calls = super().parse(
|
||||
model_output,
|
||||
request,
|
||||
enable_auto_tools,
|
||||
model_output_token_ids,
|
||||
)
|
||||
if tool_calls:
|
||||
from vllm.tool_parsers.mistral_tool_parser import MistralToolCall
|
||||
|
||||
# Named/required tool_choice builds FunctionCalls without
|
||||
# ID, backfill with Mistral-format IDs.
|
||||
for tc in tool_calls:
|
||||
if not tc.id:
|
||||
tc.id = MistralToolCall.generate_random_id()
|
||||
return reasoning, content, tool_calls
|
||||
|
||||
def parse_delta(
|
||||
self,
|
||||
delta_text: str,
|
||||
delta_token_ids: list[int],
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
prompt_token_ids: list[int] | None = None,
|
||||
*,
|
||||
finished: bool,
|
||||
) -> DeltaMessage | None:
|
||||
self._maybe_force_auto_tool_parsing(request)
|
||||
return super().parse_delta(
|
||||
delta_text,
|
||||
delta_token_ids,
|
||||
request,
|
||||
prompt_token_ids,
|
||||
finished=finished,
|
||||
)
|
||||
@@ -0,0 +1,113 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Nemotron V3 parser.
|
||||
|
||||
The Nemotron 3 Super model uses the same tool call and reasoning
|
||||
format as Qwen3 (``<think>``/``</think>`` + ``<tool_call>`` XML).
|
||||
This config reuses :func:`qwen3_config` with a distinct name.
|
||||
|
||||
When ``enable_thinking=False`` or ``force_nonempty_content=True`` and
|
||||
content is empty, reasoning and content are swapped.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
import functools
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from vllm.parser.qwen3 import Qwen3Parser, qwen3_config
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import (
|
||||
ChatCompletionRequest,
|
||||
)
|
||||
from vllm.entrypoints.openai.engine.protocol import DeltaMessage
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
from vllm.parser.engine.parser_engine import SemanticEvent
|
||||
from vllm.parser.engine.parser_engine_config import ParserEngineConfig
|
||||
from vllm.tokenizers import TokenizerLike
|
||||
from vllm.tool_parsers.abstract_tool_parser import Tool
|
||||
|
||||
|
||||
@functools.cache
|
||||
def nemotron_v3_config(thinking: bool = True) -> ParserEngineConfig:
|
||||
return dataclasses.replace(
|
||||
qwen3_config(thinking=thinking),
|
||||
name="nemotron_v3",
|
||||
strip_trailing_reasoning_whitespace=True,
|
||||
)
|
||||
|
||||
|
||||
class NemotronV3Parser(Qwen3Parser):
|
||||
"""Nemotron V3 parser: same format as Qwen3, with Nemotron-specific
|
||||
behavior: when ``enable_thinking=False`` or
|
||||
``force_nonempty_content=True`` and content is empty, swaps
|
||||
reasoning and content.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: TokenizerLike,
|
||||
tools: list[Tool] | None = None,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
chat_kwargs = kwargs.get("chat_template_kwargs", {}) or {}
|
||||
thinking = chat_kwargs.get("enable_thinking", True)
|
||||
super().__init__(
|
||||
tokenizer,
|
||||
tools,
|
||||
parser_engine_config=nemotron_v3_config(thinking=thinking),
|
||||
**kwargs,
|
||||
)
|
||||
self._streamed_reasoning: list[str] = []
|
||||
|
||||
def _reset(self, initial_state=None) -> None:
|
||||
super()._reset(initial_state=initial_state)
|
||||
self._streamed_reasoning = []
|
||||
|
||||
def _events_to_delta(
|
||||
self,
|
||||
events: list[SemanticEvent],
|
||||
finished: bool = False,
|
||||
) -> DeltaMessage | None:
|
||||
delta = super()._events_to_delta(events, finished=finished)
|
||||
if delta is not None and delta.reasoning is not None:
|
||||
self._streamed_reasoning.append(delta.reasoning)
|
||||
return delta
|
||||
|
||||
@staticmethod
|
||||
def _should_force_content(
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> bool:
|
||||
chat_template_kwargs = getattr(request, "chat_template_kwargs", None)
|
||||
return bool(
|
||||
chat_template_kwargs
|
||||
and (
|
||||
chat_template_kwargs.get("enable_thinking") is False
|
||||
or chat_template_kwargs.get("force_nonempty_content") is True
|
||||
)
|
||||
)
|
||||
|
||||
def get_streaming_fallback_content(
|
||||
self,
|
||||
text: str,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> str | None:
|
||||
if not self._should_force_content(request):
|
||||
return None
|
||||
return "".join(self._streamed_reasoning) or None
|
||||
|
||||
def extract_reasoning(
|
||||
self,
|
||||
model_output: str,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> tuple[str | None, str | None]:
|
||||
reasoning, content = super().extract_reasoning(model_output, request)
|
||||
|
||||
if self._should_force_content(request) and (
|
||||
content is None or not content.strip()
|
||||
):
|
||||
reasoning, content = content, reasoning
|
||||
|
||||
return reasoning, content
|
||||
@@ -0,0 +1,137 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from vllm.logger import init_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.parser.abstract_parser import Parser
|
||||
from vllm.reasoning import ReasoningParser
|
||||
from vllm.tool_parsers import ToolParser
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ParserManager:
|
||||
"""
|
||||
Provides a unified Parser by composing individual reasoning and tool
|
||||
parsers from their respective registries.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def get_tool_parser(
|
||||
cls,
|
||||
tool_parser_name: str | None = None,
|
||||
enable_auto_tools: bool = False,
|
||||
model_name: str | None = None,
|
||||
) -> type[ToolParser] | None:
|
||||
"""Get the tool parser based on the name."""
|
||||
from vllm.tool_parsers import ToolParserManager
|
||||
|
||||
parser: type[ToolParser] | None = None
|
||||
if not enable_auto_tools or tool_parser_name is None:
|
||||
return parser
|
||||
logger.info_once('"auto" tool choice has been enabled.')
|
||||
|
||||
try:
|
||||
if (
|
||||
tool_parser_name == "pythonic"
|
||||
and model_name
|
||||
and model_name.startswith("meta-llama/Llama-3.2")
|
||||
):
|
||||
logger.warning(
|
||||
"Llama3.2 models may struggle to emit valid pythonic tool calls"
|
||||
)
|
||||
parser = ToolParserManager.get_tool_parser(tool_parser_name)
|
||||
except Exception as e:
|
||||
raise TypeError(
|
||||
"Error: --enable-auto-tool-choice requires "
|
||||
f"tool_parser:'{tool_parser_name}' which has not "
|
||||
"been registered"
|
||||
) from e
|
||||
return parser
|
||||
|
||||
@classmethod
|
||||
def get_reasoning_parser(
|
||||
cls,
|
||||
reasoning_parser_name: str | None,
|
||||
) -> type[ReasoningParser] | None:
|
||||
"""Get the reasoning parser based on the name."""
|
||||
from vllm.reasoning import ReasoningParserManager
|
||||
|
||||
parser: type[ReasoningParser] | None = None
|
||||
if not reasoning_parser_name:
|
||||
return None
|
||||
try:
|
||||
parser = ReasoningParserManager.get_reasoning_parser(reasoning_parser_name)
|
||||
assert parser is not None
|
||||
except Exception as e:
|
||||
raise TypeError(f"{reasoning_parser_name=} has not been registered") from e
|
||||
return parser
|
||||
|
||||
@classmethod
|
||||
def get_parser(
|
||||
cls,
|
||||
tool_parser_name: str | None = None,
|
||||
reasoning_parser_name: str | None = None,
|
||||
enable_auto_tools: bool = False,
|
||||
model_name: str | None = None,
|
||||
is_harmony: bool = False,
|
||||
) -> type[Parser] | None:
|
||||
"""
|
||||
Get a Parser that handles both reasoning and tool parsing.
|
||||
|
||||
Composes individual reasoning and tool parsers into a single
|
||||
DelegatingParser subclass.
|
||||
|
||||
Args:
|
||||
tool_parser_name: The name of the tool parser.
|
||||
reasoning_parser_name: The name of the reasoning parser.
|
||||
enable_auto_tools: Whether auto tool choice is enabled.
|
||||
model_name: The model name for parser-specific warnings.
|
||||
is_harmony: Whether the selected model uses the Harmony format.
|
||||
If True, HarmonyParser is always returned.
|
||||
|
||||
Returns:
|
||||
A Parser class, or None if neither parser is specified.
|
||||
"""
|
||||
if not tool_parser_name and not reasoning_parser_name:
|
||||
return None
|
||||
|
||||
reasoning_parser_cls = cls.get_reasoning_parser(reasoning_parser_name)
|
||||
tool_parser_cls = cls.get_tool_parser(
|
||||
tool_parser_name, enable_auto_tools, model_name
|
||||
)
|
||||
|
||||
if reasoning_parser_cls is None and tool_parser_cls is None:
|
||||
return None
|
||||
|
||||
from vllm.utils.mistral import is_mistral_tool_parser
|
||||
|
||||
if is_harmony:
|
||||
from vllm.parser.harmony import HarmonyParser
|
||||
|
||||
HarmonyParser.reasoning_parser_cls = reasoning_parser_cls
|
||||
HarmonyParser.tool_parser_cls = tool_parser_cls
|
||||
return HarmonyParser
|
||||
|
||||
if is_mistral_tool_parser(tool_parser_cls):
|
||||
from vllm.parser.mistral import MistralParser
|
||||
|
||||
MistralParser.reasoning_parser_cls = reasoning_parser_cls
|
||||
MistralParser.tool_parser_cls = tool_parser_cls
|
||||
return MistralParser
|
||||
|
||||
from vllm.parser.abstract_parser import DelegatingParser
|
||||
|
||||
r_cls = reasoning_parser_cls
|
||||
t_cls = tool_parser_cls
|
||||
|
||||
class _Parser(DelegatingParser):
|
||||
reasoning_parser_cls = r_cls
|
||||
tool_parser_cls = t_cls
|
||||
|
||||
return _Parser
|
||||
@@ -0,0 +1,267 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Qwen3 parser for tool calls and reasoning.
|
||||
|
||||
Qwen3 XML tool call format::
|
||||
|
||||
<tool_call>
|
||||
<function=func_name>
|
||||
<parameter=key>value</parameter>
|
||||
</function>
|
||||
</tool_call>
|
||||
|
||||
The argument body consists of ``<parameter=NAME>VALUE</parameter>`` tags.
|
||||
The ``_qwen3_arg_converter`` parses these into a JSON object.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
import json
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import regex as re
|
||||
|
||||
from vllm.parser.engine.events import EventType
|
||||
from vllm.parser.engine.parser_engine import ParserEngine
|
||||
from vllm.parser.engine.parser_engine_config import (
|
||||
ParserEngineConfig,
|
||||
ParserState,
|
||||
Transition,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import (
|
||||
ChatCompletionRequest,
|
||||
)
|
||||
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
|
||||
from vllm.tokenizers import TokenizerLike
|
||||
from vllm.tool_parsers.abstract_tool_parser import Tool
|
||||
|
||||
THINK_START = "<think>"
|
||||
THINK_END = "</think>"
|
||||
TOOL_CALL_START = "<tool_call>"
|
||||
TOOL_CALL_END = "</tool_call>"
|
||||
FUNC_PREFIX = "<function="
|
||||
FUNC_END = "</function>"
|
||||
PARAM_START = "<parameter="
|
||||
PARAM_END = "</parameter>"
|
||||
|
||||
_PARAM_RE = re.compile(
|
||||
r"<\s*parameter\s*=\s*([^>]*)>"
|
||||
r"(.*?)"
|
||||
r"(?:<\s*/\s*parameter\s*>|(?=<\s*parameter\s*=))",
|
||||
re.DOTALL,
|
||||
)
|
||||
_PARTIAL_PARAM_RE = re.compile(r"<\s*parameter\s*=\s*([^>]+)>(.*)$", re.DOTALL)
|
||||
|
||||
|
||||
def _qwen3_arg_converter(raw_args: str, partial: bool) -> str:
|
||||
params: dict[str, object] = {}
|
||||
|
||||
for match in _PARAM_RE.finditer(raw_args):
|
||||
name = match.group(1)
|
||||
value = match.group(2)
|
||||
params[name] = value.strip()
|
||||
|
||||
if partial:
|
||||
remaining = _PARAM_RE.sub("", raw_args)
|
||||
m = _PARTIAL_PARAM_RE.search(remaining)
|
||||
if m:
|
||||
name = m.group(1)
|
||||
value = m.group(2)
|
||||
if name:
|
||||
params[name] = value.strip()
|
||||
|
||||
return json.dumps(params, ensure_ascii=False)
|
||||
|
||||
|
||||
@functools.cache
|
||||
def qwen3_config(
|
||||
thinking: bool = True,
|
||||
*,
|
||||
name: str = "qwen3",
|
||||
think_start: str = THINK_START,
|
||||
think_end: str = THINK_END,
|
||||
tool_start: str = TOOL_CALL_START,
|
||||
tool_end: str = TOOL_CALL_END,
|
||||
) -> ParserEngineConfig:
|
||||
return ParserEngineConfig(
|
||||
name=name,
|
||||
initial_state=ParserState.REASONING if thinking else ParserState.CONTENT,
|
||||
terminals={
|
||||
# Reasoning terminals
|
||||
"THINK_START": think_start,
|
||||
"THINK_END": think_end,
|
||||
# Tool call terminals
|
||||
"TOOL_START": tool_start,
|
||||
"TOOL_END": tool_end,
|
||||
"FUNC_PREFIX": FUNC_PREFIX,
|
||||
"FUNC_END": FUNC_END,
|
||||
"PARAM_START": PARAM_START,
|
||||
"PARAM_END": PARAM_END,
|
||||
"CLOSE_ANGLE": ">",
|
||||
},
|
||||
token_id_terminals={
|
||||
"THINK_START": think_start,
|
||||
"THINK_END": think_end,
|
||||
"TOOL_START": tool_start,
|
||||
"TOOL_END": tool_end,
|
||||
},
|
||||
transitions={
|
||||
# -- Reasoning transitions --
|
||||
(ParserState.REASONING, "THINK_START"): Transition(
|
||||
ParserState.REASONING,
|
||||
(),
|
||||
),
|
||||
(ParserState.REASONING, "THINK_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(EventType.REASONING_END,),
|
||||
),
|
||||
# Absorb duplicate </think> — model may emit it after
|
||||
# already transitioning to CONTENT; drop it silently.
|
||||
(ParserState.CONTENT, "THINK_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(),
|
||||
),
|
||||
# Tool call directly from reasoning (implicit end)
|
||||
(ParserState.REASONING, "TOOL_START"): Transition(
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
(EventType.REASONING_END, EventType.TOOL_CALL_START),
|
||||
),
|
||||
# -- Tool call transitions --
|
||||
(ParserState.CONTENT, "TOOL_START"): Transition(
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
(EventType.REASONING_END, EventType.TOOL_CALL_START),
|
||||
),
|
||||
# Fallback: <function= without a preceding <tool_call>
|
||||
(ParserState.CONTENT, "FUNC_PREFIX"): Transition(
|
||||
ParserState.TOOL_NAME,
|
||||
(EventType.TOOL_CALL_START,),
|
||||
),
|
||||
(ParserState.TOOL_PREAMBLE, "TOOL_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(EventType.TOOL_CALL_END,),
|
||||
),
|
||||
(ParserState.TOOL_PREAMBLE, "FUNC_PREFIX"): Transition(
|
||||
ParserState.TOOL_NAME,
|
||||
(),
|
||||
),
|
||||
(ParserState.TOOL_NAME, "CLOSE_ANGLE"): Transition(
|
||||
ParserState.TOOL_ARGS,
|
||||
(),
|
||||
),
|
||||
# Malformed: </function> while still in TOOL_NAME (no closing >)
|
||||
(ParserState.TOOL_NAME, "FUNC_END"): Transition(
|
||||
ParserState.TOOL_BETWEEN,
|
||||
(EventType.TOOL_CALL_END,),
|
||||
),
|
||||
(ParserState.TOOL_ARGS, "FUNC_END"): Transition(
|
||||
ParserState.TOOL_BETWEEN,
|
||||
(EventType.TOOL_CALL_END,),
|
||||
),
|
||||
(ParserState.TOOL_ARGS, "PARAM_START"): Transition(
|
||||
ParserState.TOOL_ARGS,
|
||||
(EventType.ARG_VALUE_CHUNK,),
|
||||
),
|
||||
(ParserState.TOOL_ARGS, "PARAM_END"): Transition(
|
||||
ParserState.TOOL_ARGS,
|
||||
(EventType.ARG_VALUE_CHUNK,),
|
||||
),
|
||||
(ParserState.TOOL_BETWEEN, "TOOL_END"): Transition(
|
||||
ParserState.CONTENT,
|
||||
(),
|
||||
),
|
||||
# Consecutive tool call without closing </tool_call>
|
||||
(ParserState.TOOL_BETWEEN, "TOOL_START"): Transition(
|
||||
ParserState.TOOL_PREAMBLE,
|
||||
(EventType.TOOL_CALL_START,),
|
||||
),
|
||||
(ParserState.TOOL_BETWEEN, "FUNC_PREFIX"): Transition(
|
||||
ParserState.TOOL_NAME,
|
||||
(EventType.TOOL_CALL_START,),
|
||||
),
|
||||
},
|
||||
arg_converter=_qwen3_arg_converter,
|
||||
stream_arg_deltas=True,
|
||||
strip_trailing_reasoning_whitespace=False,
|
||||
tool_args_json=False,
|
||||
)
|
||||
|
||||
|
||||
class Qwen3Parser(ParserEngine):
|
||||
"""Qwen3 parser: ``<think>``/``</think>`` reasoning +
|
||||
``<tool_call>`` XML tool calls in a single engine.
|
||||
|
||||
- ``<tool_call>`` as implicit reasoning end
|
||||
- Unpaired ``<tool_call>`` token ID detection for ``is_reasoning_end``
|
||||
|
||||
Subclasses that share the grammar but differ only in the four wrapper
|
||||
token strings (reasoning + tool-call) override the class attributes
|
||||
below; everything else is inherited unchanged.
|
||||
"""
|
||||
|
||||
CONFIG_NAME = "qwen3"
|
||||
THINK_START = THINK_START
|
||||
THINK_END = THINK_END
|
||||
TOOL_START = TOOL_CALL_START
|
||||
TOOL_END = TOOL_CALL_END
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: TokenizerLike,
|
||||
tools: list[Tool] | None = None,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
chat_kwargs = kwargs.get("chat_template_kwargs", {}) or {}
|
||||
self.thinking_enabled = chat_kwargs.get("enable_thinking", True)
|
||||
kwargs.setdefault(
|
||||
"parser_engine_config",
|
||||
qwen3_config(
|
||||
thinking=self.thinking_enabled,
|
||||
name=self.CONFIG_NAME,
|
||||
think_start=self.THINK_START,
|
||||
think_end=self.THINK_END,
|
||||
tool_start=self.TOOL_START,
|
||||
tool_end=self.TOOL_END,
|
||||
),
|
||||
)
|
||||
super().__init__(
|
||||
tokenizer,
|
||||
tools,
|
||||
**kwargs,
|
||||
)
|
||||
vocab = self.vocab
|
||||
self._tool_call_token_id: int | None = vocab.get(self.TOOL_START)
|
||||
self._tool_call_end_token_id: int | None = vocab.get(self.TOOL_END)
|
||||
|
||||
def extract_reasoning(
|
||||
self,
|
||||
model_output: str,
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> tuple[str | None, str | None]:
|
||||
if not self.thinking_enabled:
|
||||
return None, model_output
|
||||
return super().extract_reasoning(model_output, request)
|
||||
|
||||
def is_reasoning_end(self, input_ids: list[int]) -> bool:
|
||||
if super().is_reasoning_end(input_ids):
|
||||
return True
|
||||
tool_call_id = self._tool_call_token_id
|
||||
tool_call_end_id = self._tool_call_end_token_id
|
||||
reasoning_start_id = self._reasoning_start_token_id
|
||||
if tool_call_id is not None:
|
||||
for i in range(len(input_ids) - 1, -1, -1):
|
||||
if (
|
||||
reasoning_start_id is not None
|
||||
and input_ids[i] == reasoning_start_id
|
||||
):
|
||||
return False
|
||||
if input_ids[i] == tool_call_id:
|
||||
if tool_call_end_id is not None and any(
|
||||
input_ids[j] == tool_call_end_id
|
||||
for j in range(i + 1, len(input_ids))
|
||||
):
|
||||
continue
|
||||
return True
|
||||
return False
|
||||
@@ -0,0 +1,28 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""seed_oss parser for tool calls and reasoning.
|
||||
|
||||
seed_oss shares the Qwen3 XML grammar exactly; only the four wrapper
|
||||
token strings differ::
|
||||
|
||||
<think> -> <seed:think>
|
||||
</think> -> </seed:think>
|
||||
<tool_call> -> <seed:tool_call>
|
||||
</tool_call> -> </seed:tool_call>
|
||||
|
||||
``<function=...>`` and ``<parameter=...>`` are byte-identical, so the
|
||||
entire transition table and ``_qwen3_arg_converter`` are inherited from
|
||||
:class:`Qwen3Parser` unchanged.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from vllm.parser.qwen3 import Qwen3Parser
|
||||
|
||||
|
||||
class SeedOssParser(Qwen3Parser):
|
||||
CONFIG_NAME = "seed_oss"
|
||||
THINK_START = "<seed:think>"
|
||||
THINK_END = "</seed:think>"
|
||||
TOOL_START = "<seed:tool_call>"
|
||||
TOOL_END = "</seed:tool_call>"
|
||||
@@ -0,0 +1,65 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
from collections.abc import Iterable, Sequence
|
||||
|
||||
from openai.types.responses import ResponseFunctionToolCall
|
||||
|
||||
from vllm.entrypoints.chat_utils import ChatCompletionMessageParam
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
|
||||
from vllm.entrypoints.openai.responses.protocol import (
|
||||
ResponseInputOutputItem,
|
||||
ResponsesRequest,
|
||||
)
|
||||
|
||||
|
||||
def count_tool_calls(tool_calls: object) -> int:
|
||||
if tool_calls is None:
|
||||
return 0
|
||||
if isinstance(tool_calls, (str, bytes, dict)):
|
||||
return 1
|
||||
if isinstance(tool_calls, Iterable):
|
||||
return sum(1 for _ in tool_calls)
|
||||
return 1
|
||||
|
||||
|
||||
def count_chat_history_tool_calls(
|
||||
messages: Sequence[ChatCompletionMessageParam],
|
||||
) -> int:
|
||||
return sum(
|
||||
count_tool_calls(msg.get("tool_calls"))
|
||||
for msg in messages
|
||||
if isinstance(msg, dict) and msg.get("role") == "assistant"
|
||||
)
|
||||
|
||||
|
||||
def count_response_history_tool_calls(
|
||||
response_items: Sequence[ResponseInputOutputItem],
|
||||
) -> int:
|
||||
count = 0
|
||||
for item in response_items:
|
||||
if isinstance(item, ResponseFunctionToolCall):
|
||||
count += 1
|
||||
continue
|
||||
|
||||
if isinstance(item, dict):
|
||||
item_type = item.get("type")
|
||||
if item_type == "function_call":
|
||||
count += 1
|
||||
elif item.get("role") == "assistant":
|
||||
count += count_tool_calls(item.get("tool_calls"))
|
||||
|
||||
return count
|
||||
|
||||
|
||||
def count_history_tool_calls(
|
||||
request: ChatCompletionRequest | ResponsesRequest,
|
||||
) -> int:
|
||||
if isinstance(request, ChatCompletionRequest):
|
||||
return count_chat_history_tool_calls(request.messages)
|
||||
|
||||
request_input = request.input
|
||||
if isinstance(request_input, str):
|
||||
return 0
|
||||
|
||||
return count_response_history_tool_calls(request_input)
|
||||
Reference in New Issue
Block a user