diff --git a/lmdeploy/serve/openai/api_server.py b/lmdeploy/serve/openai/api_server.py
index 9fe802ca1f..662ad85cb0 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,62 @@ 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] | None
+ logprobs: list[dict[int, float]] | None
+
+ @classmethod
+ 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:
+ 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=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 []),
+ )
+ 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:
+ 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 +644,39 @@ def create_stream_usage_response_json(usage: UsageInfo) -> str:
async def completion_stream_generator() -> AsyncGenerator[str, None]:
streaming_tools = False
final_usage = 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:
- 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 []
+ 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,
- delta_token_ids
+ raw_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(
+ 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')
and (request.return_token_ids or request.return_routed_experts)
@@ -627,8 +692,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 +708,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 +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'] = res.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/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/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'{self.dsml_token}invoke>'
+ parameter_close_tag = f'{self.dsml_token}parameter>'
+
+ 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 c91655917d..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 # type: ignore[reportMissingImports]
+from .xml_tool_parser import XmlParseResult, XmlToolParser
@ToolParserManager.register_module(['glm47'])
@@ -21,25 +21,11 @@ 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,
)
- 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
-
@classmethod
def get_tool_open_tag(cls) -> str | None:
return ''
@@ -52,128 +38,82 @@ 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.
-
- 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
-
- 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
-
- return self._func_name, dict(self._args), False
+ 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()
+ return XmlParseResult(
+ arg_key_start,
+ next_phase='arg_start',
+ func_name=name or None,
+ )
+
+ remaining = payload[pos:]
+ if final and remaining.strip():
+ return XmlParseResult(len(payload), func_name=remaining.strip())
+ return XmlParseResult(None)
+
+ def _consume_arg_start(self, payload: str, pos: int) -> XmlParseResult:
+ arg_key_start = payload.find('', pos)
+ if arg_key_start < 0:
+ return XmlParseResult(None)
+
+ return XmlParseResult(arg_key_start + len(''), next_phase='arg_name')
+
+ def _consume_arg_name(self, payload: str, pos: int) -> XmlParseResult:
+ key_end = payload.find('', pos)
+ if key_end < 0:
+ return XmlParseResult(None)
+
+ value_start = payload.find('', key_end + len(''))
+ if value_start < 0:
+ return XmlParseResult(None)
+
+ return XmlParseResult(
+ value_start + len(''),
+ next_phase='arg_value',
+ arg_name=payload[pos:key_end].strip(),
+ )
+
+ def _consume_arg_value(self, payload: str, pos: int) -> XmlParseResult:
+ value_end = payload.find('', pos)
+
+ if value_end >= 0:
+ 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 XmlParseResult(None)
+
+ 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._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)
+ 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:
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 +122,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..47df946268
--- /dev/null
+++ b/lmdeploy/serve/parsers/tool_parser/json_tool_parser.py
@@ -0,0 +1,304 @@
+# 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.
+
+ The model protocol places ``name`` before exactly one argument field,
+ named either ``arguments`` or ``parameters``.
+ """
+
+ 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
+
+ @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)
+ 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) -> 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..87ab6746ac 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,
@@ -11,33 +10,18 @@
)
from .tool_parser import ToolParserManager
-from .xml_tool_parser import XmlToolParser
+from .xml_tool_parser import XmlParseResult, XmlToolParser
@ToolParserManager.register_module(['qwen3coder'])
class Qwen3CoderToolParser(XmlToolParser):
"""Tool parser for Qwen3Coder XML tool-call payloads."""
- func_prefix = '\n]+>\s*(?:\n]+>.*?\s*)*\s*$',
re.DOTALL,
)
- 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
-
# 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
@@ -57,137 +41,87 @@ 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.
-
- ``payload`` is the text inside ``...`` (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`.
-
- Returns:
- ``(func_name, args_dict, is_func_closed)`` where:
-
- - ``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
-
- while True:
- param_start = content.find(self.param_prefix, self._scan_pos)
- if param_start == -1:
- self._in_progress_value = False
- return
-
- name_start = param_start + len(self.param_prefix)
- name_end = content.find('>', name_start)
- if name_end == -1:
- self._in_progress_value = True
- return
-
- param_name = content[name_start:name_end].strip()
-
- val_start = name_end + 1
- val_end = content.find(self.param_suffix, val_start)
- if val_end == -1:
- self._open_param_name = param_name
- self._value_start = val_start
- self._in_progress_value = True
- return
-
- next_pos = val_end + len(self.param_suffix)
- if param_name in self._args:
- self._scan_pos = next_pos
- continue
-
- param_val_str = content[val_start:val_end].strip()
- self._args[param_name] = self._parse_param_value(param_val_str)
- self._scan_pos = next_pos
- self._in_progress_value = False
-
- @staticmethod
- def _parse_param_value(param_val_str: str) -> 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_function(self, payload: str, pos: int, final: bool) -> XmlParseResult:
+ start = payload.find('', name_start)
+ if name_end < 0:
+ return XmlParseResult(None)
+
+ 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) -> XmlParseResult:
+ param_start = payload.find('', pos)
+
+ if func_end >= 0 and (param_start < 0 or func_end < param_start):
+ return XmlParseResult(
+ func_end + len(''),
+ next_phase='done',
+ payload_closed=True,
+ )
+
+ if param_start < 0:
+ return XmlParseResult(None)
+
+ return XmlParseResult(param_start + len(' XmlParseResult:
+ name_end = payload.find('>', pos)
+ if name_end < 0:
+ return XmlParseResult(None)
+
+ return XmlParseResult(
+ name_end + 1,
+ next_phase='arg_value',
+ arg_name=payload[pos:name_end].strip(),
+ )
+
+ def _consume_arg_value(self, payload: str, pos: int) -> XmlParseResult:
+ value_end = payload.find('', pos)
+
+ if value_end >= 0:
+ 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 XmlParseResult(None)
+
+ 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))
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 +129,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 +141,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..334e351bd5 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,
)
@@ -29,11 +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._args_emitted_len: int = 0
+ self._payload_closed: bool = False
def adjust_request(self, request: ChatCompletionRequest) -> ChatCompletionRequest:
"""Adjust request payload before rendering, if needed."""
@@ -59,15 +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._args_emitted_len = 0
- 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._args_emitted_len = 0
- 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."""
@@ -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..24870350d4 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, field
from typing import TYPE_CHECKING, Any
from lmdeploy.serve.openai.protocol import (
@@ -15,21 +16,64 @@
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
+ 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):
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._state = XmlParseState()
+ self._arg_state = XmlArgState()
def adjust_request(self, request: ChatCompletionRequest) -> ChatCompletionRequest:
self._function_param_schemas = self._build_function_param_schemas(request)
@@ -37,69 +81,102 @@ 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._reset_incremental_state()
+ self._state = XmlParseState()
+ self._arg_state = XmlArgState()
+
+ def _consume_payload(self, payload: str, *, final: bool) -> tuple[XmlToolSnapshot, int]:
+ pos = 0
+ json_fragments: list[str] = []
+
+ while pos < len(payload):
+ 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 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
- def _reset_incremental_state(self) -> None:
- """Reset subclass-specific incremental parse state."""
+ return XmlToolSnapshot(self._state.func_name, ''.join(json_fragments), self._payload_closed), pos
- 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_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) -> XmlParseResult:
+ raise NotImplementedError('XmlToolParser._consume_arg_start has not been implemented!')
+
+ 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) -> XmlParseResult:
+ 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())
-
- json_fragments: list[str] = []
- if not self._xml_has_emitted_json_start and (args_dict or should_close):
+ json_fragments = [snapshot.args_delta] if snapshot.args_delta else []
+ should_close = snapshot.payload_closed or (final and self._close_json_on_final())
+ if should_close and not self._has_emitted_json_start:
json_fragments.append('{')
- self._xml_has_emitted_json_start = True
-
- 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:
+ 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 +188,108 @@ def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[Delta
))
return out
+ def _consume_arg_delta(self, raw: str, json_fragments: list[str]) -> None:
+ arg_name = self._state.arg_name
+ if arg_name is None:
+ return
+
+ arg_state = self._arg_state
+ if arg_state.mode == 'buffered':
+ arg_state.buffered_parts.append(raw)
+ return
+
+ if arg_state.mode == 'streaming':
+ self._stream_string_delta(raw, json_fragments)
+ return
+
+ 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
+
+ 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
+
+ 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 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:
+ 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:
@@ -205,45 +384,17 @@ 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, Any],
- *,
- 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, {})
- if not param_schemas:
- return raw_args_dict
- 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
- 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)
- if use_cache:
- self._coerced_args[key] = coerced_value
- coerced[key] = coerced_value
+ schema = param_schemas.get(key)
+ schema_type = self._resolve_schema_type(schema) if isinstance(schema, dict) else None
+ coerced[key] = self._coerce_value(value, schema_type)
return coerced
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..9075d443ef
--- /dev/null
+++ b/tests/test_lmdeploy/serve/openai/chat_completions/test_streaming_metadata.py
@@ -0,0 +1,293 @@
+# 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
+ 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:
+ 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):
+ _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)
+ 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']
+
+
+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_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),
- ('', False, None, None, False, None, None, None),
- ('parameter', False, None, None, False, None, None, None),
- # Tokenizer maps this `>` 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),
- ('', False, None, None, False, None, None, None),
- ('function', 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', []),
+ ('', []),
+ ('parameter', []),
+ ('>', [{'tool_emitted': True, 'type': None, 'name': None, 'arguments': '"'}]),
+ ('\n', []),
+ ('', []),
+ ('function', []),
+ ('>', [{'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.
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',
+ '',
+ 'parameter',
+ '>',
+ '',
+ 'function',
+ '>',
+ ]
+ 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'{token}parameter>\n',
+ f'{token}invoke>\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{token}parameter>\n'
+ f'<{token}parameter name="filters" string="false">{{"active":true}}{token}parameter>\n'
+ f'{token}invoke>\n'
+ f'<{token}invoke name="lookup">\n'
+ f'<{token}parameter name="name" string="true">Ada{token}parameter>\n'
+ f'{token}invoke>\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():