diff --git a/lmdeploy/serve/parsers/tool_parser/glm47_tool_parser.py b/lmdeploy/serve/parsers/tool_parser/glm47_tool_parser.py index 3b57658c14..5769822d34 100644 --- a/lmdeploy/serve/parsers/tool_parser/glm47_tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/glm47_tool_parser.py @@ -10,7 +10,7 @@ ) from .tool_parser import ToolParserManager -from .xml_tool_parser import XmlToolParser +from .xml_tool_parser import XmlParseResult, XmlToolParser @ToolParserManager.register_module(['glm47']) @@ -27,14 +27,8 @@ class Glm47ToolParser(XmlToolParser): ) def _reset_incremental_state(self) -> None: - self._func_name: str | None = None - self._args: dict[str, str] = {} - self._arg_name: str | None = None self._value_parts: list[str] = [] - self._phase = 'function' - self._stream_started = False - self._stream_pending_ws = '' - self._stream_blocked = False + self._reset_value_stream_state() @classmethod def get_tool_open_tag(cls) -> str | None: @@ -53,13 +47,11 @@ def _reset_value_stream_state(self) -> None: self._stream_pending_ws = '' self._stream_blocked = False - def _stream_arg_delta(self, raw: str) -> str: - if self._stream_blocked: - return '' + def _block_stream_arg_delta(self) -> None: + self._stream_blocked = True - schema_type = self._get_param_schema_type(self._func_name, self._arg_name or '') - if schema_type not in (None, 'string'): - self._stream_blocked = True + def _normalize_stream_arg_delta(self, raw: str, schema_type: str | None) -> str: + if self._stream_blocked: return '' text = self._stream_pending_ws + raw @@ -94,50 +86,67 @@ def _stream_arg_delta(self, raw: str) -> str: self._stream_started = True return text - def _consume_function(self, payload: str, pos: int, final: bool) -> int | None: + def _consume_function(self, payload: str, pos: int, final: bool) -> XmlParseResult: + """Read the GLM function name before the first ````. + + Returns no position when the function name is still split across chunks. A final no-argument payload can + complete with only a function name. + """ arg_key_start = payload.find('', pos) if arg_key_start >= 0: name = payload[pos:arg_key_start].strip() - if name: - self._func_name = name - self._phase = 'arg_start' - return arg_key_start + return XmlParseResult( + next_pos=arg_key_start, + next_phase='arg_start', + func_name=name or None, + ) remaining = payload[pos:] if final and remaining.strip(): - self._func_name = remaining.strip() - return len(payload) - return None + return XmlParseResult(next_pos=len(payload), func_name=remaining.strip()) + return XmlParseResult(next_pos=None) - def _consume_arg_start(self, payload: str, pos: int) -> int | None: + def _consume_arg_start(self, payload: str, pos: int) -> XmlParseResult: + """Find the next GLM ```` marker and enter arg-name + parsing.""" arg_key_start = payload.find('', pos) if arg_key_start < 0: - return None + return XmlParseResult(next_pos=None) - self._phase = 'arg_name' - return arg_key_start + len('') + return XmlParseResult( + next_pos=arg_key_start + len(''), + next_phase='arg_name', + ) - def _consume_arg_name(self, payload: str, pos: int) -> int | None: + def _consume_arg_name(self, payload: str, pos: int) -> XmlParseResult: + """Read ```` content and advance to the following value body. + + The method waits for both ```` and ```` so the base + class only receives an argument name once the value phase can start. + """ key_end = payload.find('', pos) if key_end < 0: - return None + return XmlParseResult(next_pos=None) value_start = payload.find('', key_end + len('')) if value_start < 0: - return None + return XmlParseResult(next_pos=None) - self._arg_name = payload[pos:key_end].strip() self._value_parts.clear() self._reset_value_stream_state() - self._phase = 'arg_value' - return value_start + len('') - - def _consume_arg_value(self, payload: str, pos: int, arg_delta_parts: list[str]) -> tuple[int | None, bool]: - """Consume an argument value. - - Returns ``(next_pos, should_stop)``. ``should_stop`` is true after - streaming an open value delta, because the next bytes may be the - argument close tag and must be checked with the next chunk. + return XmlParseResult( + next_pos=value_start + len(''), + next_phase='arg_value', + arg_name=payload[pos:key_end].strip(), + ) + + def _consume_arg_value(self, payload: str, pos: int) -> XmlParseResult: + """Consume GLM argument value text. + + Closed values return the completed raw value for the base class to attach + to the active argument name. Open values return only safe raw deltas and + leave a possible partial ```` suffix buffered for the next + chunk. """ value_end = payload.find('', pos) @@ -145,26 +154,26 @@ def _consume_arg_value(self, payload: str, pos: int, arg_delta_parts: list[str]) raw = payload[pos:value_end] if raw: self._value_parts.append(raw) - if self._arg_name: - self._args[self._arg_name] = ''.join(self._value_parts) - self._arg_name = None + value = ''.join(self._value_parts) self._value_parts.clear() self._reset_value_stream_state() - self._phase = 'function' - return value_end + len(''), False + return XmlParseResult( + next_pos=value_end + len(''), + next_phase='function', + completed_arg_value=value, + ) - # Open value: keep any partial "" suffix buffered instead - # of emitting it as argument text. raw_end = self._trim_partial_close_tag_suffix(payload, pos, '') if raw_end == pos: - return None, True + return XmlParseResult(next_pos=None, should_stop=True) raw_delta = payload[pos:raw_end] self._value_parts.append(raw_delta) - stream_delta = self._stream_arg_delta(raw_delta) - if stream_delta: - arg_delta_parts.append(stream_delta) - return raw_end, True + return XmlParseResult( + next_pos=raw_end, + raw_arg_delta=raw_delta, + should_stop=True, + ) def parse_tool_call_complete(self, payload: str) -> ToolCall | None: func_name, raw_args_dict = self._extract_complete_args(payload) diff --git a/lmdeploy/serve/parsers/tool_parser/qwen3coder_tool_parser.py b/lmdeploy/serve/parsers/tool_parser/qwen3coder_tool_parser.py index 10d987a1f5..4ae9509a32 100644 --- a/lmdeploy/serve/parsers/tool_parser/qwen3coder_tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/qwen3coder_tool_parser.py @@ -10,7 +10,7 @@ ) from .tool_parser import ToolParserManager -from .xml_tool_parser import XmlToolParser +from .xml_tool_parser import XmlParseResult, XmlToolParser @ToolParserManager.register_module(['qwen3coder']) @@ -23,14 +23,8 @@ class Qwen3CoderToolParser(XmlToolParser): ) def _reset_incremental_state(self) -> None: - self._func_name: str | None = None - self._args: dict[str, str] = {} - self._arg_name: str | None = None self._value_parts: list[str] = [] - self._phase = 'function' - self._stream_started = False - self._stream_pending_ws = '' - self._stream_blocked = False + self._reset_value_stream_state() # Qwen3Coder closes tool argument JSON only when the model emits the # explicit function end marker (). We intentionally avoid @@ -56,13 +50,11 @@ def _reset_value_stream_state(self) -> None: self._stream_pending_ws = '' self._stream_blocked = False - def _stream_arg_delta(self, raw: str) -> str: - if self._stream_blocked: - return '' + def _block_stream_arg_delta(self) -> None: + self._stream_blocked = True - schema_type = self._get_param_schema_type(self._func_name, self._arg_name or '') - if schema_type not in (None, 'string'): - self._stream_blocked = True + def _normalize_stream_arg_delta(self, raw: str, schema_type: str | None) -> str: + if self._stream_blocked: return '' text = self._stream_pending_ws + raw @@ -84,52 +76,72 @@ def _stream_arg_delta(self, raw: str) -> str: self._stream_started = True return stable - def _consume_function(self, payload: str, pos: int, final: bool) -> int | None: + def _consume_function(self, payload: str, pos: int, final: bool) -> XmlParseResult: + """Read a Qwen ```` opener and publish the function + name. + + The parser waits when either the opener or its closing ``>`` is split + across chunks. + """ start = payload.find('', name_start) if name_end < 0: - return None + return XmlParseResult(next_pos=None) - self._func_name = payload[name_start:name_end].strip() - self._phase = 'arg_start' - return name_end + 1 + return XmlParseResult( + next_pos=name_end + 1, + next_phase='arg_start', + func_name=payload[name_start:name_end].strip(), + ) - def _consume_arg_start(self, payload: str, pos: int) -> int | None: + def _consume_arg_start(self, payload: str, pos: int) -> XmlParseResult: + """Find the next Qwen parameter opener or the function close marker. + + ```` closes the XML payload only when it appears before the next + ``', pos) if func_end >= 0 and (param_start < 0 or func_end < param_start): - self._payload_closed = True - self._phase = 'done' - return func_end + len('') + return XmlParseResult( + next_pos=func_end + len(''), + next_phase='done', + payload_closed=True, + ) if param_start < 0: - return None + return XmlParseResult(next_pos=None) - self._phase = 'arg_name' - return param_start + len(' int | None: + def _consume_arg_name(self, payload: str, pos: int) -> XmlParseResult: + """Read the Qwen parameter name and advance to its value body.""" name_end = payload.find('>', pos) if name_end < 0: - return None + return XmlParseResult(next_pos=None) - self._arg_name = payload[pos:name_end].strip() self._value_parts.clear() self._reset_value_stream_state() - self._phase = 'arg_value' - return name_end + 1 - - def _consume_arg_value(self, payload: str, pos: int, arg_delta_parts: list[str]) -> tuple[int | None, bool]: - """Consume a parameter value. - - Returns ``(next_pos, should_stop)``. ``should_stop`` is true after - streaming an open value delta, because the next bytes may be the - parameter close tag and must be checked with the next chunk. + return XmlParseResult( + next_pos=name_end + 1, + next_phase='arg_value', + arg_name=payload[pos:name_end].strip(), + ) + + def _consume_arg_value(self, payload: str, pos: int) -> XmlParseResult: + """Consume Qwen parameter value text. + + Closed values are stripped to preserve Qwen's current XML formatting + behavior. Open values leave a possible partial ```` suffix + buffered so split close tags are not emitted as argument text. """ value_end = payload.find('', pos) @@ -137,26 +149,26 @@ def _consume_arg_value(self, payload: str, pos: int, arg_delta_parts: list[str]) raw = payload[pos:value_end] if raw: self._value_parts.append(raw) - if self._arg_name: - self._args[self._arg_name] = ''.join(self._value_parts).strip() - self._arg_name = None + value = ''.join(self._value_parts).strip() self._value_parts.clear() self._reset_value_stream_state() - self._phase = 'arg_start' - return value_end + len(''), False + return XmlParseResult( + next_pos=value_end + len(''), + next_phase='arg_start', + completed_arg_value=value, + ) - # Open value: keep any partial "" suffix buffered instead - # of emitting it as argument text. raw_end = self._trim_partial_close_tag_suffix(payload, pos, '') if raw_end == pos: - return None, True + return XmlParseResult(next_pos=None, should_stop=True) raw_delta = payload[pos:raw_end] self._value_parts.append(raw_delta) - stream_delta = self._stream_arg_delta(raw_delta) - if stream_delta: - arg_delta_parts.append(stream_delta) - return raw_end, True + return XmlParseResult( + next_pos=raw_end, + raw_arg_delta=raw_delta, + should_stop=True, + ) def parse_tool_call_complete(self, payload: str) -> ToolCall | None: func_name, raw_args_dict, _ = self._extract_params(payload) diff --git a/lmdeploy/serve/parsers/tool_parser/xml_tool_parser.py b/lmdeploy/serve/parsers/tool_parser/xml_tool_parser.py index 43e3d6acfb..a9c58dbabb 100644 --- a/lmdeploy/serve/parsers/tool_parser/xml_tool_parser.py +++ b/lmdeploy/serve/parsers/tool_parser/xml_tool_parser.py @@ -25,6 +25,27 @@ class XmlToolSnapshot: payload_closed: bool +@dataclass +class XmlParseState: + phase: str + func_name: str | None + args: dict[str, str] + arg_name: str | None + payload_closed: bool + + +@dataclass +class XmlParseResult: + next_pos: int | None + next_phase: str | None = None + func_name: str | None = None + arg_name: str | None = None + raw_arg_delta: str = '' + completed_arg_value: str | None = None + payload_closed: bool = False + should_stop: bool = False + + class XmlToolParser(ToolParser): """Base class for XML-like tool parsers. @@ -42,7 +63,7 @@ def __init__(self): self._streamed_arg_name: str | None = None self._streamed_arg_emitted_len = 0 self._streamed_arg_quote_opened = False - self._payload_closed = False + self._xml_state = XmlParseState('function', None, {}, None, False) def adjust_request(self, request: ChatCompletionRequest) -> ChatCompletionRequest: self._function_param_schemas = self._build_function_param_schemas(request) @@ -62,7 +83,7 @@ def _reset_stream_state(self) -> None: self._emitted_arg_names.clear() self._payload_parts.clear() self._coerced_args.clear() - self._payload_closed = False + self._xml_state = XmlParseState('function', None, {}, None, False) self._reset_arg() self._reset_incremental_state() @@ -72,52 +93,85 @@ def _reset_incremental_state(self) -> None: def _consume_payload(self, payload: str, *, final: bool) -> tuple[XmlToolSnapshot, int]: pos = 0 arg_delta_parts: list[str] = [] + state = self._xml_state while pos < len(payload): - if self._phase == 'function': - next_pos = self._consume_function(payload, pos, final) - elif self._phase == 'arg_start': - next_pos = self._consume_arg_start(payload, pos) - elif self._phase == 'arg_name': - next_pos = self._consume_arg_name(payload, pos) - elif self._phase == 'arg_value': - next_pos, should_stop = self._consume_arg_value(payload, pos, arg_delta_parts) - if next_pos is None: - break - pos = next_pos - if should_stop: - break - continue + if state.phase == 'function': + result = self._consume_function(payload, pos, final) + elif state.phase == 'arg_start': + result = self._consume_arg_start(payload, pos) + elif state.phase == 'arg_name': + result = self._consume_arg_name(payload, pos) + elif state.phase == 'arg_value': + result = self._consume_arg_value(payload, pos) else: break - if next_pos is None: + if result.next_pos is None: + break + + if result.func_name is not None: + state.func_name = result.func_name + if result.arg_name is not None: + state.arg_name = result.arg_name + if result.raw_arg_delta: + stream_delta = self._stream_arg_delta(result.raw_arg_delta) + if stream_delta: + arg_delta_parts.append(stream_delta) + if result.completed_arg_value is not None: + if state.arg_name is None: + raise RuntimeError('XML parser completed an argument without an active argument name') + state.args[state.arg_name] = result.completed_arg_value + state.arg_name = None + if result.payload_closed: + state.payload_closed = True + if result.next_phase is not None: + state.phase = result.next_phase + + pos = result.next_pos + if result.should_stop: break - pos = next_pos return ( XmlToolSnapshot( - self._func_name, - dict(self._args), - self._arg_name, + state.func_name, + dict(state.args), + state.arg_name, ''.join(arg_delta_parts), - self._payload_closed, + state.payload_closed, ), pos, ) - def _consume_function(self, payload: str, pos: int, final: bool) -> int | None: + def _consume_function(self, payload: str, pos: int, final: bool) -> XmlParseResult: raise NotImplementedError('XmlToolParser._consume_function has not been implemented!') - def _consume_arg_start(self, payload: str, pos: int) -> int | None: + def _consume_arg_start(self, payload: str, pos: int) -> XmlParseResult: raise NotImplementedError('XmlToolParser._consume_arg_start has not been implemented!') - def _consume_arg_name(self, payload: str, pos: int) -> int | None: + def _consume_arg_name(self, payload: str, pos: int) -> XmlParseResult: raise NotImplementedError('XmlToolParser._consume_arg_name has not been implemented!') - def _consume_arg_value(self, payload: str, pos: int, arg_delta_parts: list[str]) -> tuple[int | None, bool]: + def _consume_arg_value(self, payload: str, pos: int) -> XmlParseResult: raise NotImplementedError('XmlToolParser._consume_arg_value has not been implemented!') + def _stream_arg_delta(self, raw: str) -> str: + state = self._xml_state + if state.arg_name is None: + raise RuntimeError('XML parser streamed an argument value without an active argument name') + + schema_type = self._get_param_schema_type(state.func_name, state.arg_name) + if schema_type not in (None, 'string'): + self._block_stream_arg_delta() + return '' + return self._normalize_stream_arg_delta(raw, schema_type) + + def _block_stream_arg_delta(self) -> None: + raise NotImplementedError('XmlToolParser._block_stream_arg_delta has not been implemented!') + + def _normalize_stream_arg_delta(self, raw: str, schema_type: str | None) -> str: + raise NotImplementedError('XmlToolParser._normalize_stream_arg_delta has not been implemented!') + def decode_tool_incremental(self, added_text: str, *, final: bool) -> list[DeltaToolCall]: self._payload_parts.append(added_text) payload = ''.join(self._payload_parts)