From 762ca52272049929128c9e699bfa7207c59a3fd0 Mon Sep 17 00:00:00 2001 From: Lyu Han Date: Mon, 6 Jul 2026 14:14:15 +0800 Subject: [PATCH 1/2] Stream tool parameters during response parsing (#4732) * feat: stream tool arguments incrementally * refactor: simplify XML tool parser streaming * fix: restore empty stream delta fallback * fix: suppress empty streaming deltas * perf: stream xml tool parser incrementally * refactor: stream json tool parser incrementally * refactor: simplify XML tool parser consumers * refactor: rename XML tool parser stream state * refactor: extract JSON tool parser base Co-authored-by: Cursor --------- Co-authored-by: Cursor --- lmdeploy/serve/openai/api_server.py | 76 ++- lmdeploy/serve/parsers/response_parser.py | 19 +- .../serve/parsers/tool_parser/__init__.py | 2 + .../parsers/tool_parser/glm47_tool_parser.py | 250 +++---- .../tool_parser/internlm2_tool_parser.py | 23 +- .../parsers/tool_parser/json_tool_parser.py | 301 +++++++++ .../parsers/tool_parser/llama3_tool_parser.py | 21 +- .../tool_parser/qwen2d5_tool_parser.py | 22 +- .../parsers/tool_parser/qwen3_tool_parser.py | 21 +- .../tool_parser/qwen3coder_tool_parser.py | 242 ++++--- .../serve/parsers/tool_parser/tool_parser.py | 93 +-- .../parsers/tool_parser/xml_tool_parser.py | 236 +++++-- .../test_arguments_validation.py | 44 +- .../test_delta_tool_call_id.py | 107 ++- .../test_streaming_metadata.py | 247 +++++++ .../serve/parsers/test_glm47_parser.py | 620 ++++++++++++++++-- .../serve/parsers/test_llama3_parser.py | 24 + .../serve/parsers/test_qwen3_5_parser.py | 513 ++++++++++++--- .../serve/parsers/test_qwen3_parser.py | 113 +++- 19 files changed, 2224 insertions(+), 750 deletions(-) create mode 100644 lmdeploy/serve/parsers/tool_parser/json_tool_parser.py create mode 100644 tests/test_lmdeploy/serve/openai/chat_completions/test_streaming_metadata.py diff --git a/lmdeploy/serve/openai/api_server.py b/lmdeploy/serve/openai/api_server.py index 9fe802ca1f..f2849c8c43 100644 --- a/lmdeploy/serve/openai/api_server.py +++ b/lmdeploy/serve/openai/api_server.py @@ -9,6 +9,7 @@ import time from collections.abc import AsyncGenerator from contextlib import aclosing, asynccontextmanager +from dataclasses import dataclass from functools import partial from http import HTTPStatus from typing import TYPE_CHECKING, Literal @@ -289,6 +290,44 @@ def _create_output_token_logprobs(token_ids: list[int] | None = None, return output_token_logprobs or None +@dataclass +class _StreamTokenMetadata: + """Token metadata buffered across parser steps with no visible delta.""" + + token_ids: list[int] + logprobs: list[dict[int, float]] + + @classmethod + def from_result(cls, token_ids: list[int] | None, logprobs: list[dict[int, float]] | None): + return cls(list(token_ids or []), list(logprobs or [])) + + def extend(self, other: _StreamTokenMetadata) -> None: + self.token_ids.extend(other.token_ids) + self.logprobs.extend(other.logprobs) + + def pop_with(self, current: _StreamTokenMetadata) -> _StreamTokenMetadata: + merged = _StreamTokenMetadata( + token_ids=self.token_ids + current.token_ids, + logprobs=self.logprobs + current.logprobs, + ) + self.token_ids.clear() + self.logprobs.clear() + return merged + + def output_ids(self, enabled: bool | None) -> list[int] | None: + return self.token_ids if enabled else None + + def chat_logprobs(self, tokenizer: PreTrainedTokenizerBase, enabled: bool | None) -> ChoiceLogprobs | None: + if not enabled or not self.token_ids or not self.logprobs: + return None + return _create_chat_completion_logprobs(tokenizer, self.token_ids, self.logprobs) + + def output_token_logprobs(self, enabled: bool | None) -> list[tuple[float, int]] | None: + if not enabled: + return None + return _create_output_token_logprobs(self.token_ids, self.logprobs) + + @router.get('/health') async def health() -> JSONResponse: """Health check.""" @@ -587,31 +626,25 @@ def create_stream_usage_response_json(usage: UsageInfo) -> str: async def completion_stream_generator() -> AsyncGenerator[str, None]: streaming_tools = False final_usage = None + pending_token_metadata = _StreamTokenMetadata.from_result(None, None) async for res in result_generator: - logprobs = None - output_token_logprobs = None - if request.logprobs and res.logprobs: - logprobs = _create_chat_completion_logprobs(tokenizer, res.token_ids, res.logprobs) - if request.return_logprob: - output_token_logprobs = _create_output_token_logprobs(res.token_ids, res.logprobs) if res.finish_reason and include_usage: final_usage = UsageInfo.build( prompt_tokens=res.input_token_len, completion_tokens=res.generate_token_len, cached_tokens=res.cached_tokens, ) - delta_token_ids = res.token_ids if res.token_ids is not None else [] + current_token_metadata = _StreamTokenMetadata.from_result(res.token_ids, res.logprobs) stream_deltas = response_parser.stream_chunk( res.response, - delta_token_ids + current_token_metadata.token_ids ) if not stream_deltas: - # Parser may buffer partial protocol tags and emit no visible delta - # while the engine still produced new tokens (e.g. MTP batch). Do not - # drop those token ids; emit them once on a placeholder delta. - if res.finish_reason is None and not delta_token_ids: + pending_token_metadata.extend(current_token_metadata) + if res.finish_reason is None: continue - stream_deltas = [(DeltaMessage(role='assistant', content=''), False)] + current_token_metadata = _StreamTokenMetadata.from_result(None, None) + stream_deltas = [(DeltaMessage(role='assistant'), False)] should_validate_complete = ( res.finish_reason in ('stop', 'length') and (request.return_token_ids or request.return_routed_experts) @@ -627,8 +660,15 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]: # The chat parser may split one engine yield into multiple protocol deltas, # so attach the engine-level metadata to the last parsed delta. finish_reason = res.finish_reason if is_last_delta else None - chunk_logprobs = logprobs if is_last_delta else None - chunk_output_token_logprobs = output_token_logprobs if is_last_delta else None + if is_last_delta: + stream_token_metadata = pending_token_metadata.pop_with(current_token_metadata) + chunk_logprobs = stream_token_metadata.chat_logprobs(tokenizer, request.logprobs) + chunk_output_token_logprobs = stream_token_metadata.output_token_logprobs(request.return_logprob) + stream_output_ids = stream_token_metadata.output_ids(request.return_token_ids) + else: + chunk_logprobs = None + chunk_output_token_logprobs = None + stream_output_ids = None if (request.tool_choice != 'none' and response_parser.tool_parser is not None): if finish_reason == 'stop' and streaming_tools is True: @@ -636,10 +676,6 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]: # Only output routed_experts in the final chunk routed_experts = res.routed_experts if finish_reason is not None else None - # Emit token ids once per engine yield on the last parsed delta, when - # accumulated delta text and token ids for this step are aligned. - stream_output_ids = delta_token_ids if (request.return_token_ids and is_last_delta) else None - response_json = create_stream_response_json(index=0, delta_message=delta_message, finish_reason=finish_reason, @@ -649,7 +685,7 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]: output_ids=stream_output_ids) if res.cache_block_ids is not None and is_last_delta: response_json['cache_block_ids'] = res.cache_block_ids - response_json['remote_token_ids'] = res.token_ids + response_json['remote_token_ids'] = current_token_metadata.token_ids yield f'data: {json.dumps(response_json)}\n\n' if final_usage is not None: yield f'data: {create_stream_usage_response_json(final_usage)}\n\n' diff --git a/lmdeploy/serve/parsers/response_parser.py b/lmdeploy/serve/parsers/response_parser.py index d000e7a7d6..c762008d11 100644 --- a/lmdeploy/serve/parsers/response_parser.py +++ b/lmdeploy/serve/parsers/response_parser.py @@ -177,7 +177,7 @@ def stream_chunk(self, Returns: A list of ``(delta_message, tool_calls_emitted)`` pairs. Return ``[]`` when this engine step produces no visible delta (for example - while buffering a partial protocol tag). + while buffering protocol syntax or tool-call payload). """ raise NotImplementedError @@ -322,7 +322,7 @@ def stream_chunk( from this stream step. Multiple entries may be returned when one engine chunk contains reasoning, content, and tool-call segments. Return ``[]`` when this engine step produces no visible delta (for - example while buffering a partial protocol tag). + example while buffering protocol syntax or tool-call payload). """ # Special-case: some backends emit a leading empty delta (no text, no # tokens) before any actual content. Tests treat this as a visible empty @@ -555,8 +555,7 @@ def _consume_tool(self) -> tuple[list[DeltaToolCall], bool]: emit = self._pending self._pending = '' out = self.tool_parser.decode_tool_incremental(added_text=emit, final=False) - if (self.profile.tool_payload_format == 'json' - and self._is_complete_json_object(self.tool_parser._tool_payload)): + if self.profile.tool_payload_format == 'json' and self.tool_parser._payload_closed: out.extend(self.tool_parser.decode_tool_incremental(added_text='', final=True)) self.tool_parser.finish_tool_call() self._mode = self.MODE_PLAIN @@ -735,15 +734,3 @@ def _longest_open_tag_prefix_suffix(text: str, tags: list[str]) -> int: best = k break return best - - @staticmethod - def _is_complete_json_object(payload: str) -> bool: - payload = payload.strip() - if not payload: - return False - decoder = json.JSONDecoder() - try: - obj, end = decoder.raw_decode(payload) - except json.JSONDecodeError: - return False - return isinstance(obj, dict) and end == len(payload) diff --git a/lmdeploy/serve/parsers/tool_parser/__init__.py b/lmdeploy/serve/parsers/tool_parser/__init__.py index f9c5ceaba7..b0335b3ae1 100644 --- a/lmdeploy/serve/parsers/tool_parser/__init__.py +++ b/lmdeploy/serve/parsers/tool_parser/__init__.py @@ -4,6 +4,7 @@ from .glm47_tool_parser import Glm47ToolParser from .internlm2_tool_parser import Internlm2ToolParser from .interns2preview_tool_parser import InternS2PreviewToolParser +from .json_tool_parser import JsonToolParser from .llama3_tool_parser import Llama3JsonToolParser from .qwen2d5_tool_parser import Qwen2d5ToolParser from .qwen3_tool_parser import Qwen3ToolParser @@ -14,6 +15,7 @@ __all__ = [ 'ToolParser', 'ToolParserManager', + 'JsonToolParser', 'XmlToolParser', 'DeepSeekV32ToolParser', 'DeepSeekV4ToolParser', diff --git a/lmdeploy/serve/parsers/tool_parser/glm47_tool_parser.py b/lmdeploy/serve/parsers/tool_parser/glm47_tool_parser.py index c91655917d..3b57658c14 100644 --- a/lmdeploy/serve/parsers/tool_parser/glm47_tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/glm47_tool_parser.py @@ -10,7 +10,7 @@ ) from .tool_parser import ToolParserManager -from .xml_tool_parser import XmlToolParser # type: ignore[reportMissingImports] +from .xml_tool_parser import XmlToolParser @ToolParserManager.register_module(['glm47']) @@ -21,10 +21,6 @@ class Glm47ToolParser(XmlToolParser): ``function_namekv...`` """ - arg_key_start_token = '' - arg_key_end_token = '' - arg_value_start_token = '' - arg_value_end_token = '' _complete_payload_pattern = re.compile( r'^\s*[^\s<]+(?:\s*[^<]+\s*.*?)*\s*$', re.DOTALL, @@ -33,12 +29,12 @@ class Glm47ToolParser(XmlToolParser): def _reset_incremental_state(self) -> None: self._func_name: str | None = None self._args: dict[str, str] = {} - self._open_arg_key: str | None = None - # Offset in accumulated args text where the in-flight value begins - # (first char after ````); -1 when none is open. - self._value_start = -1 - # Resume arg scanning after the last completed ````. - self._scan_pos = 0 + self._arg_name: str | None = None + self._value_parts: list[str] = [] + self._phase = 'function' + self._stream_started = False + self._stream_pending_ws = '' + self._stream_blocked = False @classmethod def get_tool_open_tag(cls) -> str | None: @@ -52,35 +48,126 @@ def get_tool_close_tag(cls) -> str | None: def get_tool_payload_format(cls) -> str: return 'xml' - def _extract_incremental_state(self, - payload: str, - final: bool = False) -> tuple[str | None, dict[str, str], bool]: - """Update streaming parse state from accumulated inner tool payload. + def _reset_value_stream_state(self) -> None: + self._stream_started = False + self._stream_pending_ws = '' + self._stream_blocked = False + + def _stream_arg_delta(self, raw: str) -> str: + if self._stream_blocked: + return '' + + schema_type = self._get_param_schema_type(self._func_name, self._arg_name or '') + if schema_type not in (None, 'string'): + self._stream_blocked = True + return '' + + text = self._stream_pending_ws + raw + self._stream_pending_ws = '' + + if schema_type == 'string': + if not self._stream_started: + text = text.lstrip() + if not text: + return '' + if text.startswith('"'): + self._stream_blocked = True + return '' + + stable = text.rstrip() + self._stream_pending_ws = text[len(stable):] + if not stable: + return '' + + self._stream_started = True + return stable + + if not self._stream_started: + stripped = text.lstrip() + if not stripped: + self._stream_pending_ws = text + return '' + if stripped.startswith('"'): + self._stream_blocked = True + return '' + + self._stream_started = True + return text + + def _consume_function(self, payload: str, pos: int, final: bool) -> int | None: + arg_key_start = payload.find('', pos) + if arg_key_start >= 0: + name = payload[pos:arg_key_start].strip() + if name: + self._func_name = name + self._phase = 'arg_start' + return arg_key_start + + remaining = payload[pos:] + if final and remaining.strip(): + self._func_name = remaining.strip() + return len(payload) + return None + + def _consume_arg_start(self, payload: str, pos: int) -> int | None: + arg_key_start = payload.find('', pos) + if arg_key_start < 0: + return None - See :meth:`Qwen3CoderToolParser._extract_incremental_state` for the - contract. GLM-4.7 uses ``func_namekv - `` instead of Qwen3Coder XML; ``is_closed`` is always - ``False`` because argument JSON is closed by the outer ````. - """ - payload = payload.strip() - if not payload: - return self._func_name, dict(self._args), False + self._phase = 'arg_name' + return arg_key_start + len('') - args_start_idx = payload.find(self.arg_key_start_token) - if args_start_idx >= 0: - func_name = payload[:args_start_idx].strip() - if func_name: - self._func_name = func_name - self._parse_args_incremental(payload[args_start_idx:]) - elif final: - func_name = payload.strip() - if func_name: - self._func_name = func_name + def _consume_arg_name(self, payload: str, pos: int) -> int | None: + key_end = payload.find('', pos) + if key_end < 0: + return None - return self._func_name, dict(self._args), False + value_start = payload.find('', key_end + len('')) + if value_start < 0: + return None + + self._arg_name = payload[pos:key_end].strip() + self._value_parts.clear() + self._reset_value_stream_state() + self._phase = 'arg_value' + return value_start + len('') + + def _consume_arg_value(self, payload: str, pos: int, arg_delta_parts: list[str]) -> tuple[int | None, bool]: + """Consume an argument value. + + Returns ``(next_pos, should_stop)``. ``should_stop`` is true after + streaming an open value delta, because the next bytes may be the + argument close tag and must be checked with the next chunk. + """ + value_end = payload.find('', pos) + + if value_end >= 0: + raw = payload[pos:value_end] + if raw: + self._value_parts.append(raw) + if self._arg_name: + self._args[self._arg_name] = ''.join(self._value_parts) + self._arg_name = None + self._value_parts.clear() + self._reset_value_stream_state() + self._phase = 'function' + return value_end + len(''), False + + # Open value: keep any partial "" suffix buffered instead + # of emitting it as argument text. + raw_end = self._trim_partial_close_tag_suffix(payload, pos, '') + if raw_end == pos: + return None, True + + raw_delta = payload[pos:raw_end] + self._value_parts.append(raw_delta) + stream_delta = self._stream_arg_delta(raw_delta) + if stream_delta: + arg_delta_parts.append(stream_delta) + return raw_end, True def parse_tool_call_complete(self, payload: str) -> ToolCall | None: - func_name, raw_args_dict = self._parse_payload(payload, final=True) + func_name, raw_args_dict = self._extract_complete_args(payload) if not func_name: return None args_dict = self._get_coerced_args(func_name, raw_args_dict, use_cache=False) @@ -89,91 +176,16 @@ def parse_tool_call_complete(self, payload: str) -> ToolCall | None: def _validate_tool_payload(self, payload: str) -> bool: return bool(self._complete_payload_pattern.fullmatch(payload)) - def _complete_open_arg_if_ready(self, args_text: str) -> bool: - """Finalize the in-flight arg once ```` is available. - - Uses ``_value_start`` so we can locate the closing tag without re-parsing - the ```` / ```` headers on every non-fast-path chunk. - """ - if self._value_start < 0 or not self._open_arg_key: - return False - value_end = args_text.find(self.arg_value_end_token, self._value_start) - if value_end < 0: - self._in_progress_value = True - return False - self._args[self._open_arg_key] = args_text[self._value_start:value_end] - self._open_arg_key = None - self._value_start = -1 - self._in_progress_value = False - self._scan_pos = value_end + len(self.arg_value_end_token) - return True - - def _parse_args_incremental(self, args_text: str) -> None: - """Scan ``kv`` blocks and - update ``_args``. - - Incomplete arg headers or values are left open in ``_open_arg_key`` / - ``_value_start`` until the closing tag arrives in a later stream chunk. - - ``_scan_pos`` only advances past completed ```` tags; while a - value is streaming, it stays before the open ````. The block - below is an optimization (not required for correctness): skip the - while-loop header re-scan and try to close the current value directly. - If ```` is still missing, return early because the while - loop would reach the same open-tag state anyway. - """ - if self._value_start >= 0: - if not self._complete_open_arg_if_ready(args_text): - return - - while True: - key_start = args_text.find(self.arg_key_start_token, self._scan_pos) - if key_start < 0: - self._in_progress_value = False - return - - key_content_start = key_start + len(self.arg_key_start_token) - key_end = args_text.find(self.arg_key_end_token, key_content_start) - if key_end < 0: - self._in_progress_value = True - return - - key = args_text[key_content_start:key_end].strip() - value_start = args_text.find(self.arg_value_start_token, key_end + len(self.arg_key_end_token)) - if value_start < 0: - self._in_progress_value = True - return - - value_content_start = value_start + len(self.arg_value_start_token) - value_end = args_text.find(self.arg_value_end_token, value_content_start) - if value_end < 0: - self._open_arg_key = key - self._value_start = value_content_start - self._in_progress_value = True - return - - next_pos = value_end + len(self.arg_value_end_token) - if key in self._args: - self._scan_pos = next_pos - continue - - if key: - self._args[key] = args_text[value_content_start:value_end] - self._scan_pos = next_pos - self._in_progress_value = False - - def _parse_payload(self, payload: str, *, final: bool = False) -> tuple[str | None, dict[str, str]]: + def _extract_complete_args(self, payload: str) -> tuple[str | None, dict[str, str]]: payload = payload.strip() if not payload: return None, {} - args_start_idx = payload.find(self.arg_key_start_token) + args_start_idx = payload.find('') if args_start_idx >= 0: func_name = payload[:args_start_idx].strip() args_text = payload[args_start_idx:] else: - if not final: - return None, {} func_name = payload.strip() args_text = '' if not func_name: @@ -182,22 +194,22 @@ def _parse_payload(self, payload: str, *, final: bool = False) -> tuple[str | No args_dict: dict[str, str] = {} search_idx = 0 while True: - key_start = args_text.find(self.arg_key_start_token, search_idx) + key_start = args_text.find('', search_idx) if key_start < 0: break - key_content_start = key_start + len(self.arg_key_start_token) - key_end = args_text.find(self.arg_key_end_token, key_content_start) + key_content_start = key_start + len('') + key_end = args_text.find('', key_content_start) if key_end < 0: break key = args_text[key_content_start:key_end].strip() - value_start = args_text.find(self.arg_value_start_token, key_end + len(self.arg_key_end_token)) + value_start = args_text.find('', key_end + len('')) if value_start < 0: break - value_content_start = value_start + len(self.arg_value_start_token) - value_end = args_text.find(self.arg_value_end_token, value_content_start) + value_content_start = value_start + len('') + value_end = args_text.find('', value_content_start) if value_end < 0: break if key: args_dict[key] = args_text[value_content_start:value_end] - search_idx = value_end + len(self.arg_value_end_token) + search_idx = value_end + len('') return func_name, args_dict diff --git a/lmdeploy/serve/parsers/tool_parser/internlm2_tool_parser.py b/lmdeploy/serve/parsers/tool_parser/internlm2_tool_parser.py index c926602e01..1035a9b7d4 100644 --- a/lmdeploy/serve/parsers/tool_parser/internlm2_tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/internlm2_tool_parser.py @@ -3,17 +3,15 @@ from typing import TYPE_CHECKING -from .tool_parser import ToolParser, ToolParserManager +from .json_tool_parser import JsonToolParser +from .tool_parser import ToolParserManager if TYPE_CHECKING: - from lmdeploy.serve.openai.protocol import ( - ChatCompletionRequest, - DeltaToolCall, - ToolCall, - ) + from lmdeploy.serve.openai.protocol import ChatCompletionRequest + @ToolParserManager.register_module(['internlm', 'intern-s1']) -class Internlm2ToolParser(ToolParser): +class Internlm2ToolParser(JsonToolParser): """Tool parser for InternLM JSON tool-call payloads.""" def adjust_request(self, request: ChatCompletionRequest) -> ChatCompletionRequest: @@ -32,14 +30,3 @@ def get_tool_open_tag(cls) -> str | None: @classmethod def get_tool_close_tag(cls) -> str | None: return '<|action_end|>' - - @classmethod - def get_tool_payload_format(cls) -> str: - return 'json' - - def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[DeltaToolCall]: - """Decode incremental JSON tool payload.""" - return self._decode_tool_incremental_json(added_text=added_text, final=final) - - def parse_tool_call_complete(self, payload: str) -> ToolCall | None: - return self._parse_tool_call_complete_json(payload) diff --git a/lmdeploy/serve/parsers/tool_parser/json_tool_parser.py b/lmdeploy/serve/parsers/tool_parser/json_tool_parser.py new file mode 100644 index 0000000000..b7166a3abe --- /dev/null +++ b/lmdeploy/serve/parsers/tool_parser/json_tool_parser.py @@ -0,0 +1,301 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from __future__ import annotations + +import json +from dataclasses import dataclass + +from lmdeploy.serve.openai.protocol import ( + DeltaFunctionCall, + DeltaToolCall, + FunctionCall, + ToolCall, +) + +from .tool_parser import ToolParser + + +@dataclass +class JsonToolSnapshot: + func_name: str | None + args_delta: str + + +class JsonToolParser(ToolParser): + """Base class for JSON tool-call payload parsers.""" + + def __init__(self): + super().__init__() + self._payload: str = '' + self._phase: str = 'payload_start' + self._json_key: str | None = None + self._value_type: str | None = None + self._value_depth: int = 0 + self._string_open_in_container: bool = False + self._value_escaped: bool = False + self._payload_closed: bool = False + + @classmethod + def get_tool_payload_format(cls) -> str: + return 'json' + + def start_tool_call(self) -> None: + super().start_tool_call() + self._reset_stream_state() + + def finish_tool_call(self) -> None: + super().finish_tool_call() + self._reset_stream_state() + + def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[DeltaToolCall]: + """Stream raw JSON tool argument text without requiring completion. + + ``_payload`` only keeps unconsumed syntax such as a partial key or + delimiter. Argument value bytes are emitted as deltas once they are + observed, while string/container state is tracked separately. + """ + self._payload += added_text + snapshot, consumed = self._consume_payload(self._payload, final=final) + if consumed > 0: + self._payload = self._payload[consumed:] + + out: list[DeltaToolCall] = [] + if snapshot.func_name and not self._name_emitted: + out.append( + DeltaToolCall( + id=self._active_tool_call_id, + index=self._active_tool_index, + type='function', + function=DeltaFunctionCall(name=snapshot.func_name), + )) + self._name_emitted = True + if snapshot.args_delta: + out.append( + DeltaToolCall( + id=None, + index=self._active_tool_index, + type=None, + function=DeltaFunctionCall(arguments=snapshot.args_delta), + )) + return out + + def parse_tool_call_complete(self, payload: str) -> ToolCall | None: + if not payload: + return None + try: + obj = json.loads(payload) + except json.JSONDecodeError: + return None + if not isinstance(obj, dict): + return None + name = obj.get('name') + if not isinstance(name, str) or not name: + return None + args_obj = obj.get('arguments', obj.get('parameters', {})) + args_json = json.dumps(args_obj, ensure_ascii=False) + return ToolCall(function=FunctionCall(name=name, arguments=args_json)) + + def _validate_tool_payload(self, payload: str) -> bool: + try: + obj = json.loads(payload) + except json.JSONDecodeError: + return False + if not isinstance(obj, dict): + return False + name = obj.get('name') + return isinstance(name, str) and bool(name) + + def _consume_payload(self, payload: str, *, final: bool) -> tuple[JsonToolSnapshot, int]: + pos = 0 + args_delta_parts: list[str] = [] + func_name: str | None = None + n = len(payload) + + while pos < n: + if self._phase == 'payload_start': + pos = self._skip_ws(payload, pos) + if pos >= n: + break + if payload[pos] != '{': + break + pos += 1 + self._phase = 'key' + continue + + if self._phase == 'key': + pos = self._skip_key_prefix(payload, pos) + if pos >= n: + break + if payload[pos] == '}': + pos += 1 + self._phase = 'done' + self._payload_closed = True + break + if payload[pos] != '"': + break + key, end = self._read_string(payload, pos) + if end < 0: + break + self._json_key = key + pos = end + self._phase = 'colon' + continue + + if self._phase == 'colon': + pos = self._skip_ws(payload, pos) + if pos >= n: + break + if payload[pos] != ':': + break + pos += 1 + self._phase = 'value_start' + continue + + if self._phase == 'value_start': + pos = self._skip_ws(payload, pos) + if pos >= n: + break + if self._json_key == 'name': + if payload[pos] == '"': + name, end = self._read_string(payload, pos) + if end < 0: + break + func_name = name + pos = end + self._json_key = None + self._phase = 'key' + continue + self._phase = 'skip_value' + continue + if self._json_key in ('arguments', 'parameters'): + self._phase = 'args_value' + continue + self._phase = 'skip_value' + continue + + if self._phase == 'args_value': + delta, consumed, complete = self._consume_value(payload, pos, emit=True) + if delta: + args_delta_parts.append(delta) + pos += consumed + if complete: + self._json_key = None + self._phase = 'key' + self._reset_value_state() + continue + break + + if self._phase == 'skip_value': + _, consumed, complete = self._consume_value(payload, pos, emit=False) + pos += consumed + if complete: + self._json_key = None + self._phase = 'key' + self._reset_value_state() + continue + break + + break + + return JsonToolSnapshot(func_name, ''.join(args_delta_parts)), pos + + def _consume_value(self, payload: str, start: int, *, emit: bool) -> tuple[str, int, bool]: + if start >= len(payload): + return '', 0, False + + i = start + if self._value_type is None: + ch = payload[i] + if ch == '"': + self._value_type = 'string' + self._value_escaped = False + i += 1 + elif ch in '{[': + self._value_type = 'container' + self._value_depth = 1 + self._string_open_in_container = False + self._value_escaped = False + i += 1 + else: + self._value_type = 'scalar' + + complete = False + while i < len(payload): + ch = payload[i] + if self._value_type == 'string': + if self._value_escaped: + self._value_escaped = False + elif ch == '\\': + self._value_escaped = True + elif ch == '"': + i += 1 + complete = True + break + elif self._value_type == 'container': + if self._string_open_in_container: + if self._value_escaped: + self._value_escaped = False + elif ch == '\\': + self._value_escaped = True + elif ch == '"': + self._string_open_in_container = False + elif ch == '"': + self._string_open_in_container = True + elif ch in '{[': + self._value_depth += 1 + elif ch in '}]': + self._value_depth -= 1 + if self._value_depth == 0: + i += 1 + complete = True + break + elif ch in ',}]': + complete = True + break + i += 1 + + consumed = i - start + return payload[start:i] if emit else '', consumed, complete + + @staticmethod + def _read_string(payload: str, start: int) -> tuple[str, int]: + chars: list[str] = [] + escaped = False + i = start + 1 + while i < len(payload): + ch = payload[i] + if escaped: + chars.append(ch) + escaped = False + elif ch == '\\': + escaped = True + elif ch == '"': + return ''.join(chars), i + 1 + else: + chars.append(ch) + i += 1 + return '', -1 + + @staticmethod + def _skip_ws(payload: str, pos: int) -> int: + while pos < len(payload) and payload[pos].isspace(): + pos += 1 + return pos + + @staticmethod + def _skip_key_prefix(payload: str, pos: int) -> int: + while pos < len(payload) and (payload[pos].isspace() or payload[pos] == ','): + pos += 1 + return pos + + def _reset_stream_state(self) -> None: + self._payload = '' + self._phase = 'payload_start' + self._json_key = None + self._payload_closed = False + self._reset_value_state() + + def _reset_value_state(self) -> None: + self._value_type = None + self._value_depth = 0 + self._string_open_in_container = False + self._value_escaped = False diff --git a/lmdeploy/serve/parsers/tool_parser/llama3_tool_parser.py b/lmdeploy/serve/parsers/tool_parser/llama3_tool_parser.py index 60c2785a68..44a94815a4 100644 --- a/lmdeploy/serve/parsers/tool_parser/llama3_tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/llama3_tool_parser.py @@ -1,15 +1,11 @@ # Copyright (c) OpenMMLab. All rights reserved. -from lmdeploy.serve.openai.protocol import ( - DeltaToolCall, - ToolCall, -) - -from .tool_parser import ToolParser, ToolParserManager +from .json_tool_parser import JsonToolParser +from .tool_parser import ToolParserManager @ToolParserManager.register_module('llama3') -class Llama3JsonToolParser(ToolParser): +class Llama3JsonToolParser(JsonToolParser): """Tool parser for Llama3 JSON tool-call payloads.""" @classmethod @@ -20,16 +16,5 @@ def get_tool_open_tag(cls) -> str | None: def get_tool_close_tag(cls) -> str | None: return None - @classmethod - def get_tool_payload_format(cls) -> str: - return 'json' - - def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[DeltaToolCall]: - """Decode incremental JSON tool payload.""" - return self._decode_tool_incremental_json(added_text=added_text, final=final) - - def parse_tool_call_complete(self, payload: str) -> ToolCall | None: - return self._parse_tool_call_complete_json(payload) - def validate_complete(self, text: str) -> bool: return True diff --git a/lmdeploy/serve/parsers/tool_parser/qwen2d5_tool_parser.py b/lmdeploy/serve/parsers/tool_parser/qwen2d5_tool_parser.py index 8c50805bb9..9e6c5ab8ac 100644 --- a/lmdeploy/serve/parsers/tool_parser/qwen2d5_tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/qwen2d5_tool_parser.py @@ -1,16 +1,11 @@ # Copyright (c) OpenMMLab. All rights reserved. - -from lmdeploy.serve.openai.protocol import ( - DeltaToolCall, - ToolCall, -) - -from .tool_parser import ToolParser, ToolParserManager +from .json_tool_parser import JsonToolParser +from .tool_parser import ToolParserManager @ToolParserManager.register_module(['qwen2d5']) -class Qwen2d5ToolParser(ToolParser): +class Qwen2d5ToolParser(JsonToolParser): """Tool parser for Qwen2.5 JSON tool-call payloads.""" @classmethod @@ -20,14 +15,3 @@ def get_tool_open_tag(cls) -> str | None: @classmethod def get_tool_close_tag(cls) -> str | None: return '' - - @classmethod - def get_tool_payload_format(cls) -> str: - return 'json' - - def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[DeltaToolCall]: - """Decode incremental JSON tool payload.""" - return self._decode_tool_incremental_json(added_text=added_text, final=final) - - def parse_tool_call_complete(self, payload: str) -> ToolCall | None: - return self._parse_tool_call_complete_json(payload) diff --git a/lmdeploy/serve/parsers/tool_parser/qwen3_tool_parser.py b/lmdeploy/serve/parsers/tool_parser/qwen3_tool_parser.py index b76e6c72df..caae16e23f 100644 --- a/lmdeploy/serve/parsers/tool_parser/qwen3_tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/qwen3_tool_parser.py @@ -1,15 +1,11 @@ # Copyright (c) OpenMMLab. All rights reserved. -from lmdeploy.serve.openai.protocol import ( - DeltaToolCall, - ToolCall, -) - -from .tool_parser import ToolParser, ToolParserManager +from .json_tool_parser import JsonToolParser +from .tool_parser import ToolParserManager @ToolParserManager.register_module(['qwen', 'qwen3']) -class Qwen3ToolParser(ToolParser): +class Qwen3ToolParser(JsonToolParser): """Tool parser for Qwen3 JSON tool-call payloads.""" @classmethod @@ -19,14 +15,3 @@ def get_tool_open_tag(cls) -> str | None: @classmethod def get_tool_close_tag(cls) -> str | None: return '' - - @classmethod - def get_tool_payload_format(cls) -> str: - return 'json' - - def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[DeltaToolCall]: - """Decode incremental JSON tool payload.""" - return self._decode_tool_incremental_json(added_text=added_text, final=final) - - def parse_tool_call_complete(self, payload: str) -> ToolCall | None: - return self._parse_tool_call_complete_json(payload) diff --git a/lmdeploy/serve/parsers/tool_parser/qwen3coder_tool_parser.py b/lmdeploy/serve/parsers/tool_parser/qwen3coder_tool_parser.py index 6cd343bceb..10d987a1f5 100644 --- a/lmdeploy/serve/parsers/tool_parser/qwen3coder_tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/qwen3coder_tool_parser.py @@ -3,7 +3,6 @@ import json import re -from typing import Any from lmdeploy.serve.openai.protocol import ( FunctionCall, @@ -18,10 +17,6 @@ class Qwen3CoderToolParser(XmlToolParser): """Tool parser for Qwen3Coder XML tool-call payloads.""" - func_prefix = '\n]+>\s*(?:\n]+>.*?\s*)*\s*$', re.DOTALL, @@ -29,14 +24,13 @@ class Qwen3CoderToolParser(XmlToolParser): def _reset_incremental_state(self) -> None: self._func_name: str | None = None - self._args: dict[str, Any] = {} - self._func_closed = False - self._open_param_name: str | None = None - # Offset in accumulated payload where the in-flight parameter value begins - # (first char after ``>`` in ````); -1 when none is open. - self._value_start = -1 - # Resume parameter scanning after the last completed ````. - self._scan_pos = 0 + self._args: dict[str, str] = {} + self._arg_name: str | None = None + self._value_parts: list[str] = [] + self._phase = 'function' + self._stream_started = False + self._stream_pending_ws = '' + self._stream_blocked = False # Qwen3Coder closes tool argument JSON only when the model emits the # explicit function end marker (). We intentionally avoid @@ -57,117 +51,112 @@ def get_tool_close_tag(cls) -> str | None: def get_tool_payload_format(cls) -> str: return 'xml' - def _extract_incremental_state(self, - payload: str, - final: bool = False) -> tuple[str | None, dict[str, Any], bool]: - """Update streaming parse state from accumulated inner tool payload. + def _reset_value_stream_state(self) -> None: + self._stream_started = False + self._stream_pending_ws = '' + self._stream_blocked = False + + def _stream_arg_delta(self, raw: str) -> str: + if self._stream_blocked: + return '' + + schema_type = self._get_param_schema_type(self._func_name, self._arg_name or '') + if schema_type not in (None, 'string'): + self._stream_blocked = True + return '' + + text = self._stream_pending_ws + raw + self._stream_pending_ws = '' + + if not self._stream_started: + text = text.lstrip() + if not text: + return '' + if text.startswith('"'): + self._stream_blocked = True + return '' + + stable = text.rstrip() + self._stream_pending_ws = text[len(stable):] + if not stable: + return '' + + self._stream_started = True + return stable + + def _consume_function(self, payload: str, pos: int, final: bool) -> int | None: + start = payload.find('...`` (outer tags - are stripped by :class:`BaseResponseParser` before tool mode). This - method mutates incremental parse state across chunks and returns the current - snapshot for :meth:`XmlToolParser.decode_tool_incremental`. + name_start = start + len('', name_start) + if name_end < 0: + return None - Returns: - ``(func_name, args_dict, is_func_closed)`` where: + self._func_name = payload[name_start:name_end].strip() + self._phase = 'arg_start' + return name_end + 1 - - ``func_name``: callee parsed from ````, or ``None`` - - ``args_dict``: parameters whose ```` has been seen - - ``is_func_closed``: whether ```` is present; used to - emit the closing ``}`` of streamed OpenAI arguments JSON - """ - content = payload.strip() - if not content: - return self._func_name, dict(self._args), self._func_closed - - if self._func_name is None: - func_start = content.find(self.func_prefix) - if func_start != -1: - name_start = func_start + len(self.func_prefix) - name_end = content.find('>', name_start) - if name_end != -1: - self._func_name = content[name_start:name_end].strip() - - self._parse_params_incremental(content) - self._func_closed = self.func_suffix in content - return self._func_name, dict(self._args), self._func_closed - - def _complete_open_param_if_ready(self, content: str) -> bool: - """Finalize the in-flight parameter once ```` is available. - - Uses ``_value_start`` so we can locate the closing tag without re-parsing - the ```` header on every non-fast-path chunk. - """ - if self._value_start < 0 or not self._open_param_name: - return False - val_end = content.find(self.param_suffix, self._value_start) - if val_end == -1: - self._in_progress_value = True - return False - param_val_str = content[self._value_start:val_end].strip() - self._args[self._open_param_name] = self._parse_param_value(param_val_str) - self._open_param_name = None - self._value_start = -1 - self._in_progress_value = False - self._scan_pos = val_end + len(self.param_suffix) - return True - - def _parse_params_incremental(self, content: str) -> None: - """Scan ``value`` blocks and update - ``_args``. - - Incomplete parameter headers or values are left open in ``_open_param_name`` - / ``_value_start`` until the closing tag arrives in a later stream chunk. - - ``_scan_pos`` only advances past completed ```` tags; while a - value is streaming, it stays before the open tag. The block below is an - optimization (not required for correctness): skip the while-loop header - re-scan and try to close the current value directly. If ```` - is still missing, return early because the while loop would reach the - same open-tag state anyway. - """ - if self._value_start >= 0: - if not self._complete_open_param_if_ready(content): - return + def _consume_arg_start(self, payload: str, pos: int) -> int | None: + param_start = payload.find('', pos) - while True: - param_start = content.find(self.param_prefix, self._scan_pos) - if param_start == -1: - self._in_progress_value = False - return + if func_end >= 0 and (param_start < 0 or func_end < param_start): + self._payload_closed = True + self._phase = 'done' + return func_end + len('') - name_start = param_start + len(self.param_prefix) - name_end = content.find('>', name_start) - if name_end == -1: - self._in_progress_value = True - return + if param_start < 0: + return None - param_name = content[name_start:name_end].strip() + self._phase = 'arg_name' + return param_start + len(' Any: - try: - parsed_val = json.loads(param_val_str) - return parsed_val if isinstance(parsed_val, str) else param_val_str - except json.JSONDecodeError: - return param_val_str + def _consume_arg_name(self, payload: str, pos: int) -> int | None: + name_end = payload.find('>', pos) + if name_end < 0: + return None + + self._arg_name = payload[pos:name_end].strip() + self._value_parts.clear() + self._reset_value_stream_state() + self._phase = 'arg_value' + return name_end + 1 + + def _consume_arg_value(self, payload: str, pos: int, arg_delta_parts: list[str]) -> tuple[int | None, bool]: + """Consume a parameter value. + + Returns ``(next_pos, should_stop)``. ``should_stop`` is true after + streaming an open value delta, because the next bytes may be the + parameter close tag and must be checked with the next chunk. + """ + value_end = payload.find('', pos) + + if value_end >= 0: + raw = payload[pos:value_end] + if raw: + self._value_parts.append(raw) + if self._arg_name: + self._args[self._arg_name] = ''.join(self._value_parts).strip() + self._arg_name = None + self._value_parts.clear() + self._reset_value_stream_state() + self._phase = 'arg_start' + return value_end + len(''), False + + # Open value: keep any partial "" suffix buffered instead + # of emitting it as argument text. + raw_end = self._trim_partial_close_tag_suffix(payload, pos, '') + if raw_end == pos: + return None, True + + raw_delta = payload[pos:raw_end] + self._value_parts.append(raw_delta) + stream_delta = self._stream_arg_delta(raw_delta) + if stream_delta: + arg_delta_parts.append(stream_delta) + return raw_end, True def parse_tool_call_complete(self, payload: str) -> ToolCall | None: func_name, raw_args_dict, _ = self._extract_params(payload) @@ -180,14 +169,14 @@ def parse_tool_call_complete(self, payload: str) -> ToolCall | None: def _validate_tool_payload(self, payload: str) -> bool: return bool(self._complete_payload_pattern.fullmatch(payload)) - def _extract_params(self, content: str) -> tuple[str | None, dict[str, Any], bool]: + def _extract_params(self, content: str) -> tuple[str | None, dict[str, str], bool]: """Extract function name, parameter map, and close status from XML.""" content = content.strip() func_name = None - func_start = content.find(self.func_prefix) + func_start = content.find('', name_start) if name_end != -1: func_name = content[name_start:name_end].strip() @@ -195,11 +184,11 @@ def _extract_params(self, content: str) -> tuple[str | None, dict[str, Any], boo args_dict = {} search_idx = 0 while True: - param_start = content.find(self.param_prefix, search_idx) + param_start = content.find('', name_start) if name_end == -1: break @@ -207,13 +196,12 @@ def _extract_params(self, content: str) -> tuple[str | None, dict[str, Any], boo param_name = content[name_start:name_end].strip() val_start = name_end + 1 - val_end = content.find(self.param_suffix, val_start) + val_end = content.find('', val_start) if val_end == -1: break - param_val_str = content[val_start:val_end].strip() - args_dict[param_name] = self._parse_param_value(param_val_str) - search_idx = val_end + len(self.param_suffix) + args_dict[param_name] = content[val_start:val_end].strip() + search_idx = val_end + len('') - is_func_closed = self.func_suffix in content + is_func_closed = '' in content return func_name, args_dict, is_func_closed diff --git a/lmdeploy/serve/parsers/tool_parser/tool_parser.py b/lmdeploy/serve/parsers/tool_parser/tool_parser.py index fdf8e1c1da..d631630225 100644 --- a/lmdeploy/serve/parsers/tool_parser/tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/tool_parser.py @@ -1,19 +1,13 @@ # Copyright (c) OpenMMLab. All rights reserved. -# modified from https://github.com/vllm-project/vllm/tree/v0.7.3/vllm/entrypoints/openai/tool_parsers from __future__ import annotations -import json from typing import TYPE_CHECKING -import partial_json_parser import shortuuid from mmengine import Registry -from partial_json_parser.core.options import Allow from lmdeploy.serve.openai.protocol import ( - DeltaFunctionCall, DeltaToolCall, - FunctionCall, ToolCall, ) @@ -33,7 +27,6 @@ def __init__(self): self._active_tool_call_id: str = '' self._active_tool_index: int = -1 self._name_emitted: bool = False - self._args_emitted_len: int = 0 def adjust_request(self, request: ChatCompletionRequest) -> ChatCompletionRequest: """Adjust request payload before rendering, if needed.""" @@ -59,14 +52,12 @@ def start_tool_call(self) -> None: self._active_tool_index += 1 self._active_tool_call_id = f'chatcmpl-tool-{shortuuid.random()}' self._name_emitted = False - self._args_emitted_len = 0 self._tool_payload = '' def finish_tool_call(self) -> None: """Mark end of a tool-call block.""" self._active_tool_call_id = '' self._name_emitted = False - self._args_emitted_len = 0 self._tool_payload = '' def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[DeltaToolCall]: @@ -102,85 +93,5 @@ def validate_complete(self, text: str) -> bool: return True def _validate_tool_payload(self, payload: str) -> bool: - """Return whether one complete JSON tool payload is structurally - valid.""" - try: - obj = json.loads(payload) - except json.JSONDecodeError: - return False - if not isinstance(obj, dict): - return False - name = obj.get('name') - return isinstance(name, str) and bool(name) - - def _decode_tool_incremental_json(self, added_text: str, *, final: bool) -> list[DeltaToolCall]: - self._tool_payload += added_text - payload = self._tool_payload.strip() - if not payload: - return [] - - # After the function name is emitted, arguments are only surfaced at - # final=True. Skip repeated partial_json_parser.loads on growing payload. - if self._name_emitted and not final: - return [] - - flags = Allow.ALL if final else Allow.ALL & ~Allow.STR - try: - obj = partial_json_parser.loads(payload, flags) - except partial_json_parser.core.exceptions.MalformedJSON: - return [] - if not isinstance(obj, dict): - return [] - - out: list[DeltaToolCall] = [] - if not self._name_emitted: - fn_name = obj.get('name') - if isinstance(fn_name, str) and fn_name: - out.append( - DeltaToolCall( - id=self._active_tool_call_id, - index=self._active_tool_index, - type='function', - function=DeltaFunctionCall(name=fn_name), - )) - self._name_emitted = True - - args_obj = obj.get('arguments', obj.get('parameters', None)) - if args_obj is None: - return out - - args_json = json.dumps(args_obj, ensure_ascii=False) - if args_json in ('{}', '[]'): - return out - - # Emit argument text only when the tool payload is complete. This keeps - # streamed argument chunks valid JSON and avoids malformed intermediate - # fragments when partial parsers expose transient dict states. - if final and len(args_json) > self._args_emitted_len: - diff = args_json[self._args_emitted_len:] - out.append( - DeltaToolCall( - id=None, - index=self._active_tool_index, - type=None, - function=DeltaFunctionCall(arguments=diff), - )) - self._args_emitted_len = len(args_json) - return out - - @staticmethod - def _parse_tool_call_complete_json(payload: str) -> ToolCall | None: - if not payload: - return None - try: - obj = json.loads(payload) - except json.JSONDecodeError: - return None - if not isinstance(obj, dict): - return None - name = obj.get('name') - if not isinstance(name, str) or not name: - return None - args_obj = obj.get('arguments', obj.get('parameters', {})) - args_json = json.dumps(args_obj, ensure_ascii=False) - return ToolCall(function=FunctionCall(name=name, arguments=args_json)) + """Return whether one complete tool payload is structurally valid.""" + raise NotImplementedError('ToolParser._validate_tool_payload has not been implemented!') diff --git a/lmdeploy/serve/parsers/tool_parser/xml_tool_parser.py b/lmdeploy/serve/parsers/tool_parser/xml_tool_parser.py index d7a44c06b8..43e3d6acfb 100644 --- a/lmdeploy/serve/parsers/tool_parser/xml_tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/xml_tool_parser.py @@ -2,6 +2,7 @@ from __future__ import annotations import json +from dataclasses import dataclass from typing import TYPE_CHECKING, Any from lmdeploy.serve.openai.protocol import ( @@ -15,6 +16,15 @@ from lmdeploy.serve.openai.protocol import ChatCompletionRequest +@dataclass +class XmlToolSnapshot: + func_name: str | None + completed_args: dict[str, str] + arg_name: str | None + arg_delta: str + payload_closed: bool + + class XmlToolParser(ToolParser): """Base class for XML-like tool parsers. @@ -24,12 +34,15 @@ class XmlToolParser(ToolParser): def __init__(self): super().__init__() self._function_param_schemas: dict[str, dict[str, dict[str, Any]]] = {} - self._xml_has_emitted_json_start = False - self._xml_json_closed = False - self._xml_emitted_param_names: set[str] = set() + self._has_emitted_json_start = False + self._json_closed = False + self._emitted_arg_names: set[str] = set() self._payload_parts: list[str] = [] self._coerced_args: dict[str, Any] = {} - self._in_progress_value = False + self._streamed_arg_name: str | None = None + self._streamed_arg_emitted_len = 0 + self._streamed_arg_quote_opened = False + self._payload_closed = False def adjust_request(self, request: ChatCompletionRequest) -> ChatCompletionRequest: self._function_param_schemas = self._build_function_param_schemas(request) @@ -37,69 +50,110 @@ def adjust_request(self, request: ChatCompletionRequest) -> ChatCompletionReques def start_tool_call(self) -> None: super().start_tool_call() - self._reset_xml_stream_state() + self._reset_stream_state() def finish_tool_call(self) -> None: super().finish_tool_call() - self._reset_xml_stream_state() + self._reset_stream_state() - def _reset_xml_stream_state(self) -> None: - self._xml_has_emitted_json_start = False - self._xml_json_closed = False - self._xml_emitted_param_names.clear() + def _reset_stream_state(self) -> None: + self._has_emitted_json_start = False + self._json_closed = False + self._emitted_arg_names.clear() self._payload_parts.clear() self._coerced_args.clear() - self._in_progress_value = False + self._payload_closed = False + self._reset_arg() self._reset_incremental_state() def _reset_incremental_state(self) -> None: """Reset subclass-specific incremental parse state.""" - def _should_buffer_value_chunk(self, added_text: str, final: bool) -> bool: - """Fast-path plain value fragments that cannot close an XML tag.""" - if final or not self._in_progress_value: - return False - return not any(ch in added_text for ch in '<>/') + def _consume_payload(self, payload: str, *, final: bool) -> tuple[XmlToolSnapshot, int]: + pos = 0 + arg_delta_parts: list[str] = [] + + while pos < len(payload): + if self._phase == 'function': + next_pos = self._consume_function(payload, pos, final) + elif self._phase == 'arg_start': + next_pos = self._consume_arg_start(payload, pos) + elif self._phase == 'arg_name': + next_pos = self._consume_arg_name(payload, pos) + elif self._phase == 'arg_value': + next_pos, should_stop = self._consume_arg_value(payload, pos, arg_delta_parts) + if next_pos is None: + break + pos = next_pos + if should_stop: + break + continue + else: + break + + if next_pos is None: + break + pos = next_pos + + return ( + XmlToolSnapshot( + self._func_name, + dict(self._args), + self._arg_name, + ''.join(arg_delta_parts), + self._payload_closed, + ), + pos, + ) + + def _consume_function(self, payload: str, pos: int, final: bool) -> int | None: + raise NotImplementedError('XmlToolParser._consume_function has not been implemented!') + + def _consume_arg_start(self, payload: str, pos: int) -> int | None: + raise NotImplementedError('XmlToolParser._consume_arg_start has not been implemented!') + + def _consume_arg_name(self, payload: str, pos: int) -> int | None: + raise NotImplementedError('XmlToolParser._consume_arg_name has not been implemented!') + + def _consume_arg_value(self, payload: str, pos: int, arg_delta_parts: list[str]) -> tuple[int | None, bool]: + raise NotImplementedError('XmlToolParser._consume_arg_value has not been implemented!') def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[DeltaToolCall]: self._payload_parts.append(added_text) - if self._should_buffer_value_chunk(added_text, final): - return [] + payload = ''.join(self._payload_parts) + snapshot, consumed = self._consume_payload(payload, final=final) - func_name, raw_args_dict, is_closed = self._extract_incremental_state( - ''.join(self._payload_parts), - final=final, - ) - args_dict = self._get_coerced_args(func_name, raw_args_dict) + if consumed > 0: + left = payload[consumed:] + self._payload_parts.clear() + if left: + self._payload_parts.append(left) out: list[DeltaToolCall] = [] - if func_name and not self._name_emitted: + if snapshot.func_name and not self._name_emitted: out.append( DeltaToolCall( id=self._active_tool_call_id, index=self._active_tool_index, type='function', - function=DeltaFunctionCall(name=func_name), + function=DeltaFunctionCall(name=snapshot.func_name), )) self._name_emitted = True - should_close = is_closed or (final and self._close_json_on_final()) + should_close = snapshot.payload_closed or (final and self._close_json_on_final()) json_fragments: list[str] = [] - if not self._xml_has_emitted_json_start and (args_dict or should_close): - json_fragments.append('{') - self._xml_has_emitted_json_start = True + completed_args = self._get_coerced_args(snapshot.func_name, snapshot.completed_args) + self._append_finished_arg(json_fragments, completed_args) + self._append_completed_args(json_fragments, completed_args) + self._append_open_arg(json_fragments, snapshot) - for key, value in args_dict.items(): - if key in self._xml_emitted_param_names: - continue - prefix = ', ' if len(self._xml_emitted_param_names) > 0 else '' - json_fragments.append(f'{prefix}\"{key}\": {json.dumps(value, ensure_ascii=False)}') - self._xml_emitted_param_names.add(key) - - if should_close and self._xml_has_emitted_json_start and not self._xml_json_closed: + if should_close and not self._has_emitted_json_start: + json_fragments.append('{') + self._has_emitted_json_start = True + if should_close and self._has_emitted_json_start and not self._json_closed: json_fragments.append('}') - self._xml_json_closed = True + self._json_closed = True if json_fragments: out.append( @@ -111,6 +165,85 @@ def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[Delta )) return out + def _append_json_start(self, json_fragments: list[str]) -> None: + if not self._has_emitted_json_start: + json_fragments.append('{') + self._has_emitted_json_start = True + + def _append_finished_arg(self, json_fragments: list[str], completed_args: dict[str, Any]) -> None: + arg_name = self._streamed_arg_name + if arg_name is None or arg_name not in completed_args or arg_name not in self._emitted_arg_names: + return + value = completed_args[arg_name] + if self._streamed_arg_quote_opened: + if isinstance(value, str) and len(value) > self._streamed_arg_emitted_len: + diff = value[self._streamed_arg_emitted_len:] + json_fragments.append(json.dumps(diff, ensure_ascii=False)[1:-1]) + json_fragments.append('"') + else: + value_text = json.dumps(value, ensure_ascii=False) + if len(value_text) > self._streamed_arg_emitted_len: + json_fragments.append(value_text[self._streamed_arg_emitted_len:]) + self._reset_arg() + + def _append_completed_args(self, json_fragments: list[str], completed_args: dict[str, Any]) -> None: + for key, value in completed_args.items(): + if key in self._emitted_arg_names: + continue + self._append_json_start(json_fragments) + prefix = ', ' if len(self._emitted_arg_names) > 0 else '' + json_fragments.append(f'{prefix}"{key}": {json.dumps(value, ensure_ascii=False)}') + self._emitted_arg_names.add(key) + + def _append_open_arg(self, json_fragments: list[str], snapshot: XmlToolSnapshot) -> None: + if snapshot.arg_name is None or not snapshot.arg_delta: + return + + if self._streamed_arg_name == snapshot.arg_name: + json_fragments.append(json.dumps(snapshot.arg_delta, ensure_ascii=False)[1:-1]) + self._streamed_arg_emitted_len += len(snapshot.arg_delta) + return + + if snapshot.arg_name in self._emitted_arg_names: + return + + schema_type = self._get_param_schema_type(snapshot.func_name, snapshot.arg_name) + if schema_type not in (None, 'string'): + return + + self._append_json_start(json_fragments) + prefix = ', ' if len(self._emitted_arg_names) > 0 else '' + json_fragments.append(f'{prefix}"{snapshot.arg_name}": "') + diff = json.dumps(snapshot.arg_delta, ensure_ascii=False)[1:-1] + json_fragments.append(diff) + self._emitted_arg_names.add(snapshot.arg_name) + self._streamed_arg_name = snapshot.arg_name + self._streamed_arg_emitted_len = len(snapshot.arg_delta) + self._streamed_arg_quote_opened = True + + def _reset_arg(self) -> None: + self._streamed_arg_name = None + self._streamed_arg_emitted_len = 0 + self._streamed_arg_quote_opened = False + + def _get_param_schema_type(self, func_name: str | None, param_name: str) -> str | None: + if func_name is None: + return None + param_schema = self._function_param_schemas.get(func_name, {}).get(param_name) + if not isinstance(param_schema, dict): + return None + return self._resolve_schema_type(param_schema) + + @staticmethod + def _trim_partial_close_tag_suffix(payload: str, start: int, close_tag: str) -> int: + """Return safe value end before any partial close-tag suffix.""" + max_len = min(len(payload) - start, len(close_tag) - 1) + for suffix_len in range(max_len, 0, -1): + suffix_start = len(payload) - suffix_len + if close_tag.startswith(payload[suffix_start:]): + return suffix_start + return len(payload) + def _build_function_param_schemas(self, request: ChatCompletionRequest) -> dict[str, dict[str, dict[str, Any]]]: """Build function->parameter schema map from request tools.""" if not request.tools: @@ -207,28 +340,20 @@ def _coerce_value(raw_value: str, schema_type: str | None) -> Any: def _get_coerced_args(self, func_name: str | None, - raw_args_dict: dict[str, Any], + raw_args_dict: dict[str, str], *, use_cache: bool = True) -> dict[str, Any]: if not func_name or not raw_args_dict: return raw_args_dict param_schemas = self._function_param_schemas.get(func_name, {}) - if not param_schemas: - return raw_args_dict coerced = dict(self._coerced_args) if use_cache else {} for key, value in raw_args_dict.items(): if use_cache and key in self._coerced_args: continue - if not isinstance(value, str): - coerced_value = value - else: - schema = param_schemas.get(key) - if not isinstance(schema, dict): - coerced_value = value - else: - schema_type = self._resolve_schema_type(schema) - coerced_value = self._coerce_value(value, schema_type) + schema = param_schemas.get(key) + schema_type = self._resolve_schema_type(schema) if isinstance(schema, dict) else None + coerced_value = self._coerce_value(value, schema_type) if use_cache: self._coerced_args[key] = coerced_value coerced[key] = coerced_value @@ -236,14 +361,3 @@ def _get_coerced_args(self, def _close_json_on_final(self) -> bool: return True - - def _extract_incremental_state(self, - payload: str, - final: bool = False) -> tuple[str | None, dict[str, Any], bool]: - """Parse accumulated inner tool payload and return the current - snapshot. - - Subclasses update their incremental state from ``payload`` and return - ``(func_name, raw_args_dict, is_closed)`` for delta emission. - """ - raise NotImplementedError diff --git a/tests/test_lmdeploy/serve/openai/chat_completions/test_arguments_validation.py b/tests/test_lmdeploy/serve/openai/chat_completions/test_arguments_validation.py index 947ec392aa..bcc33cb0e3 100644 --- a/tests/test_lmdeploy/serve/openai/chat_completions/test_arguments_validation.py +++ b/tests/test_lmdeploy/serve/openai/chat_completions/test_arguments_validation.py @@ -6,27 +6,11 @@ def test_parse_tool_call_complete_json_normalizes_arguments(): - """_parse_tool_call_complete_json validates and re-serializes arguments.""" - from lmdeploy.serve.parsers.tool_parser.tool_parser import ToolParser - - class TestToolParser(ToolParser): - def get_tool_open_tag(self): - return None - - def get_tool_close_tag(self): - return None - - def get_tool_payload_format(self): - return 'json' - - def decode_tool_incremental(self, added_text, *, final): - return [] - - def parse_tool_call_complete(self, payload): - return None + """parse_tool_call_complete validates and re-serializes arguments.""" + from lmdeploy.serve.parsers.tool_parser.json_tool_parser import JsonToolParser payload = '{"name": "get_weather", "arguments": {"city": "NYC" } }' - result = TestToolParser._parse_tool_call_complete_json(payload) + result = JsonToolParser().parse_tool_call_complete(payload) assert result is not None assert result.function.name == 'get_weather' parsed_args = json.loads(result.function.arguments) @@ -34,26 +18,10 @@ def parse_tool_call_complete(self, payload): def test_parse_tool_call_complete_json_invalid_returns_none(): - """_parse_tool_call_complete_json returns None for invalid JSON.""" - from lmdeploy.serve.parsers.tool_parser.tool_parser import ToolParser - - class TestToolParser(ToolParser): - def get_tool_open_tag(self): - return None - - def get_tool_close_tag(self): - return None - - def get_tool_payload_format(self): - return 'json' - - def decode_tool_incremental(self, added_text, *, final): - return [] - - def parse_tool_call_complete(self, payload): - return None + """parse_tool_call_complete returns None for invalid JSON.""" + from lmdeploy.serve.parsers.tool_parser.json_tool_parser import JsonToolParser - result = TestToolParser._parse_tool_call_complete_json( + result = JsonToolParser().parse_tool_call_complete( '{"name": "get_weather", "arguments": {"city":' ) assert result is None diff --git a/tests/test_lmdeploy/serve/openai/chat_completions/test_delta_tool_call_id.py b/tests/test_lmdeploy/serve/openai/chat_completions/test_delta_tool_call_id.py index c13da3bad2..8065fb77f5 100644 --- a/tests/test_lmdeploy/serve/openai/chat_completions/test_delta_tool_call_id.py +++ b/tests/test_lmdeploy/serve/openai/chat_completions/test_delta_tool_call_id.py @@ -1,22 +1,43 @@ # Copyright (c) OpenMMLab. All rights reserved. +import json + +from lmdeploy.serve.parsers.tool_parser.json_tool_parser import JsonToolParser + + +class _TestToolParser(JsonToolParser): + @classmethod + def get_tool_open_tag(cls): + return None + + @classmethod + def get_tool_close_tag(cls): + return None + + +def _stream_argument_fragments(chunks, *, final_on_last): + parser = _TestToolParser() + parser.start_tool_call() + fragments = [] + for idx, chunk in enumerate(chunks): + deltas = parser.decode_tool_incremental(chunk, final=final_on_last and idx == len(chunks) - 1) + fragments.extend(delta.function.arguments for delta in deltas if delta.function and delta.function.arguments) + return fragments + + +def _complete_arguments(payload): + call = _TestToolParser().parse_tool_call_complete(payload) + return json.loads(call.function.arguments) + def test_decode_tool_incremental_json_id_only_on_first_chunk(): """When streaming a tool call, id should appear only on the name-delta - chunk, not on subsequent argument-delta chunks.""" - from lmdeploy.serve.parsers.tool_parser.tool_parser import ToolParser - - class TestToolParser(ToolParser): - def get_tool_open_tag(self): return None - def get_tool_close_tag(self): return None - def get_tool_payload_format(self): return 'json' - def decode_tool_incremental(self, added_text, *, final): return [] - def parse_tool_call_complete(self, payload): return None + chunk, not on subsequent argument chunks.""" - parser = TestToolParser() + parser = _TestToolParser() parser.start_tool_call() # Step 1: feed partial JSON with name - deltas = parser._decode_tool_incremental_json('{"name": "get_weather", ', final=False) + deltas = parser.decode_tool_incremental('{"name": "get_weather", ', final=False) assert len(deltas) == 1 name_delta = deltas[0] assert name_delta.function.name == 'get_weather' @@ -24,18 +45,72 @@ def parse_tool_call_complete(self, payload): return None assert name_delta.id.startswith('chatcmpl-tool-') assert name_delta.type == 'function' - # Step 2: feed final chunk with arguments - deltas = parser._decode_tool_incremental_json('"arguments": {"city": "NYC"}}', final=True) + deltas = parser.decode_tool_incremental('"arguments": {"city": "NY', final=False) assert len(deltas) == 1 args_delta = deltas[0] - assert args_delta.function.arguments is not None + assert args_delta.id is None + assert args_delta.type is None + assert args_delta.function.arguments + + +def test_decode_tool_incremental_json_streams_empty_arguments(): + for arguments in ('{}', '[]', 'null'): + argument_fragments = _stream_argument_fragments( + ['{"name":"f","arguments":' + arguments + '}'], + final_on_last=True, + ) + + assert argument_fragments == [arguments] + + +def test_decode_tool_incremental_json_streams_arguments_before_payload_complete(): + payload = '{"name":"f","arguments":{"city":"New York","units":"c"}}' + fragments = _stream_argument_fragments( + [ + '{"name":"f","arguments":{"city":"Ne', + 'w York","units":"c"}', + ], + final_on_last=False, + ) + + assert fragments + assert json.loads(''.join(fragments)) == _complete_arguments(payload) + + +def test_decode_tool_incremental_json_streams_nested_and_escaped_arguments(): + args = {'outer': {'items': [1, {'text': 'a"b'}], 'path': 'C:\\tmp'}} + payload = '{"name":"f","arguments":' + json.dumps(args) + '}' + body_without_outer_close = payload[:-1] + split_at = body_without_outer_close.find('a\\"b') + fragments = _stream_argument_fragments( + [ + body_without_outer_close[:split_at + 2], + body_without_outer_close[split_at + 2:], + ], + final_on_last=False, + ) + + assert fragments + assert json.loads(''.join(fragments)) == _complete_arguments(payload) + + +def test_decode_tool_incremental_json_streams_parameters_fallback_before_payload_complete(): + payload = '{"name":"f","parameters":{"p":1}}' + fragments = _stream_argument_fragments( + [ + '{"name":"f","parameters":{"p":', + '1}', + ], + final_on_last=False, + ) + + assert fragments + assert json.loads(''.join(fragments)) == _complete_arguments(payload) def test_stream_delta_tool_call_omits_null_id_and_type_in_json(): """Serialized stream chunks should omit null id/type, not emit them as JSON null.""" - import json - from lmdeploy.serve.openai.protocol import ( ChatCompletionResponseStreamChoice, ChatCompletionStreamResponse, diff --git a/tests/test_lmdeploy/serve/openai/chat_completions/test_streaming_metadata.py b/tests/test_lmdeploy/serve/openai/chat_completions/test_streaming_metadata.py new file mode 100644 index 0000000000..c7c738ca1d --- /dev/null +++ b/tests/test_lmdeploy/serve/openai/chat_completions/test_streaming_metadata.py @@ -0,0 +1,247 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from __future__ import annotations + +import asyncio +import json +from types import SimpleNamespace + +import pytest + +from lmdeploy.serve.openai import api_server +from lmdeploy.serve.openai.protocol import ChatCompletionRequest, DeltaMessage + + +class _FakeTokenizer: + + vocab_size = 1000 + + def convert_ids_to_tokens(self, token_id): + return f'tok{token_id}' + + +class _FakeSession: + + def __init__(self, session_id): + self.session_id = session_id + self.epoch = None + + async def async_abort(self): + pass + + +class _FakeSessionManager: + + def __init__(self): + self.sessions = [] + self.removed = [] + + def get(self, session_id=None, create_if_not_exists=True): + session = _FakeSession(session_id if session_id is not None else len(self.sessions) + 1) + self.sessions.append(session) + return session + + def has(self, session_id): + return False + + def map_user_session_id(self, session_id): + return session_id + + def remove(self, session): + self.removed.append(session) + + +class _FakeAsyncEngine: + + model_name = 'fake-model' + epoch = 0 + + def __init__(self, outputs): + self.outputs = outputs + self.backend_config = SimpleNamespace( + adapters=[], + logprobs_mode='raw', + enable_return_routed_experts=False, + ) + self.session_mgr = _FakeSessionManager() + self.tokenizer = SimpleNamespace(model=SimpleNamespace(model=_FakeTokenizer())) + + def generate(self, *args, **kwargs): + + async def _generator(): + for output in self.outputs: + yield SimpleNamespace( + response=output.get('response', ''), + token_ids=output.get('token_ids'), + logprobs=output.get('logprobs'), + input_token_len=1, + generate_token_len=1, + cached_tokens=0, + finish_reason=output.get('finish_reason'), + cache_block_ids=None, + routed_experts=None, + ) + + return _generator() + + +class _FakeRawRequest: + + async def json(self): + return {} + + async def is_disconnected(self): + return False + + +class _StreamingMetadataParser: + + tool_parser_cls = None + + def __init__(self, request): + self.request = request + self.tool_parser = None + + def stream_chunk(self, delta_text, delta_token_ids, **kwargs): + if delta_text in ('', ''): + return [] + if delta_text: + return [(DeltaMessage(content=delta_text), False)] + return [] + + def parse_complete(self, text, token_ids=None, **kwargs): + return text, None, None + + def validate_complete(self, text=None): + return True + + +@pytest.fixture +def install_fake_chat_server(monkeypatch): + + def _install(outputs): + engine = _FakeAsyncEngine(outputs) + monkeypatch.setattr(api_server.VariableInterface, 'async_engine', engine) + monkeypatch.setattr(api_server.VariableInterface, 'response_parser_cls', _StreamingMetadataParser) + return engine + + return _install + + +def _chat_stream_payloads(*, logprobs=True, return_logprob=True, return_token_ids=True): + request = ChatCompletionRequest( + model='fake-model', + messages=[{ + 'role': 'user', + 'content': 'hi', + }], + stream=True, + logprobs=logprobs, + return_logprob=return_logprob, + return_token_ids=return_token_ids, + ) + + async def _collect(): + response = await api_server.chat_completions_v1(request, _FakeRawRequest()) + events = [event async for event in response.body_iterator] + payloads = [] + for event in events: + if isinstance(event, bytes): + event = event.decode() + for line in event.splitlines(): + if not line.startswith('data: '): + continue + data = line.removeprefix('data: ') + if data != '[DONE]': + payloads.append(json.loads(data)) + return payloads + + return asyncio.run(_collect()) + + +def _choice(payload): + return payload['choices'][0] + + +def test_empty_parser_result_with_token_metadata_waits_for_visible_delta(install_fake_chat_server): + install_fake_chat_server([ + { + 'response': '', + 'token_ids': [101], + 'logprobs': [{ + 101: -0.1, + }], + 'finish_reason': None, + }, + { + 'response': 'visible', + 'token_ids': [102], + 'logprobs': [{ + 102: -0.2, + }], + 'finish_reason': None, + }, + ]) + + payloads = _chat_stream_payloads() + + assert len(payloads) == 1 + choice = _choice(payloads[0]) + assert choice['delta']['content'] == 'visible' + assert choice['output_ids'] == [101, 102] + assert choice['output_token_logprobs'] == [[-0.1, 101], [-0.2, 102]] + assert [item['token'] for item in choice['logprobs']['content']] == ['tok101', 'tok102'] + + +def test_suppressed_parser_result_carries_aligned_metadata_to_next_delta(install_fake_chat_server): + install_fake_chat_server([ + { + 'response': '', + 'token_ids': [101, 102], + 'logprobs': [{ + 101: -0.1, + }, { + 102: -0.2, + }], + 'finish_reason': None, + }, + { + 'response': 'visible', + 'token_ids': [103], + 'logprobs': [{ + 103: -0.3, + }], + 'finish_reason': None, + }, + ]) + + payloads = _chat_stream_payloads() + + assert len(payloads) == 1 + choice = _choice(payloads[0]) + assert choice['delta']['content'] == 'visible' + assert choice['output_ids'] == [101, 102, 103] + assert choice['output_token_logprobs'] == [[-0.1, 101], [-0.2, 102], [-0.3, 103]] + assert [item['token'] for item in choice['logprobs']['content']] == ['tok101', 'tok102', 'tok103'] + + +def test_terminal_empty_parser_result_emits_finish_reason_without_empty_content(install_fake_chat_server): + install_fake_chat_server([ + { + 'response': '', + 'token_ids': [101], + 'logprobs': [{ + 101: -0.1, + }], + 'finish_reason': 'stop', + }, + ]) + + payloads = _chat_stream_payloads() + + assert len(payloads) == 1 + choice = _choice(payloads[0]) + assert choice['delta'] == {'role': 'assistant'} + assert choice['finish_reason'] == 'stop' + assert choice['output_ids'] == [101] + assert choice['output_token_logprobs'] == [[-0.1, 101]] + assert [item['token'] for item in choice['logprobs']['content']] == ['tok101'] diff --git a/tests/test_lmdeploy/serve/parsers/test_glm47_parser.py b/tests/test_lmdeploy/serve/parsers/test_glm47_parser.py index e85b716592..8c7e66302c 100644 --- a/tests/test_lmdeploy/serve/parsers/test_glm47_parser.py +++ b/tests/test_lmdeploy/serve/parsers/test_glm47_parser.py @@ -7,8 +7,6 @@ from lmdeploy.serve.parsers.reasoning_parser import ReasoningParserManager from lmdeploy.serve.parsers.tool_parser import Glm47ToolParser, ToolParserManager -from .helpers import first_stream_delta - MODEL_ID = 'zai-org/GLM-4.7' @@ -41,15 +39,62 @@ def response_parser_with_reasoning(): return cls(request=request) +def _flatten_stream_deltas(deltas): + events = [] + for delta_msg, tool_emitted in deltas: + if delta_msg is None: + continue + if delta_msg.reasoning_content is not None: + events.append({'reasoning_content': delta_msg.reasoning_content, 'tool_emitted': tool_emitted}) + if delta_msg.content is not None: + events.append({'content': delta_msg.content, 'tool_emitted': tool_emitted}) + if delta_msg.tool_calls: + for call in delta_msg.tool_calls: + events.append({ + 'tool_emitted': tool_emitted, + 'type': call.type, + 'name': call.function.name if call.function else None, + 'arguments': call.function.arguments if call.function else None, + }) + return events + + +def _stream_tool_arguments(parser, chunks): + parser.start_tool_call() + argument_fragments = [] + for chunk, final in chunks: + for call in parser.decode_tool_incremental(chunk, final=final): + if call.function and call.function.arguments is not None: + argument_fragments.append(call.function.arguments) + parser.finish_tool_call() + return ''.join(argument_fragments) + + +def _stream_tool_arguments_by_chunk(parser, chunks): + parser.start_tool_call() + argument_fragments = [] + per_chunk = [] + for chunk, final in chunks: + chunk_fragments = [] + for call in parser.decode_tool_incremental(chunk, final=final): + if call.function and call.function.arguments is not None: + argument_fragments.append(call.function.arguments) + chunk_fragments.append(call.function.arguments) + per_chunk.append(''.join(chunk_fragments)) + parser.finish_tool_call() + return ''.join(argument_fragments), per_chunk + + REFERENCE_CHUNKS = [ - # (delta_text, emitted_delta_msg, content, tool_emitted, function_name, function_arguments, tool_call_type) - ('prefix ', True, 'prefix ', False, None, None, None), - ('', False, None, False, None, None, None), - # Name is deferred until the first ```` appears in the payload. - ('get_weather', False, None, False, None, None, None), - ('location', True, None, True, 'get_weather', None, 'function'), - ('Beijing', True, None, True, None, '{"location": "Beijing"', None), - ('', True, None, True, None, '}', None), + ('prefix ', [{'content': 'prefix ', 'tool_emitted': False}]), + ('', []), + ('get_weather', []), + ('location', + [{'tool_emitted': True, 'type': 'function', 'name': 'get_weather', 'arguments': None}]), + ('Bei', [{'tool_emitted': True, 'type': None, 'name': None, 'arguments': '{"location": "Bei'}]), + ('jing', [{'tool_emitted': True, 'type': None, 'name': None, 'arguments': 'jing'}]), + ('', [{'tool_emitted': True, 'type': None, 'name': None, 'arguments': '"'}]), + ('', [{'tool_emitted': True, 'type': None, 'name': None, 'arguments': '}'}]), ] @@ -58,24 +103,35 @@ class TestGlm47ResponseParserStreaming: parser.""" def test_stream_chunk_matches_reference(self, response_parser): - for (delta_text, exp_delta_msg, exp_content, exp_tool_emitted, - exp_function_name, exp_function_arguments, exp_type) in REFERENCE_CHUNKS: - delta_msg, tool_emitted = first_stream_delta(response_parser.stream_chunk( - delta_text=delta_text, delta_token_ids=[])) - if not exp_delta_msg: - assert delta_msg is None - continue - assert delta_msg is not None - assert delta_msg.content == exp_content - assert tool_emitted == exp_tool_emitted - if tool_emitted: - assert delta_msg.tool_calls is not None - assert len(delta_msg.tool_calls) == 1 - call = delta_msg.tool_calls[0] - assert call.type == exp_type - assert call.function is not None - assert call.function.name == exp_function_name - assert call.function.arguments == exp_function_arguments + actual = [] + expected = [] + for delta_text, expected_events in REFERENCE_CHUNKS: + actual.extend( + _flatten_stream_deltas(response_parser.stream_chunk(delta_text=delta_text, delta_token_ids=[]))) + expected.extend(expected_events) + assert actual == expected + + def test_stream_chunk_emits_arg_value_before_arg_value_close(self, response_parser): + chunks = [ + '', + 'get_weather', + 'location', + 'San', + ' Francisco', + ', CA', + ] + + argument_fragments = [] + emitted_before_close = False + for chunk in chunks: + for event in _flatten_stream_deltas(response_parser.stream_chunk(delta_text=chunk, delta_token_ids=[])): + fragment = event.get('arguments') + if fragment: + argument_fragments.append(fragment) + emitted_before_close = True + + assert emitted_before_close is True + assert ''.join(argument_fragments) == '{"location": "San Francisco, CA' def test_stream_chunk_function_name_split_before_arg_key(self, response_parser): """Callee name streamed in many deltas before ```` must not @@ -91,14 +147,13 @@ def test_stream_chunk_function_name_split_before_arg_key(self, response_parser): emitted_name = None emitted_args = '' for chunk in chunks: - delta, tool_emitted = first_stream_delta(response_parser.stream_chunk(delta_text=chunk, delta_token_ids=[])) - if not tool_emitted or delta is None or not delta.tool_calls: - continue - for call in delta.tool_calls: - if call.function and call.function.name: - emitted_name = call.function.name - if call.function and call.function.arguments: - emitted_args += call.function.arguments + for event in _flatten_stream_deltas(response_parser.stream_chunk(delta_text=chunk, delta_token_ids=[])): + if not event.get('tool_emitted'): + continue + if event.get('name'): + emitted_name = event['name'] + if event.get('arguments'): + emitted_args += event['arguments'] assert emitted_name == 'get_current_temperature' assert json.loads(emitted_args) == {'location': '北京'} @@ -117,34 +172,28 @@ def test_stream_chunk_mixed_default_reasoning_and_glm47_tool(self, response_pars emitted_args = '' for chunk in chunks: - delta, tool_emitted = first_stream_delta(response_parser_with_reasoning.stream_chunk( - delta_text=chunk, delta_token_ids=[])) - if delta is not None: - if delta.reasoning_content: - reasoning_seen.append(delta.reasoning_content) - if delta.content: - content_seen.append(delta.content) - if tool_emitted and delta and delta.tool_calls: - for call in delta.tool_calls: - if call.function and call.function.name: - emitted_name = call.function.name - if call.function and call.function.arguments: - emitted_args += call.function.arguments + for event in _flatten_stream_deltas( + response_parser_with_reasoning.stream_chunk(delta_text=chunk, delta_token_ids=[])): + if event.get('reasoning_content'): + reasoning_seen.append(event['reasoning_content']) + if event.get('content'): + content_seen.append(event['content']) + if event.get('tool_emitted') and event.get('name'): + emitted_name = event['name'] + if event.get('tool_emitted') and event.get('arguments'): + emitted_args += event['arguments'] for _ in range(3): - delta, tool_emitted = first_stream_delta(response_parser_with_reasoning.stream_chunk( - delta_text='', delta_token_ids=[])) - if delta is not None: - if delta.reasoning_content: - reasoning_seen.append(delta.reasoning_content) - if delta.content: - content_seen.append(delta.content) - if tool_emitted and delta and delta.tool_calls: - for call in delta.tool_calls: - if call.function and call.function.name: - emitted_name = call.function.name - if call.function and call.function.arguments: - emitted_args += call.function.arguments + for event in _flatten_stream_deltas( + response_parser_with_reasoning.stream_chunk(delta_text='', delta_token_ids=[])): + if event.get('reasoning_content'): + reasoning_seen.append(event['reasoning_content']) + if event.get('content'): + content_seen.append(event['content']) + if event.get('tool_emitted') and event.get('name'): + emitted_name = event['name'] + if event.get('tool_emitted') and event.get('arguments'): + emitted_args += event['arguments'] assert ''.join(reasoning_seen) == 'first reason' assert ''.join(content_seen) == '\nAnswer: ' @@ -162,14 +211,13 @@ def test_stream_chunk_keeps_string_without_schema(self, response_parser): emitted_name = None emitted_args = '' for chunk in chunks: - delta, tool_emitted = first_stream_delta(response_parser.stream_chunk(delta_text=chunk, delta_token_ids=[])) - if not tool_emitted or delta is None or not delta.tool_calls: - continue - for call in delta.tool_calls: - if call.function and call.function.name: - emitted_name = call.function.name - if call.function and call.function.arguments: - emitted_args += call.function.arguments + for event in _flatten_stream_deltas(response_parser.stream_chunk(delta_text=chunk, delta_token_ids=[])): + if not event.get('tool_emitted'): + continue + if event.get('name'): + emitted_name = event['name'] + if event.get('arguments'): + emitted_args += event['arguments'] assert emitted_name == 'no_schema_tool' assert emitted_args == '{"zip": "77004", "active": "true"}' @@ -280,3 +328,435 @@ def test_parse_tool_call_complete_keeps_string_without_schema(self): 'active': 'true', 'meta': '{"city":"Houston"}', } + + def test_streamed_arguments_match_complete_parse_for_quoted_string_value(self): + parser = Glm47ToolParser() + request = ChatCompletionRequest( + model=MODEL_ID, + messages=[], + tools=[{ + 'type': 'function', + 'function': { + 'name': 'typed_tool', + 'parameters': { + 'type': 'object', + 'properties': { + 'name': { + 'type': 'string' + }, + }, + }, + }, + }], + tool_choice='auto', + ) + parser.adjust_request(request) + payload = 'typed_toolname"Chen"' + + streamed_arguments = _stream_tool_arguments( + parser, + [ + ('typed_tool', False), + ('name', False), + ('', False), + ('"Chen"', False), + ('', False), + ('', True), + ], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_match_complete_parse_for_string_schema_value_with_whitespace(self): + parser = Glm47ToolParser() + request = ChatCompletionRequest( + model=MODEL_ID, + messages=[], + tools=[{ + 'type': 'function', + 'function': { + 'name': 'typed_tool', + 'parameters': { + 'type': 'object', + 'properties': { + 'name': { + 'type': 'string' + }, + }, + }, + }, + }], + tool_choice='auto', + ) + parser.adjust_request(request) + payload = 'typed_toolname abc ' + + streamed_arguments = _stream_tool_arguments( + parser, + [ + ('typed_tool', False), + ('name a', False), + ('bc ', False), + ('', True), + ], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_emit_typed_arg_after_arg_value_close(self): + parser = Glm47ToolParser() + request = ChatCompletionRequest( + model=MODEL_ID, + messages=[], + tools=[{ + 'type': 'function', + 'function': { + 'name': 'typed_tool', + 'parameters': { + 'type': 'object', + 'properties': { + 'age': { + 'type': 'integer' + }, + }, + }, + }, + }], + tool_choice='auto', + ) + parser.adjust_request(request) + payload = 'typed_toolage12' + + streamed_arguments, per_chunk = _stream_tool_arguments_by_chunk( + parser, + [ + ('typed_tool', False), + ('age', False), + ('12', False), + ('', False), + ('', True), + ], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert per_chunk[1] == '' + assert per_chunk[2] == '' + assert per_chunk[3] == '{"age": 12' + assert per_chunk[4] == '}' + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + assert json.loads(streamed_arguments) == {'age': 12} + + def test_streamed_string_arg_after_typed_arg_uses_own_value(self): + parser = Glm47ToolParser() + request = ChatCompletionRequest( + model=MODEL_ID, + messages=[], + tools=[{ + 'type': 'function', + 'function': { + 'name': 'typed_tool', + 'parameters': { + 'type': 'object', + 'properties': { + 'age': { + 'type': 'integer' + }, + 'name': { + 'type': 'string' + }, + }, + }, + }, + }], + tool_choice='auto', + ) + parser.adjust_request(request) + + streamed_arguments = _stream_tool_arguments( + parser, + [ + ('typed_toolage', False), + ('1', False), + ('2name', False), + ('Alice', False), + ('', True), + ], + ) + + assert json.loads(streamed_arguments) == {'age': 12, 'name': 'Alice'} + + def test_streamed_arguments_match_complete_parse_for_newline_escaped_quoted_string_value(self): + parser = Glm47ToolParser() + request = ChatCompletionRequest( + model=MODEL_ID, + messages=[], + tools=[{ + 'type': 'function', + 'function': { + 'name': 'typed_tool', + 'parameters': { + 'type': 'object', + 'properties': { + 'name': { + 'type': 'string' + }, + }, + }, + }, + }], + tool_choice='auto', + ) + parser.adjust_request(request) + payload = r'typed_toolname"A\nB"' + + streamed_arguments, per_chunk = _stream_tool_arguments_by_chunk( + parser, + [ + ('typed_tool', False), + ('name', False), + ('"A\\', False), + ('nB"', False), + ('', False), + ('', True), + ], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert per_chunk[2] == '' + assert per_chunk[3] == '' + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_match_complete_parse_for_quote_escaped_quoted_string_value(self): + parser = Glm47ToolParser() + request = ChatCompletionRequest( + model=MODEL_ID, + messages=[], + tools=[{ + 'type': 'function', + 'function': { + 'name': 'typed_tool', + 'parameters': { + 'type': 'object', + 'properties': { + 'name': { + 'type': 'string' + }, + }, + }, + }, + }], + tool_choice='auto', + ) + parser.adjust_request(request) + payload = r'typed_toolname"A\"B"' + + streamed_arguments = _stream_tool_arguments( + parser, + [ + ('typed_tool', False), + ('name', False), + ('"A\\', False), + ('"B"', False), + ('', False), + ('', True), + ], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_match_complete_parse_for_invalid_integer_value(self): + parser = Glm47ToolParser() + request = ChatCompletionRequest( + model=MODEL_ID, + messages=[], + tools=[{ + 'type': 'function', + 'function': { + 'name': 'typed_tool', + 'parameters': { + 'type': 'object', + 'properties': { + 'age': { + 'type': 'integer' + }, + }, + }, + }, + }], + tool_choice='auto', + ) + parser.adjust_request(request) + payload = 'typed_toolageabc' + + streamed_arguments = _stream_tool_arguments( + parser, + [ + ('typed_tool', False), + ('age', False), + ('', False), + ('abc', False), + ('', False), + ('', True), + ], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_match_complete_parse_for_invalid_integer_after_numeric_prefix(self): + parser = Glm47ToolParser() + request = ChatCompletionRequest( + model=MODEL_ID, + messages=[], + tools=[{ + 'type': 'function', + 'function': { + 'name': 'typed_tool', + 'parameters': { + 'type': 'object', + 'properties': { + 'age': { + 'type': 'integer' + }, + }, + }, + }, + }], + tool_choice='auto', + ) + parser.adjust_request(request) + payload = 'typed_toolage2a' + + streamed_arguments, per_chunk = _stream_tool_arguments_by_chunk( + parser, + [ + ('typed_tool', False), + ('age', False), + ('2', False), + ('a', False), + ('', False), + ('', True), + ], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert per_chunk[2] == '' + assert per_chunk[3] == '' + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + assert json.loads(streamed_arguments) == {'age': '2a'} + + def test_streamed_arguments_match_complete_parse_when_next_arg_starts_with_previous_close(self): + parser = Glm47ToolParser() + payload = 'two_argsaonebtwo' + + streamed_arguments = _stream_tool_arguments( + parser, + [ + ('two_args', False), + ('a', False), + ('one', False), + ('btwo', False), + ('', False), + ('', True), + ], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_match_complete_parse_when_close_chunk_has_value_tail(self): + parser = Glm47ToolParser() + payload = 'faSan Francisco' + + streamed_arguments = _stream_tool_arguments( + parser, + [ + ('f', False), + ('aSan ', False), + ('Francisco', False), + ('', True), + ], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_match_complete_parse_for_unquoted_newline_value(self): + parser = Glm47ToolParser() + payload = 'faA\nB' + + streamed_arguments = _stream_tool_arguments( + parser, + [ + ('f', False), + ('aA\n', False), + ('B', False), + ('', True), + ], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_match_complete_parse_for_unquoted_quote_value(self): + parser = Glm47ToolParser() + payload = 'faA"B' + + streamed_arguments = _stream_tool_arguments( + parser, + [ + ('f', False), + ('aA"', False), + ('B', False), + ('', True), + ], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_match_complete_parse_when_value_contains_arg_like_text(self): + parser = Glm47ToolParser() + payload = 'fafoo bar baz' + + streamed_arguments = _stream_tool_arguments( + parser, + [ + ('f', False), + ('afoo ', False), + ('bar baz', False), + ('', False), + ('', True), + ], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_decode_incremental_keeps_open_value_buffer_bounded(self): + parser = Glm47ToolParser() + parser.start_tool_call() + try: + parser.decode_tool_incremental('write_filecontent', final=False) + for _ in range(200): + parser.decode_tool_incremental('x' * 32, final=False) + + buffered = ''.join(parser._payload_parts) + assert len(buffered) <= len('') - 1 + finally: + parser.finish_tool_call() diff --git a/tests/test_lmdeploy/serve/parsers/test_llama3_parser.py b/tests/test_lmdeploy/serve/parsers/test_llama3_parser.py index 8436292bb3..0f47b3a207 100644 --- a/tests/test_lmdeploy/serve/parsers/test_llama3_parser.py +++ b/tests/test_lmdeploy/serve/parsers/test_llama3_parser.py @@ -64,6 +64,30 @@ def test_llama3_streaming_without_close_tag(): } +def test_llama3_streaming_emits_arguments_before_json_payload_complete(): + parser = _build_parser() + + parser.stream_chunk('<|python_tag|>', []) + chunks = [ + '{"name":"find_user_id_by_name_zip","parameters":{"first_name":"Ch', + 'en","last_name":"Johnson","zip":77004', + ] + + argument_fragments = [] + for chunk in chunks: + for delta_msg, tool_emitted in parser.stream_chunk(chunk, []): + if not tool_emitted or delta_msg is None or not delta_msg.tool_calls: + continue + for call in delta_msg.tool_calls: + if call.function and call.function.arguments: + argument_fragments.append(call.function.arguments) + + assert argument_fragments + joined = ''.join(argument_fragments) + assert 'Chen' in joined + assert joined != '{"first_name":"Chen","last_name":"Johnson","zip":77004}' + + def test_llama3_parse_complete_without_close_tag(): parser = _build_parser() text = ('<|python_tag|>{"name":"find_user_id_by_name_zip","parameters":{"first_name":"Chen",' diff --git a/tests/test_lmdeploy/serve/parsers/test_qwen3_5_parser.py b/tests/test_lmdeploy/serve/parsers/test_qwen3_5_parser.py index 90c35c68da..f4f48ba455 100644 --- a/tests/test_lmdeploy/serve/parsers/test_qwen3_5_parser.py +++ b/tests/test_lmdeploy/serve/parsers/test_qwen3_5_parser.py @@ -1,13 +1,11 @@ import json -from lmdeploy.serve.openai.protocol import ChatCompletionRequest, DeltaToolCall +from lmdeploy.serve.openai.protocol import ChatCompletionRequest from lmdeploy.serve.parsers import ResponseParserManager from lmdeploy.serve.parsers.reasoning_parser import ReasoningParserManager from lmdeploy.serve.parsers.tool_parser import ToolParserManager from lmdeploy.serve.parsers.tool_parser.qwen3coder_tool_parser import Qwen3CoderToolParser -from .helpers import first_stream_delta - MODEL_ID = 'Qwen/Qwen3.5-35B-A3B' @@ -26,54 +24,95 @@ def _build_response_parser(): return cls(request=request) +def _flatten_stream_deltas(deltas): + events = [] + for delta_msg, tool_emitted in deltas: + if delta_msg is None: + continue + if delta_msg.reasoning_content is not None: + events.append({'reasoning_content': delta_msg.reasoning_content, 'tool_emitted': tool_emitted}) + if delta_msg.content is not None: + events.append({'content': delta_msg.content, 'tool_emitted': tool_emitted}) + if delta_msg.tool_calls: + for call in delta_msg.tool_calls: + events.append({ + 'tool_emitted': tool_emitted, + 'type': call.type, + 'name': call.function.name if call.function else None, + 'arguments': call.function.arguments if call.function else None, + }) + return events + + +def _stream_tool_arguments(parser, chunks): + parser.start_tool_call() + argument_fragments = [] + for chunk in chunks: + for call in parser.decode_tool_incremental(chunk, final=False): + if call.function and call.function.arguments is not None: + argument_fragments.append(call.function.arguments) + parser.finish_tool_call() + return ''.join(argument_fragments) + + +def _stream_tool_arguments_by_chunk(parser, chunks): + parser.start_tool_call() + argument_fragments = [] + per_chunk = [] + for chunk in chunks: + chunk_fragments = [] + for call in parser.decode_tool_incremental(chunk, final=False): + if call.function and call.function.arguments is not None: + argument_fragments.append(call.function.arguments) + chunk_fragments.append(call.function.arguments) + per_chunk.append(''.join(chunk_fragments)) + parser.finish_tool_call() + return ''.join(argument_fragments), per_chunk + + REFERENCE_CHUNKS = [ - # (delta_text, emitted_delta_msg, reasoning_content, content, - # tool_emitted, function_name, function_arguments, tool_call_type) - # Short representative reasoning stream; literal text is irrelevant. - ('计划', True, '计划', None, False, None, None, None), - ('调用', True, '调用', None, False, None, None, None), - ('get', True, 'get', None, False, None, None, None), - ('_current', True, '_current', None, False, None, None, None), - ('_temperature', True, '_temperature', None, False, None, None, None), - ('函数', True, '函数', None, False, None, None, None), - ('并提供', True, '并提供', None, False, None, None, None), - ('location', True, 'location', None, False, None, None, None), - ('参数', True, '参数', None, False, None, None, None), - ('。', True, '。', None, False, None, None, None), - ('\n', True, '\n', None, False, None, None, None), - ('', False, None, None, False, None, None, None), - ('\n\n', True, None, '\n\n', False, None, None, None), - # Tool call section: placeholder; will be updated to match Qwen3.5 XML-style. - ('', False, None, None, False, None, None, None), - ('\n', False, None, None, False, None, None, None), - ('<', False, None, None, False, None, None, None), - ('function', False, None, None, False, None, None, None), - ('=get', False, None, None, False, None, None, None), - ('_current', False, None, None, False, None, None, None), - ('_temperature', False, None, None, False, None, None, None), - ('>', True, None, None, True, 'get_current_temperature', None, 'function'), - ('\n', False, None, None, False, None, None, None), - ('<', False, None, None, False, None, None, None), - ('parameter', False, None, None, False, None, None, None), - ('=location', False, None, None, False, None, None, None), - ('>', False, None, None, False, None, None, None), - ('\n', False, None, None, False, None, None, None), - ('Be', False, None, None, False, None, None, None), - ('ijing', False, None, None, False, None, None, None), - (',', False, None, None, False, None, None, None), - (' China', False, None, None, False, None, None, None), - ('\n', False, None, None, False, None, None, None), - ('` to a single id; Qwen3Coder may emit accumulated JSON args in one delta. - ('>', True, None, None, True, None, '{"location": "Beijing, China"', None), - ('\n', False, None, None, False, None, None, None), - ('', True, None, None, True, None, '}', None), - ('\n', False, None, None, False, None, None, None), - ('', False, None, None, False, None, None, None), - ('', True, None, '', False, None, None, None), + ('计划', [{'reasoning_content': '计划', 'tool_emitted': False}]), + ('调用', [{'reasoning_content': '调用', 'tool_emitted': False}]), + ('get', [{'reasoning_content': 'get', 'tool_emitted': False}]), + ('_current', [{'reasoning_content': '_current', 'tool_emitted': False}]), + ('_temperature', [{'reasoning_content': '_temperature', 'tool_emitted': False}]), + ('函数', [{'reasoning_content': '函数', 'tool_emitted': False}]), + ('并提供', [{'reasoning_content': '并提供', 'tool_emitted': False}]), + ('location', [{'reasoning_content': 'location', 'tool_emitted': False}]), + ('参数', [{'reasoning_content': '参数', 'tool_emitted': False}]), + ('。', [{'reasoning_content': '。', 'tool_emitted': False}]), + ('\n', [{'reasoning_content': '\n', 'tool_emitted': False}]), + ('', []), + ('\n\n', [{'content': '\n\n', 'tool_emitted': False}]), + ('', []), + ('\n', []), + ('<', []), + ('function', []), + ('=get', []), + ('_current', []), + ('_temperature', []), + ('>', [{'tool_emitted': True, 'type': 'function', 'name': 'get_current_temperature', 'arguments': None}]), + ('\n', []), + ('<', []), + ('parameter', []), + ('=location', []), + ('>', []), + ('\n', []), + ('Be', [{'tool_emitted': True, 'type': None, 'name': None, 'arguments': '{"location": "Be'}]), + ('ijing', [{'tool_emitted': True, 'type': None, 'name': None, 'arguments': 'ijing'}]), + (',', [{'tool_emitted': True, 'type': None, 'name': None, 'arguments': ','}]), + (' China', [{'tool_emitted': True, 'type': None, 'name': None, 'arguments': ' China'}]), + ('\n', []), + ('', [{'tool_emitted': True, 'type': None, 'name': None, 'arguments': '"'}]), + ('\n', []), + ('', [{'tool_emitted': True, 'type': None, 'name': None, 'arguments': '}'}]), + ('\n', []), + ('', []), + ('', [{'content': '', 'tool_emitted': False}]), ] @@ -82,50 +121,38 @@ class TestQwen3_5ResponseParserStreaming: parsers.""" def test_stream_chunk_matches_reference(self): - """Feed the real streaming sequence into ResponseParser.stream_chunk - and verify each parsed chunk. - - Expectations for tool_calls will be refined once the Qwen3.5 ground-truth stream is finalized. - """ + response_parser = _build_response_parser() + actual = [] + expected = [] + for delta_text, expected_events in REFERENCE_CHUNKS: + actual.extend( + _flatten_stream_deltas(response_parser.stream_chunk(delta_text=delta_text, delta_token_ids=[]))) + expected.extend(expected_events) + assert actual == expected + def test_stream_chunk_emits_parameter_value_before_parameter_close(self): response_parser = _build_response_parser() - for (delta_text, exp_delta_msg, exp_reasoning, exp_content, exp_tool_emitted, - exp_function_name, exp_function_arguments, - exp_type) in REFERENCE_CHUNKS: - delta_msg, tool_emitted = first_stream_delta(response_parser.stream_chunk( - delta_text=delta_text, - delta_token_ids=[], - )) - if exp_delta_msg is False: - assert delta_msg is None - continue - - assert delta_msg.reasoning_content == exp_reasoning - assert delta_msg.content == exp_content - - # Tool-call expectations in this fixture are placeholders for now. - # Only enforce the exact tool_emitted flag when an explicit tool - # delta shape is provided. - if ( - exp_function_name is None - and exp_function_arguments is None - and exp_type is None - and exp_reasoning is None - and exp_content is None - ): - continue - - assert tool_emitted == exp_tool_emitted - - if tool_emitted: - assert delta_msg.tool_calls is not None - assert len(delta_msg.tool_calls) == 1 - call = delta_msg.tool_calls[0] - assert isinstance(call, DeltaToolCall) - assert call.type == exp_type - assert call.function is not None - assert call.function.name == exp_function_name - assert call.function.arguments == exp_function_arguments + chunks = [ + '', + '', + '', + '', + 'San', + ' Francisco', + ', CA', + ] + + argument_fragments = [] + emitted_before_close = False + for chunk in chunks: + for event in _flatten_stream_deltas(response_parser.stream_chunk(delta_text=chunk, delta_token_ids=[])): + fragment = event.get('arguments') + if fragment: + argument_fragments.append(fragment) + emitted_before_close = True + + assert emitted_before_close is True + assert ''.join(argument_fragments) == '{"location": "San Francisco, CA' def test_parse_complete_parallel_tool_calls_keep_distinct_arguments(self): """Regression: parallel tool calls must not reuse the first call's args.""" @@ -265,3 +292,305 @@ def test_parse_tool_call_complete_coerces_types_by_schema(self): 'scores': [98, 87], 'misc': None, } + + def test_streamed_arguments_match_complete_parse_for_quoted_string_value(self): + parser = Qwen3CoderToolParser() + payload = '"Chen"' + + streamed_arguments = _stream_tool_arguments( + parser, + ['', '', '"Chen"', '', ''], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_match_complete_parse_for_string_schema_value_with_whitespace(self): + parser = Qwen3CoderToolParser() + request = ChatCompletionRequest( + model=MODEL_ID, + messages=[], + tools=[{ + 'type': 'function', + 'function': { + 'name': 'typed_tool', + 'parameters': { + 'type': 'object', + 'properties': { + 'name': { + 'type': 'string' + }, + }, + }, + }, + }], + tool_choice='auto', + ) + parser.adjust_request(request) + payload = ' abc ' + + streamed_arguments = _stream_tool_arguments( + parser, + ['', ' a', 'bc '], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_emit_typed_parameter_after_parameter_close(self): + parser = Qwen3CoderToolParser() + request = ChatCompletionRequest( + model=MODEL_ID, + messages=[], + tools=[{ + 'type': 'function', + 'function': { + 'name': 'typed_tool', + 'parameters': { + 'type': 'object', + 'properties': { + 'age': { + 'type': 'integer' + }, + }, + }, + }, + }], + tool_choice='auto', + ) + parser.adjust_request(request) + payload = '12' + + streamed_arguments, per_chunk = _stream_tool_arguments_by_chunk( + parser, + ['', '', '12', ''], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert per_chunk[1] == '' + assert per_chunk[2] == '' + assert per_chunk[3] == '{"age": 12}' + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + assert json.loads(streamed_arguments) == {'age': 12} + + def test_streamed_string_parameter_after_typed_parameter_uses_own_value(self): + parser = Qwen3CoderToolParser() + request = ChatCompletionRequest( + model=MODEL_ID, + messages=[], + tools=[{ + 'type': 'function', + 'function': { + 'name': 'typed_tool', + 'parameters': { + 'type': 'object', + 'properties': { + 'age': { + 'type': 'integer' + }, + 'name': { + 'type': 'string' + }, + }, + }, + }, + }], + tool_choice='auto', + ) + parser.adjust_request(request) + + streamed_arguments = _stream_tool_arguments( + parser, + [ + '', + '1', + '2', + 'Alice', + ], + ) + + assert json.loads(streamed_arguments) == {'age': 12, 'name': 'Alice'} + + def test_streamed_arguments_match_complete_parse_for_newline_escaped_quoted_string_value(self): + parser = Qwen3CoderToolParser() + payload = r'"A\nB"' + + streamed_arguments, per_chunk = _stream_tool_arguments_by_chunk( + parser, + ['', '', '"A\\', 'nB"', ''], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert per_chunk[2] == '' + assert per_chunk[3] == '' + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_match_complete_parse_for_quote_escaped_quoted_string_value(self): + parser = Qwen3CoderToolParser() + payload = r'"A\"B"' + + streamed_arguments = _stream_tool_arguments( + parser, + ['', '', '"A\\', '"B"', ''], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_match_complete_parse_for_invalid_integer_value(self): + parser = Qwen3CoderToolParser() + request = ChatCompletionRequest( + model=MODEL_ID, + messages=[], + tools=[{ + 'type': 'function', + 'function': { + 'name': 'typed_tool', + 'parameters': { + 'type': 'object', + 'properties': { + 'age': { + 'type': 'integer' + }, + }, + }, + }, + }], + tool_choice='auto', + ) + parser.adjust_request(request) + payload = 'abc' + + streamed_arguments = _stream_tool_arguments( + parser, + ['', '', 'abc', '', ''], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_match_complete_parse_for_invalid_integer_after_numeric_prefix(self): + parser = Qwen3CoderToolParser() + request = ChatCompletionRequest( + model=MODEL_ID, + messages=[], + tools=[{ + 'type': 'function', + 'function': { + 'name': 'typed_tool', + 'parameters': { + 'type': 'object', + 'properties': { + 'age': { + 'type': 'integer' + }, + }, + }, + }, + }], + tool_choice='auto', + ) + parser.adjust_request(request) + payload = '2a' + + streamed_arguments, per_chunk = _stream_tool_arguments_by_chunk( + parser, + ['', '', '2', 'a', ''], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert per_chunk[2] == '' + assert per_chunk[3] == '' + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + assert json.loads(streamed_arguments) == {'age': '2a'} + + def test_streamed_arguments_match_complete_parse_when_next_param_starts_with_previous_close(self): + parser = Qwen3CoderToolParser() + payload = 'onetwo' + + streamed_arguments = _stream_tool_arguments( + parser, + [ + '', + '', + 'one', + 'two', + '', + ], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_match_complete_parse_when_close_chunk_has_value_tail(self): + parser = Qwen3CoderToolParser() + payload = 'San Francisco' + + streamed_arguments = _stream_tool_arguments( + parser, + ['', 'San ', 'Francisco'], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_match_complete_parse_for_unquoted_newline_value(self): + parser = Qwen3CoderToolParser() + payload = 'A\nB' + + streamed_arguments = _stream_tool_arguments( + parser, + ['', 'A\n', 'B'], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_match_complete_parse_for_unquoted_quote_value(self): + parser = Qwen3CoderToolParser() + payload = 'A"B' + + streamed_arguments = _stream_tool_arguments( + parser, + ['', 'A"', 'B'], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_streamed_arguments_match_complete_parse_when_value_contains_parameter_like_text(self): + parser = Qwen3CoderToolParser() + payload = 'foo baz' + + streamed_arguments = _stream_tool_arguments( + parser, + ['', 'foo ', ' baz', ''], + ) + complete_tool_call = parser.parse_tool_call_complete(payload) + + assert complete_tool_call is not None + assert streamed_arguments == complete_tool_call.function.arguments + + def test_decode_incremental_keeps_open_value_buffer_bounded(self): + parser = Qwen3CoderToolParser() + parser.start_tool_call() + try: + parser.decode_tool_incremental('', final=False) + for _ in range(200): + parser.decode_tool_incremental('x' * 32, final=False) + + buffered = ''.join(parser._payload_parts) + assert len(buffered) <= len('') - 1 + finally: + parser.finish_tool_call() diff --git a/tests/test_lmdeploy/serve/parsers/test_qwen3_parser.py b/tests/test_lmdeploy/serve/parsers/test_qwen3_parser.py index 4cea4be68f..bae56c0339 100644 --- a/tests/test_lmdeploy/serve/parsers/test_qwen3_parser.py +++ b/tests/test_lmdeploy/serve/parsers/test_qwen3_parser.py @@ -152,40 +152,60 @@ def test_stream_chunk_matches_reference(self, response_parser, reference_chunks) after streaming completes. """ - expected = [ - (exp_reasoning, exp_content, exp_tool_emitted, exp_function_name, exp_function_arguments, exp_type) - for (_, exp_delta_msg, exp_reasoning, exp_content, exp_tool_emitted, exp_function_name, - exp_function_arguments, exp_type) in reference_chunks - if exp_delta_msg - ] - - actual = [] + expected_reasoning = [row[2] for row in reference_chunks if row[1] and row[2] is not None] + expected_content = [row[3] for row in reference_chunks if row[1] and row[3] is not None] + expected_names = [row[5] for row in reference_chunks if row[1] and row[5] is not None] + expected_types = [row[7] for row in reference_chunks if row[1] and row[7] is not None] + expected_arguments = ''.join(row[6] or '' for row in reference_chunks if row[1]) + + actual_reasoning = [] + actual_content = [] + actual_names = [] + actual_types = [] + actual_arguments = [] for (delta_text, *_) in reference_chunks: if delta_text is None: continue - for delta_msg, tool_emitted in response_parser.stream_chunk( + deltas = response_parser.stream_chunk( delta_text=delta_text, delta_token_ids=[], - ): + ) + for delta_msg, tool_emitted in deltas: if delta_msg is None: continue - actual.append((delta_msg, tool_emitted)) - - assert len(actual) == len(expected) - for (delta_msg, tool_emitted), (exp_reasoning, exp_content, exp_tool_emitted, exp_function_name, - exp_function_arguments, exp_type) in zip(actual, expected): - assert delta_msg.reasoning_content == exp_reasoning - assert delta_msg.content == exp_content - assert tool_emitted == exp_tool_emitted - if tool_emitted: - assert delta_msg.tool_calls is not None - assert len(delta_msg.tool_calls) == 1 - call = delta_msg.tool_calls[0] - assert isinstance(call, DeltaToolCall) - assert call.type == exp_type - assert call.function is not None - assert call.function.name == exp_function_name - assert call.function.arguments == exp_function_arguments + if delta_msg.reasoning_content is not None: + actual_reasoning.append(delta_msg.reasoning_content) + assert tool_emitted is False + if delta_msg.content is not None: + actual_content.append(delta_msg.content) + assert tool_emitted is False + if delta_msg.tool_calls: + assert tool_emitted is True + assert len(delta_msg.tool_calls) == 1 + call = delta_msg.tool_calls[0] + assert isinstance(call, DeltaToolCall) + assert call.function is not None + if call.type is not None: + actual_types.append(call.type) + if call.function.name is not None: + actual_names.append(call.function.name) + if call.function.arguments is not None: + actual_arguments.append(call.function.arguments) + else: + assert tool_emitted is False + + assert actual_reasoning == expected_reasoning + assert actual_content == expected_content + assert actual_names == expected_names + assert actual_types == expected_types + assert ''.join(actual_arguments) == expected_arguments + if actual_arguments: + assert actual_arguments[0] != expected_arguments + assert len(actual_arguments) > 1 + assert all(fragment for fragment in actual_arguments) + complete_text = ''.join(row[0] or '' for row in reference_chunks) + _, complete_tool_calls, _ = response_parser.parse_complete(complete_text) + assert complete_tool_calls[0].function.arguments == expected_arguments def test_stream_chunk_handles_mixed_reasoning_content_tool(self, response_parser): """A single delta may contain reasoning/content/tool segments together. @@ -275,6 +295,45 @@ def test_stream_chunk_tool_enabled_without_reasoning_parser(self): cls.reasoning_parser_cls = old_reasoning_cls cls.tool_parser_cls = old_tool_cls + def test_stream_chunk_returns_empty_list_when_tool_syntax_has_no_visible_delta(self): + cls = ResponseParserManager.get('default') + old_reasoning_cls = cls.reasoning_parser_cls + old_tool_cls = cls.tool_parser_cls + try: + cls.reasoning_parser_cls = None + cls.tool_parser_cls = ToolParserManager.get('qwen3') + request = ChatCompletionRequest( + model=MODEL_ID, + messages=[], + stream=True, + tool_choice='auto', + chat_template_kwargs={'enable_thinking': False}, + ) + parser = cls(request=request) + + assert parser.stream_chunk(delta_text='', delta_token_ids=[1]) == [] + assert parser.stream_chunk(delta_text='\n{"name": "get', delta_token_ids=[2, 3]) == [] + + deltas = parser.stream_chunk(delta_text='_weather"', delta_token_ids=[4]) + delta_msg, tool_emitted = first_stream_delta(deltas) + assert tool_emitted is True + assert delta_msg.tool_calls[0].function.name == 'get_weather' + finally: + cls.reasoning_parser_cls = old_reasoning_cls + cls.tool_parser_cls = old_tool_cls + + def test_decode_incremental_keeps_json_argument_buffer_bounded(self): + parser = ToolParserManager.get('qwen3')() + parser.start_tool_call() + try: + parser.decode_tool_incremental('{"name":"write_file","arguments":{"content":"', final=False) + for _ in range(200): + parser.decode_tool_incremental('x' * 32, final=False) + + assert len(parser._payload) <= 1 + finally: + parser.finish_tool_call() + def test_stream_chunk_reasoning_without_open_tag(self, response_parser): """Qwen thinking mode may omit ```` and start directly with reasoning. From 5d1883503a5307bdb754abb82e7e8ab03fc0ed82 Mon Sep 17 00:00:00 2001 From: lvhan028 Date: Wed, 29 Jul 2026 13:45:58 +0000 Subject: [PATCH 2/2] refactor: stream tool parameters incrementally --- lmdeploy/serve/openai/api_server.py | 62 +++- .../tool_parser/deepseek_v32_tool_parser.py | 188 ++++++++++-- .../parsers/tool_parser/glm47_tool_parser.py | 132 ++------- .../parsers/tool_parser/json_tool_parser.py | 11 +- .../tool_parser/qwen3coder_tool_parser.py | 123 +++----- .../serve/parsers/tool_parser/tool_parser.py | 6 +- .../parsers/tool_parser/xml_tool_parser.py | 273 ++++++++++-------- .../test_streaming_metadata.py | 46 +++ .../parsers/test_tool_parser_incremental.py | 195 +++++++++++++ .../test_deepseek_v32_encoding.py | 3 +- .../test_deepseek_v4_encoding.py | 3 +- 11 files changed, 684 insertions(+), 358 deletions(-) create mode 100644 tests/test_lmdeploy/serve/parsers/test_tool_parser_incremental.py diff --git a/lmdeploy/serve/openai/api_server.py b/lmdeploy/serve/openai/api_server.py index f2849c8c43..662ad85cb0 100644 --- a/lmdeploy/serve/openai/api_server.py +++ b/lmdeploy/serve/openai/api_server.py @@ -294,24 +294,42 @@ def _create_output_token_logprobs(token_ids: list[int] | None = None, class _StreamTokenMetadata: """Token metadata buffered across parser steps with no visible delta.""" - token_ids: list[int] - logprobs: list[dict[int, float]] + token_ids: list[int] | None + logprobs: list[dict[int, float]] | None @classmethod - def from_result(cls, token_ids: list[int] | None, logprobs: list[dict[int, float]] | None): - return cls(list(token_ids or []), list(logprobs or [])) + def from_result( + cls, + token_ids: list[int] | None, + logprobs: list[dict[int, float]] | None, + *, + keep_token_ids: bool, + keep_logprobs: bool, + ) -> _StreamTokenMetadata: + if keep_logprobs and logprobs is not None and len(token_ids or []) != len(logprobs): + raise ValueError('Token ids and logprobs must have the same length.') + return cls( + list(token_ids or []) if keep_token_ids else None, + list(logprobs or []) if keep_logprobs else None, + ) def extend(self, other: _StreamTokenMetadata) -> None: - self.token_ids.extend(other.token_ids) - self.logprobs.extend(other.logprobs) + if self.token_ids is not None: + assert other.token_ids is not None + self.token_ids.extend(other.token_ids) + if self.logprobs is not None: + assert other.logprobs is not None + self.logprobs.extend(other.logprobs) def pop_with(self, current: _StreamTokenMetadata) -> _StreamTokenMetadata: merged = _StreamTokenMetadata( - token_ids=self.token_ids + current.token_ids, - logprobs=self.logprobs + current.logprobs, + token_ids=None if self.token_ids is None else self.token_ids + (current.token_ids or []), + logprobs=None if self.logprobs is None else self.logprobs + (current.logprobs or []), ) - self.token_ids.clear() - self.logprobs.clear() + if self.token_ids is not None: + self.token_ids.clear() + if self.logprobs is not None: + self.logprobs.clear() return merged def output_ids(self, enabled: bool | None) -> list[int] | None: @@ -626,7 +644,12 @@ def create_stream_usage_response_json(usage: UsageInfo) -> str: async def completion_stream_generator() -> AsyncGenerator[str, None]: streaming_tools = False final_usage = None - pending_token_metadata = _StreamTokenMetadata.from_result(None, None) + keep_logprobs = bool(request.logprobs or request.return_logprob) + keep_token_ids = bool(request.return_token_ids or keep_logprobs) + pending_token_metadata = _StreamTokenMetadata( + token_ids=[] if keep_token_ids else None, + logprobs=[] if keep_logprobs else None, + ) async for res in result_generator: if res.finish_reason and include_usage: final_usage = UsageInfo.build( @@ -634,16 +657,25 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]: completion_tokens=res.generate_token_len, cached_tokens=res.cached_tokens, ) - current_token_metadata = _StreamTokenMetadata.from_result(res.token_ids, res.logprobs) + raw_token_ids = res.token_ids or [] + current_token_metadata = _StreamTokenMetadata.from_result( + res.token_ids, + res.logprobs, + keep_token_ids=keep_token_ids, + keep_logprobs=keep_logprobs, + ) stream_deltas = response_parser.stream_chunk( res.response, - current_token_metadata.token_ids + raw_token_ids, ) if not stream_deltas: pending_token_metadata.extend(current_token_metadata) if res.finish_reason is None: continue - current_token_metadata = _StreamTokenMetadata.from_result(None, None) + current_token_metadata = _StreamTokenMetadata( + token_ids=[] if keep_token_ids else None, + logprobs=[] if keep_logprobs else None, + ) stream_deltas = [(DeltaMessage(role='assistant'), False)] should_validate_complete = ( res.finish_reason in ('stop', 'length') @@ -685,7 +717,7 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]: output_ids=stream_output_ids) if res.cache_block_ids is not None and is_last_delta: response_json['cache_block_ids'] = res.cache_block_ids - response_json['remote_token_ids'] = current_token_metadata.token_ids + response_json['remote_token_ids'] = raw_token_ids yield f'data: {json.dumps(response_json)}\n\n' if final_usage is not None: yield f'data: {create_stream_usage_response_json(final_usage)}\n\n' diff --git a/lmdeploy/serve/parsers/tool_parser/deepseek_v32_tool_parser.py b/lmdeploy/serve/parsers/tool_parser/deepseek_v32_tool_parser.py index 62f7f03a1d..b84cb62524 100644 --- a/lmdeploy/serve/parsers/tool_parser/deepseek_v32_tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/deepseek_v32_tool_parser.py @@ -1,6 +1,9 @@ # Copyright (c) OpenMMLab. All rights reserved. from __future__ import annotations +import json +import re + import shortuuid from lmdeploy.deepseek_v32_encoding import dsml_token, parse_tool_calls @@ -24,6 +27,15 @@ class DeepSeekV32ToolParser(ToolParser): tool_calls_block_name = TOOL_CALLS_BLOCK_NAME parse_tool_calls_func = staticmethod(parse_tool_calls) + def __init__(self): + super().__init__() + self._buffer = '' + self._phase = 'invoke_start' + self._invoke_count = 0 + self._current_tool_index = -1 + self._current_param_is_string = False + self._emitted_param_names: set[str] = set() + @classmethod def get_tool_open_tag(cls) -> str | None: return f'\n\n<{cls.dsml_token}{cls.tool_calls_block_name}>' @@ -36,36 +48,162 @@ def get_tool_close_tag(cls) -> str | None: def get_tool_payload_format(cls) -> str: return 'dsml' - def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[DeltaToolCall]: - self._tool_payload += added_text - if not final: - return [] + def start_tool_call(self) -> None: + super().start_tool_call() + self._reset_stream_state() - tool_calls = self.parse_tool_call_complete(self._tool_payload) - if not tool_calls: - return [] + def finish_tool_call(self) -> None: + if self._invoke_count > 0: + self._active_tool_index += self._invoke_count - 1 + super().finish_tool_call() + self._reset_stream_state() + def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[DeltaToolCall]: + """Emit each DSML function name and parameter fragment immediately.""" + self._buffer += added_text out: list[DeltaToolCall] = [] - for offset, tool_call in enumerate(tool_calls): - index = self._active_tool_index + offset - out.append( - DeltaToolCall( - id=f'chatcmpl-tool-{shortuuid.random()}', - index=index, - type='function', - function=DeltaFunctionCall(name=tool_call.function.name), - )) - out.append( - DeltaToolCall( - id=None, - index=index, - type=None, - function=DeltaFunctionCall(arguments=tool_call.function.arguments), - )) - - self._active_tool_index += len(tool_calls) - 1 + pos = 0 + invoke_tag = f'<{self.dsml_token}invoke' + parameter_tag = f'<{self.dsml_token}parameter' + invoke_close_tag = f'' + parameter_close_tag = f'' + + while pos < len(self._buffer): + if self._phase == 'invoke_start': + start = self._buffer.find(invoke_tag, pos) + if start < 0: + pos = self._trim_partial_marker_suffix(self._buffer, pos, (invoke_tag, )) + break + pos = start + len(invoke_tag) + self._phase = 'invoke_header' + continue + + if self._phase == 'invoke_header': + header_end = self._buffer.find('>\n', pos) + if header_end < 0: + break + header = self._buffer[pos:header_end] + match = re.fullmatch(r'\s*name="(.*?)"', header, flags=re.DOTALL) + if match is None: + break + self._current_tool_index = self._active_tool_index + self._invoke_count + tool_id = ( + self._active_tool_call_id + if self._invoke_count == 0 else f'chatcmpl-tool-{shortuuid.random()}' + ) + out.append( + DeltaToolCall( + id=tool_id, + index=self._current_tool_index, + type='function', + function=DeltaFunctionCall(name=match.group(1)), + )) + self._emitted_param_names.clear() + pos = header_end + 2 + self._phase = 'parameter_or_invoke_end' + continue + + if self._phase == 'parameter_or_invoke_end': + parameter_start = self._buffer.find(parameter_tag, pos) + invoke_end = self._buffer.find(invoke_close_tag, pos) + if invoke_end >= 0 and (parameter_start < 0 or invoke_end < parameter_start): + self._append_arguments_delta(out, '}' if self._emitted_param_names else '{}') + self._invoke_count += 1 + self._current_tool_index = -1 + pos = invoke_end + len(invoke_close_tag) + self._phase = 'invoke_start' + continue + if parameter_start < 0: + pos = self._trim_partial_marker_suffix( + self._buffer, + pos, + (parameter_tag, invoke_close_tag), + ) + break + pos = parameter_start + len(parameter_tag) + self._phase = 'parameter_header' + continue + + if self._phase == 'parameter_header': + header_end = self._buffer.find('>', pos) + if header_end < 0: + break + header = self._buffer[pos:header_end] + match = re.fullmatch(r'\s*name="(.*?)"\s+string="(true|false)"', header, flags=re.DOTALL) + if match is None: + break + param_name, string_flag = match.groups() + prefix = '{' if not self._emitted_param_names else ', ' + key = json.dumps(param_name, ensure_ascii=False) + quote = '"' if string_flag == 'true' else '' + self._append_arguments_delta(out, f'{prefix}{key}: {quote}') + self._emitted_param_names.add(param_name) + self._current_param_is_string = string_flag == 'true' + pos = header_end + 1 + self._phase = 'parameter_value' + continue + + if self._phase == 'parameter_value': + value_end = self._buffer.find(parameter_close_tag, pos) + if value_end >= 0: + value_delta = self._encode_param_delta(self._buffer[pos:value_end]) + if self._current_param_is_string: + value_delta += '"' + self._append_arguments_delta(out, value_delta) + pos = value_end + len(parameter_close_tag) + self._phase = 'parameter_or_invoke_end' + continue + + raw_end = self._trim_partial_marker_suffix(self._buffer, pos, (parameter_close_tag, )) + if raw_end == pos: + break + self._append_arguments_delta(out, self._encode_param_delta(self._buffer[pos:raw_end])) + pos = raw_end + break + + break + + if pos > 0: + self._buffer = self._buffer[pos:] return out + def _append_arguments_delta(self, out: list[DeltaToolCall], arguments: str) -> None: + if not arguments: + return + out.append( + DeltaToolCall( + id=None, + index=self._current_tool_index, + type=None, + function=DeltaFunctionCall(arguments=arguments), + )) + + def _encode_param_delta(self, raw: str) -> str: + if not self._current_param_is_string: + return raw + return json.dumps(raw, ensure_ascii=False)[1:-1] + + @staticmethod + def _trim_partial_marker_suffix(payload: str, start: int, markers: tuple[str, ...]) -> int: + """Keep only a suffix that can grow into one of ``markers``.""" + keep_from = len(payload) + for marker in markers: + max_len = min(len(payload) - start, len(marker) - 1) + for suffix_len in range(max_len, 0, -1): + suffix_start = len(payload) - suffix_len + if marker.startswith(payload[suffix_start:]): + keep_from = min(keep_from, suffix_start) + break + return keep_from + + def _reset_stream_state(self) -> None: + self._buffer = '' + self._phase = 'invoke_start' + self._invoke_count = 0 + self._current_tool_index = -1 + self._current_param_is_string = False + self._emitted_param_names.clear() + def parse_tool_call_complete(self, payload: str) -> list[ToolCall] | None: payload = payload.strip() if not payload: diff --git a/lmdeploy/serve/parsers/tool_parser/glm47_tool_parser.py b/lmdeploy/serve/parsers/tool_parser/glm47_tool_parser.py index 3b57658c14..798943678e 100644 --- a/lmdeploy/serve/parsers/tool_parser/glm47_tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/glm47_tool_parser.py @@ -10,7 +10,7 @@ ) from .tool_parser import ToolParserManager -from .xml_tool_parser import XmlToolParser +from .xml_tool_parser import XmlParseResult, XmlToolParser @ToolParserManager.register_module(['glm47']) @@ -26,16 +26,6 @@ class Glm47ToolParser(XmlToolParser): re.DOTALL, ) - def _reset_incremental_state(self) -> None: - self._func_name: str | None = None - self._args: dict[str, str] = {} - self._arg_name: str | None = None - self._value_parts: list[str] = [] - self._phase = 'function' - self._stream_started = False - self._stream_pending_ws = '' - self._stream_blocked = False - @classmethod def get_tool_open_tag(cls) -> str | None: return '' @@ -48,129 +38,67 @@ def get_tool_close_tag(cls) -> str | None: def get_tool_payload_format(cls) -> str: return 'xml' - def _reset_value_stream_state(self) -> None: - self._stream_started = False - self._stream_pending_ws = '' - self._stream_blocked = False - - def _stream_arg_delta(self, raw: str) -> str: - if self._stream_blocked: - return '' - - schema_type = self._get_param_schema_type(self._func_name, self._arg_name or '') - if schema_type not in (None, 'string'): - self._stream_blocked = True - return '' - - text = self._stream_pending_ws + raw - self._stream_pending_ws = '' - - if schema_type == 'string': - if not self._stream_started: - text = text.lstrip() - if not text: - return '' - if text.startswith('"'): - self._stream_blocked = True - return '' - - stable = text.rstrip() - self._stream_pending_ws = text[len(stable):] - if not stable: - return '' - - self._stream_started = True - return stable - - if not self._stream_started: - stripped = text.lstrip() - if not stripped: - self._stream_pending_ws = text - return '' - if stripped.startswith('"'): - self._stream_blocked = True - return '' - - self._stream_started = True - return text - - def _consume_function(self, payload: str, pos: int, final: bool) -> int | None: + def _consume_function(self, payload: str, pos: int, final: bool) -> XmlParseResult: arg_key_start = payload.find('', pos) if arg_key_start >= 0: name = payload[pos:arg_key_start].strip() - if name: - self._func_name = name - self._phase = 'arg_start' - return arg_key_start + return XmlParseResult( + arg_key_start, + next_phase='arg_start', + func_name=name or None, + ) remaining = payload[pos:] if final and remaining.strip(): - self._func_name = remaining.strip() - return len(payload) - return None + return XmlParseResult(len(payload), func_name=remaining.strip()) + return XmlParseResult(None) - def _consume_arg_start(self, payload: str, pos: int) -> int | None: + def _consume_arg_start(self, payload: str, pos: int) -> XmlParseResult: arg_key_start = payload.find('', pos) if arg_key_start < 0: - return None + return XmlParseResult(None) - self._phase = 'arg_name' - return arg_key_start + len('') + return XmlParseResult(arg_key_start + len(''), next_phase='arg_name') - def _consume_arg_name(self, payload: str, pos: int) -> int | None: + def _consume_arg_name(self, payload: str, pos: int) -> XmlParseResult: key_end = payload.find('', pos) if key_end < 0: - return None + return XmlParseResult(None) value_start = payload.find('', key_end + len('')) if value_start < 0: - return None - - self._arg_name = payload[pos:key_end].strip() - self._value_parts.clear() - self._reset_value_stream_state() - self._phase = 'arg_value' - return value_start + len('') + return XmlParseResult(None) - def _consume_arg_value(self, payload: str, pos: int, arg_delta_parts: list[str]) -> tuple[int | None, bool]: - """Consume an argument value. + return XmlParseResult( + value_start + len(''), + next_phase='arg_value', + arg_name=payload[pos:key_end].strip(), + ) - Returns ``(next_pos, should_stop)``. ``should_stop`` is true after - streaming an open value delta, because the next bytes may be the - argument close tag and must be checked with the next chunk. - """ + def _consume_arg_value(self, payload: str, pos: int) -> XmlParseResult: value_end = payload.find('', pos) if value_end >= 0: - raw = payload[pos:value_end] - if raw: - self._value_parts.append(raw) - if self._arg_name: - self._args[self._arg_name] = ''.join(self._value_parts) - self._arg_name = None - self._value_parts.clear() - self._reset_value_stream_state() - self._phase = 'function' - return value_end + len(''), False + return XmlParseResult( + value_end + len(''), + next_phase='function', + arg_delta=payload[pos:value_end], + arg_closed=True, + ) # Open value: keep any partial "" suffix buffered instead # of emitting it as argument text. raw_end = self._trim_partial_close_tag_suffix(payload, pos, '') if raw_end == pos: - return None, True + return XmlParseResult(None) - raw_delta = payload[pos:raw_end] - self._value_parts.append(raw_delta) - stream_delta = self._stream_arg_delta(raw_delta) - if stream_delta: - arg_delta_parts.append(stream_delta) - return raw_end, True + return XmlParseResult(raw_end, arg_delta=payload[pos:raw_end], should_stop=True) def parse_tool_call_complete(self, payload: str) -> ToolCall | None: func_name, raw_args_dict = self._extract_complete_args(payload) if not func_name: return None - args_dict = self._get_coerced_args(func_name, raw_args_dict, use_cache=False) + args_dict = self._get_coerced_args(func_name, raw_args_dict) return ToolCall(function=FunctionCall(name=func_name, arguments=json.dumps(args_dict, ensure_ascii=False))) def _validate_tool_payload(self, payload: str) -> bool: diff --git a/lmdeploy/serve/parsers/tool_parser/json_tool_parser.py b/lmdeploy/serve/parsers/tool_parser/json_tool_parser.py index b7166a3abe..47df946268 100644 --- a/lmdeploy/serve/parsers/tool_parser/json_tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/json_tool_parser.py @@ -21,7 +21,11 @@ class JsonToolSnapshot: class JsonToolParser(ToolParser): - """Base class for JSON tool-call payload parsers.""" + """Base class for JSON tool-call payload parsers. + + The model protocol places ``name`` before exactly one argument field, + named either ``arguments`` or ``parameters``. + """ def __init__(self): super().__init__() @@ -32,7 +36,6 @@ def __init__(self): self._value_depth: int = 0 self._string_open_in_container: bool = False self._value_escaped: bool = False - self._payload_closed: bool = False @classmethod def get_tool_payload_format(cls) -> str: @@ -54,7 +57,7 @@ def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[Delta observed, while string/container state is tracked separately. """ self._payload += added_text - snapshot, consumed = self._consume_payload(self._payload, final=final) + snapshot, consumed = self._consume_payload(self._payload) if consumed > 0: self._payload = self._payload[consumed:] @@ -104,7 +107,7 @@ def _validate_tool_payload(self, payload: str) -> bool: name = obj.get('name') return isinstance(name, str) and bool(name) - def _consume_payload(self, payload: str, *, final: bool) -> tuple[JsonToolSnapshot, int]: + def _consume_payload(self, payload: str) -> tuple[JsonToolSnapshot, int]: pos = 0 args_delta_parts: list[str] = [] func_name: str | None = None diff --git a/lmdeploy/serve/parsers/tool_parser/qwen3coder_tool_parser.py b/lmdeploy/serve/parsers/tool_parser/qwen3coder_tool_parser.py index 10d987a1f5..87ab6746ac 100644 --- a/lmdeploy/serve/parsers/tool_parser/qwen3coder_tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/qwen3coder_tool_parser.py @@ -10,7 +10,7 @@ ) from .tool_parser import ToolParserManager -from .xml_tool_parser import XmlToolParser +from .xml_tool_parser import XmlParseResult, XmlToolParser @ToolParserManager.register_module(['qwen3coder']) @@ -22,16 +22,6 @@ class Qwen3CoderToolParser(XmlToolParser): re.DOTALL, ) - def _reset_incremental_state(self) -> None: - self._func_name: str | None = None - self._args: dict[str, str] = {} - self._arg_name: str | None = None - self._value_parts: list[str] = [] - self._phase = 'function' - self._stream_started = False - self._stream_pending_ws = '' - self._stream_blocked = False - # Qwen3Coder closes tool argument JSON only when the model emits the # explicit function end marker (). We intentionally avoid # auto-closing on stream final to prevent producing a syntactically @@ -51,118 +41,73 @@ def get_tool_close_tag(cls) -> str | None: def get_tool_payload_format(cls) -> str: return 'xml' - def _reset_value_stream_state(self) -> None: - self._stream_started = False - self._stream_pending_ws = '' - self._stream_blocked = False - - def _stream_arg_delta(self, raw: str) -> str: - if self._stream_blocked: - return '' - - schema_type = self._get_param_schema_type(self._func_name, self._arg_name or '') - if schema_type not in (None, 'string'): - self._stream_blocked = True - return '' - - text = self._stream_pending_ws + raw - self._stream_pending_ws = '' - - if not self._stream_started: - text = text.lstrip() - if not text: - return '' - if text.startswith('"'): - self._stream_blocked = True - return '' - - stable = text.rstrip() - self._stream_pending_ws = text[len(stable):] - if not stable: - return '' - - self._stream_started = True - return stable - - def _consume_function(self, payload: str, pos: int, final: bool) -> int | None: + def _consume_function(self, payload: str, pos: int, final: bool) -> XmlParseResult: start = payload.find('', name_start) if name_end < 0: - return None + return XmlParseResult(None) - self._func_name = payload[name_start:name_end].strip() - self._phase = 'arg_start' - return name_end + 1 + return XmlParseResult( + name_end + 1, + next_phase='arg_start', + func_name=payload[name_start:name_end].strip(), + ) - def _consume_arg_start(self, payload: str, pos: int) -> int | None: + def _consume_arg_start(self, payload: str, pos: int) -> XmlParseResult: param_start = payload.find('', pos) if func_end >= 0 and (param_start < 0 or func_end < param_start): - self._payload_closed = True - self._phase = 'done' - return func_end + len('') + return XmlParseResult( + func_end + len(''), + next_phase='done', + payload_closed=True, + ) if param_start < 0: - return None + return XmlParseResult(None) - self._phase = 'arg_name' - return param_start + len(' int | None: + def _consume_arg_name(self, payload: str, pos: int) -> XmlParseResult: name_end = payload.find('>', pos) if name_end < 0: - return None - - self._arg_name = payload[pos:name_end].strip() - self._value_parts.clear() - self._reset_value_stream_state() - self._phase = 'arg_value' - return name_end + 1 + return XmlParseResult(None) - def _consume_arg_value(self, payload: str, pos: int, arg_delta_parts: list[str]) -> tuple[int | None, bool]: - """Consume a parameter value. + return XmlParseResult( + name_end + 1, + next_phase='arg_value', + arg_name=payload[pos:name_end].strip(), + ) - Returns ``(next_pos, should_stop)``. ``should_stop`` is true after - streaming an open value delta, because the next bytes may be the - parameter close tag and must be checked with the next chunk. - """ + def _consume_arg_value(self, payload: str, pos: int) -> XmlParseResult: value_end = payload.find('', pos) if value_end >= 0: - raw = payload[pos:value_end] - if raw: - self._value_parts.append(raw) - if self._arg_name: - self._args[self._arg_name] = ''.join(self._value_parts).strip() - self._arg_name = None - self._value_parts.clear() - self._reset_value_stream_state() - self._phase = 'arg_start' - return value_end + len(''), False + return XmlParseResult( + value_end + len(''), + next_phase='arg_start', + arg_delta=payload[pos:value_end], + arg_closed=True, + ) # Open value: keep any partial "" suffix buffered instead # of emitting it as argument text. raw_end = self._trim_partial_close_tag_suffix(payload, pos, '') if raw_end == pos: - return None, True + return XmlParseResult(None) - raw_delta = payload[pos:raw_end] - self._value_parts.append(raw_delta) - stream_delta = self._stream_arg_delta(raw_delta) - if stream_delta: - arg_delta_parts.append(stream_delta) - return raw_end, True + return XmlParseResult(raw_end, arg_delta=payload[pos:raw_end], should_stop=True) def parse_tool_call_complete(self, payload: str) -> ToolCall | None: func_name, raw_args_dict, _ = self._extract_params(payload) if not func_name: return None - args_dict = self._get_coerced_args(func_name, raw_args_dict, use_cache=False) + args_dict = self._get_coerced_args(func_name, raw_args_dict) args_json = json.dumps(args_dict, ensure_ascii=False) if args_dict else '{}' return ToolCall(function=FunctionCall(name=func_name, arguments=args_json)) diff --git a/lmdeploy/serve/parsers/tool_parser/tool_parser.py b/lmdeploy/serve/parsers/tool_parser/tool_parser.py index d631630225..334e351bd5 100644 --- a/lmdeploy/serve/parsers/tool_parser/tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/tool_parser.py @@ -23,10 +23,10 @@ class ToolParser: """Base class for model-specific tool parsers.""" def __init__(self): - self._tool_payload: str = '' self._active_tool_call_id: str = '' self._active_tool_index: int = -1 self._name_emitted: bool = False + self._payload_closed: bool = False def adjust_request(self, request: ChatCompletionRequest) -> ChatCompletionRequest: """Adjust request payload before rendering, if needed.""" @@ -52,13 +52,13 @@ def start_tool_call(self) -> None: self._active_tool_index += 1 self._active_tool_call_id = f'chatcmpl-tool-{shortuuid.random()}' self._name_emitted = False - self._tool_payload = '' + self._payload_closed = False def finish_tool_call(self) -> None: """Mark end of a tool-call block.""" self._active_tool_call_id = '' self._name_emitted = False - self._tool_payload = '' + self._payload_closed = False def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[DeltaToolCall]: """Decode incremental tool payload emitted between tool tags.""" diff --git a/lmdeploy/serve/parsers/tool_parser/xml_tool_parser.py b/lmdeploy/serve/parsers/tool_parser/xml_tool_parser.py index 43e3d6acfb..24870350d4 100644 --- a/lmdeploy/serve/parsers/tool_parser/xml_tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/xml_tool_parser.py @@ -2,7 +2,7 @@ from __future__ import annotations import json -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any from lmdeploy.serve.openai.protocol import ( @@ -16,19 +16,53 @@ from lmdeploy.serve.openai.protocol import ChatCompletionRequest +@dataclass +class XmlParseState: + """Syntax state shared by XML-like parser implementations.""" + + phase: str = 'function' + func_name: str | None = None + arg_name: str | None = None + + +@dataclass +class XmlParseResult: + """One explicit syntax transition returned by a format adapter.""" + + next_pos: int | None + next_phase: str | None = None + func_name: str | None = None + arg_name: str | None = None + arg_delta: str = '' + arg_closed: bool = False + payload_closed: bool = False + should_stop: bool = False + + +@dataclass +class XmlArgState: + """State needed to decide whether a value can be streamed safely.""" + + mode: str = 'undecided' + pending_ws: str = '' + buffered_parts: list[str] = field(default_factory=list) + + @dataclass class XmlToolSnapshot: func_name: str | None - completed_args: dict[str, str] - arg_name: str | None - arg_delta: str + args_delta: str payload_closed: bool class XmlToolParser(ToolParser): - """Base class for XML-like tool parsers. + """Base class for incremental XML-like tool parsers. - Subclasses only need to implement XML payload extraction. + Format adapters identify syntax boundaries and return ``XmlParseResult``. + This class owns JSON emission, schema coercion, and all stream lifecycle + state. Unquoted string values are emitted immediately and discarded; + only undecided syntax, trailing whitespace, and non-streamable values are + retained. """ def __init__(self): @@ -38,11 +72,8 @@ def __init__(self): self._json_closed = False self._emitted_arg_names: set[str] = set() self._payload_parts: list[str] = [] - self._coerced_args: dict[str, Any] = {} - self._streamed_arg_name: str | None = None - self._streamed_arg_emitted_len = 0 - self._streamed_arg_quote_opened = False - self._payload_closed = False + self._state = XmlParseState() + self._arg_state = XmlArgState() def adjust_request(self, request: ChatCompletionRequest) -> ChatCompletionRequest: self._function_param_schemas = self._build_function_param_schemas(request) @@ -61,61 +92,59 @@ def _reset_stream_state(self) -> None: self._json_closed = False self._emitted_arg_names.clear() self._payload_parts.clear() - self._coerced_args.clear() - self._payload_closed = False - self._reset_arg() - self._reset_incremental_state() - - def _reset_incremental_state(self) -> None: - """Reset subclass-specific incremental parse state.""" + self._state = XmlParseState() + self._arg_state = XmlArgState() def _consume_payload(self, payload: str, *, final: bool) -> tuple[XmlToolSnapshot, int]: pos = 0 - arg_delta_parts: list[str] = [] + json_fragments: list[str] = [] while pos < len(payload): - if self._phase == 'function': - next_pos = self._consume_function(payload, pos, final) - elif self._phase == 'arg_start': - next_pos = self._consume_arg_start(payload, pos) - elif self._phase == 'arg_name': - next_pos = self._consume_arg_name(payload, pos) - elif self._phase == 'arg_value': - next_pos, should_stop = self._consume_arg_value(payload, pos, arg_delta_parts) - if next_pos is None: - break - pos = next_pos - if should_stop: - break - continue + if self._state.phase == 'function': + result = self._consume_function(payload, pos, final) + elif self._state.phase == 'arg_start': + result = self._consume_arg_start(payload, pos) + elif self._state.phase == 'arg_name': + result = self._consume_arg_name(payload, pos) + elif self._state.phase == 'arg_value': + result = self._consume_arg_value(payload, pos) else: break - if next_pos is None: + if result.next_pos is None: + break + + if result.func_name is not None: + self._state.func_name = result.func_name + if result.arg_name is not None: + self._state.arg_name = result.arg_name + self._arg_state = XmlArgState() + if result.arg_delta: + self._consume_arg_delta(result.arg_delta, json_fragments) + if result.arg_closed: + self._finish_arg(json_fragments) + self._state.arg_name = None + if result.payload_closed: + self._payload_closed = True + if result.next_phase is not None: + self._state.phase = result.next_phase + + pos = result.next_pos + if result.should_stop: break - pos = next_pos - - return ( - XmlToolSnapshot( - self._func_name, - dict(self._args), - self._arg_name, - ''.join(arg_delta_parts), - self._payload_closed, - ), - pos, - ) - - def _consume_function(self, payload: str, pos: int, final: bool) -> int | None: + + return XmlToolSnapshot(self._state.func_name, ''.join(json_fragments), self._payload_closed), pos + + def _consume_function(self, payload: str, pos: int, final: bool) -> XmlParseResult: raise NotImplementedError('XmlToolParser._consume_function has not been implemented!') - def _consume_arg_start(self, payload: str, pos: int) -> int | None: + def _consume_arg_start(self, payload: str, pos: int) -> XmlParseResult: raise NotImplementedError('XmlToolParser._consume_arg_start has not been implemented!') - def _consume_arg_name(self, payload: str, pos: int) -> int | None: + def _consume_arg_name(self, payload: str, pos: int) -> XmlParseResult: raise NotImplementedError('XmlToolParser._consume_arg_name has not been implemented!') - def _consume_arg_value(self, payload: str, pos: int, arg_delta_parts: list[str]) -> tuple[int | None, bool]: + def _consume_arg_value(self, payload: str, pos: int) -> XmlParseResult: raise NotImplementedError('XmlToolParser._consume_arg_value has not been implemented!') def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[DeltaToolCall]: @@ -140,14 +169,8 @@ def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[Delta )) self._name_emitted = True + json_fragments = [snapshot.args_delta] if snapshot.args_delta else [] should_close = snapshot.payload_closed or (final and self._close_json_on_final()) - - json_fragments: list[str] = [] - completed_args = self._get_coerced_args(snapshot.func_name, snapshot.completed_args) - self._append_finished_arg(json_fragments, completed_args) - self._append_completed_args(json_fragments, completed_args) - self._append_open_arg(json_fragments, snapshot) - if should_close and not self._has_emitted_json_start: json_fragments.append('{') self._has_emitted_json_start = True @@ -165,66 +188,89 @@ def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[Delta )) return out - def _append_json_start(self, json_fragments: list[str]) -> None: - if not self._has_emitted_json_start: - json_fragments.append('{') - self._has_emitted_json_start = True + def _consume_arg_delta(self, raw: str, json_fragments: list[str]) -> None: + arg_name = self._state.arg_name + if arg_name is None: + return - def _append_finished_arg(self, json_fragments: list[str], completed_args: dict[str, Any]) -> None: - arg_name = self._streamed_arg_name - if arg_name is None or arg_name not in completed_args or arg_name not in self._emitted_arg_names: + arg_state = self._arg_state + if arg_state.mode == 'buffered': + arg_state.buffered_parts.append(raw) return - value = completed_args[arg_name] - if self._streamed_arg_quote_opened: - if isinstance(value, str) and len(value) > self._streamed_arg_emitted_len: - diff = value[self._streamed_arg_emitted_len:] - json_fragments.append(json.dumps(diff, ensure_ascii=False)[1:-1]) - json_fragments.append('"') - else: - value_text = json.dumps(value, ensure_ascii=False) - if len(value_text) > self._streamed_arg_emitted_len: - json_fragments.append(value_text[self._streamed_arg_emitted_len:]) - self._reset_arg() - - def _append_completed_args(self, json_fragments: list[str], completed_args: dict[str, Any]) -> None: - for key, value in completed_args.items(): - if key in self._emitted_arg_names: - continue - self._append_json_start(json_fragments) - prefix = ', ' if len(self._emitted_arg_names) > 0 else '' - json_fragments.append(f'{prefix}"{key}": {json.dumps(value, ensure_ascii=False)}') - self._emitted_arg_names.add(key) - def _append_open_arg(self, json_fragments: list[str], snapshot: XmlToolSnapshot) -> None: - if snapshot.arg_name is None or not snapshot.arg_delta: + if arg_state.mode == 'streaming': + self._stream_string_delta(raw, json_fragments) return - if self._streamed_arg_name == snapshot.arg_name: - json_fragments.append(json.dumps(snapshot.arg_delta, ensure_ascii=False)[1:-1]) - self._streamed_arg_emitted_len += len(snapshot.arg_delta) + schema_type = self._get_param_schema_type(self._state.func_name, arg_name) + if schema_type not in (None, 'string'): + arg_state.mode = 'buffered' + arg_state.buffered_parts.append(raw) return - if snapshot.arg_name in self._emitted_arg_names: + text = arg_state.pending_ws + raw + arg_state.pending_ws = '' + stripped = text.lstrip() + if not stripped: + arg_state.pending_ws = text + return + if stripped.startswith('"'): + arg_state.mode = 'buffered' + arg_state.buffered_parts.append(stripped) return - schema_type = self._get_param_schema_type(snapshot.func_name, snapshot.arg_name) - if schema_type not in (None, 'string'): + arg_state.mode = 'streaming' + self._stream_string_delta(stripped, json_fragments) + + def _stream_string_delta(self, text: str, json_fragments: list[str]) -> None: + arg_state = self._arg_state + text = arg_state.pending_ws + text + stable = text.rstrip() + arg_state.pending_ws = text[len(stable):] + if not stable: return + arg_name = self._state.arg_name + if arg_name is None: + return + if arg_name not in self._emitted_arg_names: + self._append_json_start(json_fragments) + prefix = ', ' if self._emitted_arg_names else '' + json_fragments.append(f'{prefix}{json.dumps(arg_name, ensure_ascii=False)}: "') + self._emitted_arg_names.add(arg_name) + json_fragments.append(json.dumps(stable, ensure_ascii=False)[1:-1]) + + def _finish_arg(self, json_fragments: list[str]) -> None: + arg_name = self._state.arg_name + if arg_name is None: + self._arg_state = XmlArgState() + return + + if self._arg_state.mode == 'streaming': + json_fragments.append('"') + else: + if self._arg_state.mode == 'buffered': + raw_value = ''.join(self._arg_state.buffered_parts) + else: + raw_value = self._arg_state.pending_ws + schema_type = self._get_param_schema_type(self._state.func_name, arg_name) + value = self._coerce_value(raw_value, schema_type) + self._append_completed_arg(json_fragments, arg_name, value) + self._arg_state = XmlArgState() + + def _append_json_start(self, json_fragments: list[str]) -> None: + if not self._has_emitted_json_start: + json_fragments.append('{') + self._has_emitted_json_start = True + + def _append_completed_arg(self, json_fragments: list[str], arg_name: str, value: Any) -> None: + if arg_name in self._emitted_arg_names: + return self._append_json_start(json_fragments) - prefix = ', ' if len(self._emitted_arg_names) > 0 else '' - json_fragments.append(f'{prefix}"{snapshot.arg_name}": "') - diff = json.dumps(snapshot.arg_delta, ensure_ascii=False)[1:-1] - json_fragments.append(diff) - self._emitted_arg_names.add(snapshot.arg_name) - self._streamed_arg_name = snapshot.arg_name - self._streamed_arg_emitted_len = len(snapshot.arg_delta) - self._streamed_arg_quote_opened = True - - def _reset_arg(self) -> None: - self._streamed_arg_name = None - self._streamed_arg_emitted_len = 0 - self._streamed_arg_quote_opened = False + prefix = ', ' if self._emitted_arg_names else '' + key = json.dumps(arg_name, ensure_ascii=False) + json_fragments.append(f'{prefix}{key}: {json.dumps(value, ensure_ascii=False)}') + self._emitted_arg_names.add(arg_name) def _get_param_schema_type(self, func_name: str | None, param_name: str) -> str | None: if func_name is None: @@ -338,25 +384,16 @@ def _coerce_value(raw_value: str, schema_type: str | None) -> Any: return raw_value - def _get_coerced_args(self, - func_name: str | None, - raw_args_dict: dict[str, str], - *, - use_cache: bool = True) -> dict[str, Any]: + def _get_coerced_args(self, func_name: str | None, raw_args_dict: dict[str, str]) -> dict[str, Any]: if not func_name or not raw_args_dict: return raw_args_dict param_schemas = self._function_param_schemas.get(func_name, {}) - coerced = dict(self._coerced_args) if use_cache else {} + coerced: dict[str, Any] = {} for key, value in raw_args_dict.items(): - if use_cache and key in self._coerced_args: - continue schema = param_schemas.get(key) schema_type = self._resolve_schema_type(schema) if isinstance(schema, dict) else None - coerced_value = self._coerce_value(value, schema_type) - if use_cache: - self._coerced_args[key] = coerced_value - coerced[key] = coerced_value + coerced[key] = self._coerce_value(value, schema_type) return coerced def _close_json_on_final(self) -> bool: diff --git a/tests/test_lmdeploy/serve/openai/chat_completions/test_streaming_metadata.py b/tests/test_lmdeploy/serve/openai/chat_completions/test_streaming_metadata.py index c7c738ca1d..9075d443ef 100644 --- a/tests/test_lmdeploy/serve/openai/chat_completions/test_streaming_metadata.py +++ b/tests/test_lmdeploy/serve/openai/chat_completions/test_streaming_metadata.py @@ -96,12 +96,14 @@ async def is_disconnected(self): class _StreamingMetadataParser: tool_parser_cls = None + seen_token_ids = [] def __init__(self, request): self.request = request self.tool_parser = None def stream_chunk(self, delta_text, delta_token_ids, **kwargs): + self.seen_token_ids.append(delta_token_ids) if delta_text in ('', ''): return [] if delta_text: @@ -119,6 +121,7 @@ def validate_complete(self, text=None): def install_fake_chat_server(monkeypatch): def _install(outputs): + _StreamingMetadataParser.seen_token_ids.clear() engine = _FakeAsyncEngine(outputs) monkeypatch.setattr(api_server.VariableInterface, 'async_engine', engine) monkeypatch.setattr(api_server.VariableInterface, 'response_parser_cls', _StreamingMetadataParser) @@ -245,3 +248,46 @@ def test_terminal_empty_parser_result_emits_finish_reason_without_empty_content( assert choice['output_ids'] == [101] assert choice['output_token_logprobs'] == [[-0.1, 101]] assert [item['token'] for item in choice['logprobs']['content']] == ['tok101'] + + +def test_unrequested_metadata_is_not_buffered_but_parser_still_receives_token_ids(install_fake_chat_server): + install_fake_chat_server([ + { + 'response': '', + 'token_ids': [101, 102], + 'logprobs': [{ + 101: -0.1, + }, { + 102: -0.2, + }], + 'finish_reason': None, + }, + { + 'response': 'visible', + 'token_ids': [103], + 'logprobs': [{ + 103: -0.3, + }], + 'finish_reason': None, + }, + ]) + + payloads = _chat_stream_payloads(logprobs=False, return_logprob=False, return_token_ids=False) + + assert _StreamingMetadataParser.seen_token_ids == [[101, 102], [103]] + choice = _choice(payloads[0]) + assert 'output_ids' not in choice + assert 'output_token_logprobs' not in choice + assert 'logprobs' not in choice + + +def test_stream_metadata_rejects_misaligned_logprobs(): + with pytest.raises(ValueError, match='same length'): + api_server._StreamTokenMetadata.from_result( + [101, 102], + [{ + 101: -0.1, + }], + keep_token_ids=True, + keep_logprobs=True, + ) diff --git a/tests/test_lmdeploy/serve/parsers/test_tool_parser_incremental.py b/tests/test_lmdeploy/serve/parsers/test_tool_parser_incremental.py new file mode 100644 index 0000000000..f341f7ed64 --- /dev/null +++ b/tests/test_lmdeploy/serve/parsers/test_tool_parser_incremental.py @@ -0,0 +1,195 @@ +import json +from collections import defaultdict + +import pytest + +from lmdeploy.serve.parsers.tool_parser import ( + DeepSeekV4ToolParser, + DeepSeekV32ToolParser, + Glm47ToolParser, + Qwen3CoderToolParser, +) + + +def _arguments_from(calls): + return ''.join( + call.function.arguments or '' + for call in calls + if call.function is not None + ) + + +def test_qwen_parameter_markers_follow_token_aligned_boundaries(): + parser = Qwen3CoderToolParser() + parser.start_tool_call() + try: + chunks = [ + '<', + 'function', + '=f', + '>', + '<', + 'parameter', + '=p', + '>', + 'value', + '', + '', + ] + per_chunk = [parser.decode_tool_incremental(chunk, final=False) for chunk in chunks] + finally: + parser.finish_tool_call() + + assert per_chunk[3][0].function.name == 'f' + assert all(not calls for calls in per_chunk[4:8]) + assert _arguments_from(per_chunk[8]) == '{"p": "value' + assert all(not calls for calls in per_chunk[9:11]) + assert _arguments_from(per_chunk[11]) == '"' + assert _arguments_from(per_chunk[14]) == '}' + + +@pytest.mark.parametrize( + ('parser_cls', 'prefix', 'tail', 'tail_final', 'close_tag'), + [ + ( + Qwen3CoderToolParser, + '', + '', + False, + '', + ), + ( + Glm47ToolParser, + 'write_filecontent', + '', + True, + '', + ), + ], +) +def test_streamed_megabyte_string_does_not_remain_buffered(parser_cls, prefix, tail, tail_final, close_tag): + parser = parser_cls() + parser.start_tool_call() + fragments = [] + value_chunk = 'x' * 1024 + try: + fragments.append(_arguments_from(parser.decode_tool_incremental(prefix, final=False))) + for _ in range(1024): + fragments.append(_arguments_from(parser.decode_tool_incremental(value_chunk, final=False))) + + assert parser._arg_state.buffered_parts == [] + assert parser._arg_state.pending_ws == '' + assert len(''.join(parser._payload_parts)) <= len(close_tag) - 1 + + fragments.append(_arguments_from(parser.decode_tool_incremental(tail, final=tail_final))) + finally: + parser.finish_tool_call() + + assert json.loads(''.join(fragments)) == {'content': value_chunk * 1024} + + +@pytest.mark.parametrize( + ('parser_cls', 'payloads'), + [ + ( + Qwen3CoderToolParser, + [ + 'one', + 'two', + ], + ), + ( + Glm47ToolParser, + [ + 'firstvalueone', + 'secondvaluetwo', + ], + ), + ], +) +def test_xml_parser_lifecycle_resets_stream_state(parser_cls, payloads): + parser = parser_cls() + parsed = [] + for payload in payloads: + parser.start_tool_call() + calls = parser.decode_tool_incremental(payload, final=True) + parsed.append((calls[0].index, calls[0].function.name, json.loads(_arguments_from(calls)))) + parser.finish_tool_call() + + assert parsed == [ + (0, 'first', {'value': 'one'}), + (1, 'second', {'value': 'two'}), + ] + + +@pytest.mark.parametrize('parser_cls', [DeepSeekV32ToolParser, DeepSeekV4ToolParser]) +def test_dsml_streams_parameter_header_and_string_value_immediately(parser_cls): + parser = parser_cls() + token = parser.dsml_token + parser.start_tool_call() + try: + chunks = [ + f'\n<{token}invoke name="search">\n', + f'<{token}parameter name="query" string="true">', + 'DeepSeek ', + '"streaming"', + f'\n', + f'\n', + ] + per_chunk = [parser.decode_tool_incremental(chunk, final=False) for chunk in chunks] + finally: + parser.finish_tool_call() + + assert per_chunk[0][0].function.name == 'search' + assert _arguments_from(per_chunk[1]) == '{"query": "' + assert _arguments_from(per_chunk[2]) == 'DeepSeek ' + assert _arguments_from(per_chunk[3]) == '\\"streaming\\"' + assert _arguments_from(per_chunk[4]) == '"' + assert _arguments_from(per_chunk[5]) == '}' + assert json.loads(''.join(_arguments_from(calls) for calls in per_chunk)) == { + 'query': 'DeepSeek "streaming"' + } + + +@pytest.mark.parametrize('parser_cls', [DeepSeekV32ToolParser, DeepSeekV4ToolParser]) +def test_dsml_streams_non_string_json_and_multiple_invokes(parser_cls): + parser = parser_cls() + token = parser.dsml_token + payload = ( + f'\n<{token}invoke name="rank">\n' + f'<{token}parameter name="limit" string="false">12\n' + f'<{token}parameter name="filters" string="false">{{"active":true}}\n' + f'\n' + f'<{token}invoke name="lookup">\n' + f'<{token}parameter name="name" string="true">Ada\n' + f'\n' + ) + + parser.start_tool_call() + try: + calls = [] + for char in payload: + calls.extend(parser.decode_tool_incremental(char, final=False)) + calls.extend(parser.decode_tool_incremental('', final=True)) + finally: + parser.finish_tool_call() + + names = [call for call in calls if call.function and call.function.name] + assert [(call.index, call.function.name) for call in names] == [(0, 'rank'), (1, 'lookup')] + assert names[0].id and names[1].id and names[0].id != names[1].id + + arguments_by_index = defaultdict(str) + for call in calls: + if call.function and call.function.arguments is not None: + arguments_by_index[call.index] += call.function.arguments + assert json.loads(arguments_by_index[0]) == { + 'limit': 12, + 'filters': { + 'active': True + }, + } + assert json.loads(arguments_by_index[1]) == {'name': 'Ada'} diff --git a/tests/test_lmdeploy/test_deepseek_v32_encoding.py b/tests/test_lmdeploy/test_deepseek_v32_encoding.py index 64e6723ea7..2c8325e4bb 100644 --- a/tests/test_lmdeploy/test_deepseek_v32_encoding.py +++ b/tests/test_lmdeploy/test_deepseek_v32_encoding.py @@ -251,4 +251,5 @@ def test_deepseek_v32_response_parser_streaming_dsml_function_calls(): assert reasoning == 'need data' assert tool_deltas[0].function.name == 'search' - assert json.loads(tool_deltas[1].function.arguments) == {'query': 'DeepSeek V3.2'} + arguments = ''.join(tool_call.function.arguments or '' for tool_call in tool_deltas) + assert json.loads(arguments) == {'query': 'DeepSeek V3.2'} diff --git a/tests/test_lmdeploy/test_deepseek_v4_encoding.py b/tests/test_lmdeploy/test_deepseek_v4_encoding.py index 7efb9a13b9..a47370e21e 100644 --- a/tests/test_lmdeploy/test_deepseek_v4_encoding.py +++ b/tests/test_lmdeploy/test_deepseek_v4_encoding.py @@ -299,7 +299,8 @@ def test_deepseek_v4_response_parser_streaming_dsml_tool_call(): assert reasoning == 'need a tool' assert tool_deltas[0].function.name == 'search' - assert json.loads(tool_deltas[1].function.arguments) == {'query': 'DeepSeek V4'} + arguments = ''.join(tool_call.function.arguments or '' for tool_call in tool_deltas) + assert json.loads(arguments) == {'query': 'DeepSeek V4'} def test_deepseek_v4_response_parser_reasoning_effort_does_not_enable_thinking():