From 2d33b9279371134b530f150c3bdff792dfde6a9b Mon Sep 17 00:00:00 2001 From: guojn1 Date: Wed, 23 Sep 2026 15:00:39 +0800 Subject: [PATCH 1/3] fix(sglang): stream tool_call deltas incrementally instead of buffering until finish Client-side TTFT of pure tool-call streaming requests was effectively equal to the whole generation time: the post-processor only attached accumulated tool_calls to the finish chunk, and content-less deltas were suppressed, so no bytes reached the client during generation (Nexus TTFT/total median ~99.6%/97.4% vs ~17.2% for native SGLang). The post-processor now emits OpenAI tool_call deltas as the parser produces them: the id+name delta goes out with the first argument fragment and argument fragments stream as they arrive (the un-emitted tail of the accumulated arguments is the source of truth, so event granularity cannot produce duplicates). Correctness guards from the buffered design are retained: - names without any argument fragment are never emitted, so a misidentified prompt word cannot leak as a dangling call; - unknown tool names (not in the request's tool list) are suppressed mid-stream and purged at finish, exactly as before; - the finish-time full-text re-parse remains authoritative, now in merge mode: argument suffixes for partially streamed calls are patched (with streamed indices kept stable), calls the streaming parser missed are appended under fresh indices, and when nothing was streamed the re-parse rebuilds state exactly as before; - malformed (non-JSON) accumulated arguments of already-streamed calls can no longer be retracted, so they now trigger the finish-time re-parse for authoritative recovery instead of being silently dropped; finish_reason is still rewritten to tool_calls when a call was emitted. Tests: existing assertions now reassemble the delta stream per the client contract (semantic regression), plus new cases locking the TTFT contract (deltas arrive before the finish chunk, first entry carries id+name, distinct indices/ids for parallel calls) and the new malformed/unknown-name streaming contract. Not executed locally (no sglang env); run on gpu_1 CI. --- dingo/frontend/sglang_prepost.py | 194 +++++++++++-- .../frontend/tests/test_sglang_tool_calls.py | 260 ++++++++++++++++-- 2 files changed, 413 insertions(+), 41 deletions(-) diff --git a/dingo/frontend/sglang_prepost.py b/dingo/frontend/sglang_prepost.py index d94663e7d47d..513df55d646e 100644 --- a/dingo/frontend/sglang_prepost.py +++ b/dingo/frontend/sglang_prepost.py @@ -1147,6 +1147,14 @@ class SglangStreamingPostProcessor: - Incremental detokenization via sliding-window decode (6-token lookback) - Reasoning content extraction via SGLang ReasoningParser - Tool call parsing via SGLang FunctionCallParser or JsonArrayParser + + Tool calls are streamed incrementally (id+name first, then argument + fragments) so clients see the first tool_call delta as soon as the + parser detects it — buffering the whole call until finish would make + the client-side TTFT of pure tool-call responses equal to the full + generation time. The finish-time full-text re-parse remains as the + authoritative safety net for calls or argument suffixes the + streaming parser missed. """ # Lookback window size for incremental detokenization. UTF-8 characters @@ -1231,6 +1239,23 @@ def __init__( self._kimi_k3_raw_text_parts: list[str] = [] self._saw_normal_output = False + # Incremental tool-call streaming state (TTFT fix). Parsed events + # are emitted as OpenAI deltas as they arrive instead of being + # buffered until finish: the id+name delta goes out with the first + # argument fragment, argument fragments stream as they come, and + # the finish-time re-parse only patches what streaming missed. + # Guards retained from the buffered design: + # - names without any argument fragment are never emitted (a + # misidentified prompt word cannot leak as a dangling call); + # - names not present in the request's tool list are suppressed + # entirely and purged at finish, exactly as before. + self._known_tool_names = ( + {t.function.name for t in self._sglang_tools} if self._sglang_tools else set() + ) + self._emitted_tool_names: set[int] = set() # indices whose id+name delta was sent + self._suppressed_tool_indices: set[int] = set() # unknown-name indices withheld + self._emitted_args_len: dict[int, int] = {} # index -> args length already sent + def _strip_trailing_eos_token_ids(self, token_ids: list[int]) -> list[int]: if not self._eos_token_ids: return token_ids @@ -1256,6 +1281,59 @@ def _tool_call_id(self, name: str, index: int) -> str: self.history_tool_calls_count, ) + def _streaming_tool_deltas( + self, tool_calls: list[Any] + ) -> list[dict[str, Any]]: + """Build incremental OpenAI tool_call deltas from parser events. + + Emission rules (TTFT fix: stream tool calls instead of buffering + them until finish): + + - The id+name delta is withheld until the first argument fragment + arrives, so a detected-but-never-argued name never reaches the + client; the finish-time logic drops such calls exactly as in + the old buffered design. + - Names not present in the request's tool list are suppressed + entirely (mirrors the finish-time known-name purge). + - Argument fragments stream as they arrive. The not-yet-emitted + tail of the accumulated arguments is the source of truth, so + event granularity (name and args in one event vs. split across + invocations) cannot produce duplicates. + """ + deltas: list[dict[str, Any]] = [] + for tc in tool_calls: + idx = tc.tool_index + if idx in self._suppressed_tool_indices: + continue + name = tc.name or self._tool_call_names.get(idx) + if tc.parameters and name and idx not in self._emitted_tool_names: + if self._known_tool_names and name not in self._known_tool_names: + # Unknown tool name: withhold the whole call. The + # accumulated state is purged at finish as before. + self._suppressed_tool_indices.add(idx) + continue + deltas.append( + { + "index": idx, + "id": self._tool_call_ids[idx], + "type": "function", + "function": {"name": name, "arguments": ""}, + } + ) + self._emitted_tool_names.add(idx) + if tc.parameters and idx in self._emitted_tool_names: + accumulated = "".join(self._tool_call_args.get(idx, [])) + sent = self._emitted_args_len.get(idx, 0) + if len(accumulated) > sent: + deltas.append( + { + "index": idx, + "function": {"arguments": accumulated[sent:]}, + } + ) + self._emitted_args_len[idx] = len(accumulated) + return deltas + def _incremental_decode( self, new_token_ids: list[int], @@ -1474,6 +1552,7 @@ def process_output(self, engine_response: dict[str, Any]) -> dict[str, Any] | No # -- Tool call parsing (accumulate deltas) -- content_text = normal_text + incremental_tool_deltas: list[dict[str, Any]] = [] if self.tool_call_parser and normal_text: # Accumulate raw text for finish-time re-parse. @@ -1499,6 +1578,8 @@ def process_output(self, engine_response: dict[str, Any]) -> dict[str, Any] | No if tc.parameters: self._tool_call_args.setdefault(idx, []).append(tc.parameters) + incremental_tool_deltas = self._streaming_tool_deltas(tool_calls) + if self._is_kimi_k3: content_text = _strip_kimi_k3_control_markers(content_text) if content_text: @@ -1514,6 +1595,13 @@ def process_output(self, engine_response: dict[str, Any]) -> dict[str, Any] | No if reasoning_text: delta["reasoning_content"] = reasoning_text has_content = True + if incremental_tool_deltas: + delta["tool_calls"] = incremental_tool_deltas + has_content = True + + # Argument suffixes recovered by the finish-time re-parse for calls + # whose identity was already streamed (patch-only emissions). + arg_patches: dict[int, str] = {} # On finish, re-parse the full accumulated text to recover tool # calls or arguments that the streaming parser missed. @@ -1527,7 +1615,11 @@ def process_output(self, engine_response: dict[str, Any]) -> dict[str, Any] | No # # The re-parse uses the accumulated text (not the parser's internal # _buffer, which is consumed during streaming) and assigns - # sequential indices to match the OpenAI API convention. + # sequential indices to match the OpenAI API convention. When + # incremental deltas were already streamed, the re-parse runs in + # merge mode: it patches argument suffixes for streamed calls and + # appends calls the streaming parser missed under fresh indices, + # keeping the streamed indices stable. if ( finish_reason and self.tool_call_parser is not None @@ -1536,12 +1628,10 @@ def process_output(self, engine_response: dict[str, Any]) -> dict[str, Any] | No # Purge streaming results that don't match any known tool. # When guided decoding is not enforced the streaming parser # can misidentify words in the prompt (e.g. a person's name) - # as function names. - known_names = ( - {t.function.name for t in self._sglang_tools} - if self._sglang_tools - else set() - ) + # as function names. Unknown names were already withheld + # from the incremental stream, so purging them here can never + # retract anything the client saw. + known_names = self._known_tool_names if known_names: for idx in list(self._tool_call_names): if self._tool_call_names[idx] not in known_names: @@ -1550,13 +1640,25 @@ def process_output(self, engine_response: dict[str, Any]) -> dict[str, Any] | No self._tool_call_args.pop(idx, None) # Discard malformed (non-JSON) argument fragments that the - # streaming parser accumulated from mixed content. + # streaming parser accumulated from mixed content. For calls + # whose arguments were already streamed the fragments cannot + # be retracted; instead flag them so the re-parse below can + # recover the authoritative arguments and patch the suffix. + streamed_malformed = False for idx in list(self._tool_call_args): combined = "".join(self._tool_call_args[idx]) if combined: try: json.loads(combined) except (json.JSONDecodeError, ValueError): + if idx in self._emitted_tool_names: + logger.warning( + "Tool call %s was streamed with malformed " + "arguments; attempting finish-time recovery", + self._tool_call_ids.get(idx), + ) + streamed_malformed = True + continue del self._tool_call_args[idx] missing_names = not self._tool_call_names @@ -1565,7 +1667,7 @@ def process_output(self, engine_response: dict[str, Any]) -> dict[str, Any] | No ) should_reparse = False full_text = "" - if missing_names or missing_args: + if missing_names or missing_args or streamed_malformed: full_text = "".join(self._tool_text_parts) # Skip the re-parse when the accumulated text has no # tool-call markers. Avoids wasted `parse_non_stream` @@ -1620,10 +1722,12 @@ def process_output(self, engine_response: dict[str, Any]) -> dict[str, Any] | No # Re-index sequentially so repeated calls to the same # tool get distinct indices (parse_non_stream may assign # indices based on the tool-definition position instead). - # When the re-parse returns results, it is authoritative: - # clear streaming state first so we don't mix a name from - # the re-parse with args from streaming at the same index. - if final_calls: + if final_calls and not self._emitted_tool_names: + # Nothing was streamed (e.g. all tokens arrived in one + # batch): the re-parse is authoritative. Clear + # streaming state first so we don't mix a name from + # the re-parse with args from streaming at the same + # index. self._tool_call_ids.clear() self._tool_call_names.clear() self._tool_call_args.clear() @@ -1635,11 +1739,58 @@ def process_output(self, engine_response: dict[str, Any]) -> dict[str, Any] | No self._tool_call_names[seq_idx] = tc.name if tc.parameters: self._tool_call_args[seq_idx] = [tc.parameters] + elif final_calls: + # Merge mode: incremental deltas already went out, so + # streamed indices must stay stable. Match recovered + # calls to streamed ones by name (in order) and emit + # only the missing argument suffix; calls the + # streaming parser missed entirely are appended under + # fresh indices. + streamed_by_name: dict[str, list[int]] = {} + for idx in sorted(self._emitted_tool_names): + nm = self._tool_call_names.get(idx) + if nm: + streamed_by_name.setdefault(nm, []).append(idx) + next_idx = max(self._emitted_tool_names) + 1 + for tc in final_calls: + candidates = streamed_by_name.get(tc.name or "") + if candidates: + idx = candidates.pop(0) + final_args = tc.parameters or "" + accumulated = "".join(self._tool_call_args.get(idx, [])) + if ( + final_args + and len(final_args) > len(accumulated) + and final_args.startswith(accumulated) + ): + arg_patches[idx] = final_args[len(accumulated):] + self._tool_call_args[idx] = [final_args] + elif final_args and final_args != accumulated: + logger.warning( + "Re-parsed arguments for tool call %s do " + "not extend the streamed fragments; " + "keeping the streamed version", + self._tool_call_ids.get(idx), + ) + else: + while next_idx in self._tool_call_names: + next_idx += 1 + self._tool_call_ids[next_idx] = self._tool_call_id( + tc.name or "", next_idx + ) + if tc.name: + self._tool_call_names[next_idx] = tc.name + if tc.parameters: + self._tool_call_args[next_idx] = [tc.parameters] + next_idx += 1 # Do not emit partial tool calls. A streaming parser can detect a # tool name before the model finishes malformed JSON; if the # finish-time re-parse cannot recover valid arguments, treat the # response as plain text instead of surfacing name + empty args. + # (Incrementally streamed names always carried at least one + # argument fragment, so this drop only affects calls that were + # never emitted to the client.) dropped_names = [] for idx in list(self._tool_call_names): if not "".join(self._tool_call_args.get(idx, [])): @@ -1667,9 +1818,15 @@ def process_output(self, engine_response: dict[str, Any]) -> dict[str, Any] | No has_content = True self._saw_normal_output = True - if finish_reason and self._tool_call_names: + # On finish, emit only what the incremental stream has not sent: + # complete calls recovered by the re-parse (never streamed) and + # argument suffix patches for streamed calls whose arguments the + # streaming parser only partially detected. + if finish_reason and (self._tool_call_names or arg_patches): tool_calls_out: list[dict[str, Any]] = [] for idx in sorted(self._tool_call_names): + if idx in self._emitted_tool_names: + continue # identity and arguments already streamed tool_calls_out.append( { "index": idx, @@ -1681,8 +1838,13 @@ def process_output(self, engine_response: dict[str, Any]) -> dict[str, Any] | No }, } ) - delta["tool_calls"] = tool_calls_out - has_content = True + for idx in sorted(arg_patches): + tool_calls_out.append( + {"index": idx, "function": {"arguments": arg_patches[idx]}} + ) + if tool_calls_out: + delta["tool_calls"] = tool_calls_out + has_content = True # Rewrite finish_reason "stop" → "tool_calls" when tool calls were # detected, matching the OpenAI API spec and official SGLang behaviour. diff --git a/dingo/frontend/tests/test_sglang_tool_calls.py b/dingo/frontend/tests/test_sglang_tool_calls.py index f0105940527c..8da8f94d3649 100644 --- a/dingo/frontend/tests/test_sglang_tool_calls.py +++ b/dingo/frontend/tests/test_sglang_tool_calls.py @@ -4,8 +4,9 @@ """Tests for tool call parsing in SglangStreamingPostProcessor. Covers the interaction between SGLang's FunctionCallParser, ReasoningParser, -and our post-processor's accumulate-and-emit-on-finish logic, including the -parse_non_stream fallback for the chunking-sensitivity issue in +and our post-processor's incremental tool-call streaming (id+name first, +then argument fragments), including the finish-time parse_non_stream +fallback/merge for the chunking-sensitivity issue in BaseFormatDetector.parse_streaming_increment. """ @@ -100,13 +101,41 @@ def _run_postprocessor(tokenizer, full_text, batch_size, *, use_reasoning=True): return results +def _merge_tool_call_entries(entries): + """Merge OpenAI streaming tool_call delta entries into complete calls. + + Mirrors the client-side reassembly contract: an entry carrying ``id`` + starts the call (name + empty arguments), entries without ``id`` + append argument fragments, and a finish-time recovered call arrives + as a single complete entry. + """ + calls: dict[int, dict] = {} + order: list[int] = [] + for e in entries: + idx = e["index"] + if idx not in calls: + calls[idx] = { + "index": idx, + "id": None, + "type": "function", + "function": {"name": None, "arguments": ""}, + } + order.append(idx) + fn = e.get("function", {}) + if e.get("id"): + calls[idx]["id"] = e["id"] + if fn.get("name"): + calls[idx]["function"]["name"] = fn["name"] + calls[idx]["function"]["arguments"] += fn.get("arguments", "") + return [calls[i] for i in order] + + def _extract_tool_calls(results): - """Extract tool_calls from the list of choices.""" + """Reassemble complete tool calls from incremental deltas across choices.""" + entries = [] for r in results: - tc = r.get("delta", {}).get("tool_calls") - if tc: - return tc - return [] + entries.extend(r.get("delta", {}).get("tool_calls") or []) + return _merge_tool_call_entries(entries) # --------------------------------------------------------------------------- @@ -194,7 +223,7 @@ def parse_stream_chunk(self, text): } ) - tc = choice["delta"]["tool_calls"] + tc = _merge_tool_call_entries(choice["delta"]["tool_calls"]) assert [item["id"] for item in tc] == [ "functions.get_weather:3", "functions.search_gutenberg_books:4", @@ -254,7 +283,7 @@ def parse_non_stream(self, text): } ) - tc = choice["delta"]["tool_calls"] + tc = _merge_tool_call_entries(choice["delta"]["tool_calls"]) # IDs must use seq_idx (0, 1) + history (3), not tool_index (5, 2). assert [item["id"] for item in tc] == [ "functions.get_weather:3", @@ -417,7 +446,7 @@ def test_all_tokens_plus_finish_in_one_batch(self, tokenizer): # Feed ALL tokens at once with finish_reason choice = post.process_output({"token_ids": token_ids, "finish_reason": "stop"}) assert choice is not None - tc = choice.get("delta", {}).get("tool_calls", []) + tc = _merge_tool_call_entries(choice.get("delta", {}).get("tool_calls", [])) assert len(tc) == 1, f"Expected 1 tool call, got {len(tc)}" assert tc[0]["function"]["name"] == "search_gutenberg_books" args = json.loads(tc[0]["function"]["arguments"]) @@ -442,7 +471,7 @@ def test_multiple_tools_single_chunk(self, tokenizer): token_ids = tokenizer.encode(text) choice = post.process_output({"token_ids": token_ids, "finish_reason": "stop"}) assert choice is not None - tc = choice.get("delta", {}).get("tool_calls", []) + tc = _merge_tool_call_entries(choice.get("delta", {}).get("tool_calls", [])) assert len(tc) == 2, f"Expected 2 tool calls, got {len(tc)}" names = {t["function"]["name"] for t in tc} assert names == {"search_gutenberg_books", "get_weather"} @@ -469,20 +498,36 @@ def test_finish_reason_rewritten_to_tool_calls(self, tokenizer): class TestMalformedToolCalls: # FRONTEND.4 — malformed model output → graceful degradation - def test_incomplete_arguments_are_not_emitted(self): - class DummyTokenizer: - def decode(self, token_ids, skip_special_tokens=True): - return "".join(chr(x) for x in token_ids) + """Contract under incremental streaming: + + - A name detected without any argument fragment is never emitted + (neither mid-stream nor at finish) and does not rewrite the + finish_reason. + - An unknown tool name is suppressed mid-stream and purged at + finish; nothing reaches the client. + - A known name with malformed (non-JSON) arguments IS streamed + optimistically — matching native SGLang behaviour — and the + finish-time re-parse attempts authoritative recovery. + """ - class DummyToolCall: - def __init__(self, tool_index, name, parameters): - self.tool_index = tool_index - self.name = name - self.parameters = parameters + class DummyTokenizer: + def decode(self, token_ids, skip_special_tokens=True): + return "".join(chr(x) for x in token_ids) + + class DummyToolCall: + def __init__(self, tool_index, name, parameters): + self.tool_index = tool_index + self.name = name + self.parameters = parameters + + def test_name_without_arguments_is_never_emitted(self): + dummy_tokenizer = self.DummyTokenizer() + dummy_tc = self.DummyToolCall class DummyParser: def parse_stream_chunk(self, text): - return "", [DummyToolCall(0, "get_weather", '{"city": "Paris"')] + # Name event only — no argument fragment ever arrives. + return "", [dummy_tc(0, "get_weather", None)] def has_tool_call(self, text): return "" in text @@ -491,7 +536,7 @@ def parse_non_stream(self, text): return "", [] post = SglangStreamingPostProcessor( - tokenizer=DummyTokenizer(), + tokenizer=dummy_tokenizer, tool_call_parser=DummyParser(), reasoning_parser=None, ) @@ -511,6 +556,81 @@ def parse_non_stream(self, text): assert choice["finish_reason"] == "stop" assert choice.get("delta", {}).get("tool_calls", []) == [] + def test_unknown_tool_name_is_never_streamed(self): + dummy_tokenizer = self.DummyTokenizer() + dummy_tc = self.DummyToolCall + + class DummyParser: + def parse_stream_chunk(self, text): + return "", [dummy_tc(0, "evil_tool", '{"x": 1}')] + + def has_tool_call(self, text): + return True + + def parse_non_stream(self, text): + return "", [] + + post = SglangStreamingPostProcessor( + tokenizer=dummy_tokenizer, + tool_call_parser=DummyParser(), + reasoning_parser=None, + sglang_tools=TOOLS, + ) + + text = '\n{"name": "evil_tool", "arguments": {"x": 1}}\n' + choice = post.process_output( + { + "token_ids": [ord(c) for c in text], + "finish_reason": "stop", + } + ) + + assert choice is not None + assert choice["finish_reason"] == "stop" + assert choice.get("delta", {}).get("tool_calls", []) == [] + + def test_malformed_arguments_stream_optimistically(self): + dummy_tokenizer = self.DummyTokenizer() + dummy_tc = self.DummyToolCall + + class DummyParser: + def parse_stream_chunk(self, text): + # Known name with malformed (unrecoverable) arguments. + return "", [dummy_tc(0, "get_weather", '{"city": "Paris"')] + + def has_tool_call(self, text): + return True + + def parse_non_stream(self, text): + return "", [] + + post = SglangStreamingPostProcessor( + tokenizer=dummy_tokenizer, + tool_call_parser=DummyParser(), + reasoning_parser=None, + sglang_tools=TOOLS, + ) + + malformed = ( + '\n{"name": "get_weather", ' + '"arguments": {"city": "Paris"}\n' + ) + choice = post.process_output( + { + "token_ids": [ord(c) for c in malformed], + "finish_reason": "stop", + } + ) + + # Native-compatible tradeoff: the fragments went out before the + # JSON could be validated; the call counts as emitted, so the + # finish_reason is rewritten to tool_calls. + assert choice is not None + assert choice["finish_reason"] == "tool_calls" + tc = _merge_tool_call_entries(choice.get("delta", {}).get("tool_calls", [])) + assert tc[0]["function"]["name"] == "get_weather" + assert tc[0]["function"]["arguments"] == '{"city": "Paris"' + # --------------------------------------------------------------------------- # JsonArrayParser path (tool_choice="required" / named function) @@ -541,7 +661,7 @@ def test_single_call_reparse(self, tokenizer): token_ids = tokenizer.encode(text) choice = post.process_output({"token_ids": token_ids, "finish_reason": "stop"}) assert choice is not None - tc = choice.get("delta", {}).get("tool_calls", []) + tc = _merge_tool_call_entries(choice.get("delta", {}).get("tool_calls", [])) assert len(tc) == 1 assert tc[0]["function"]["name"] == "get_weather" assert json.loads(tc[0]["function"]["arguments"]) == {"city": "NYC"} @@ -563,7 +683,7 @@ def test_multiple_calls_reparse(self, tokenizer): token_ids = tokenizer.encode(text) choice = post.process_output({"token_ids": token_ids, "finish_reason": "stop"}) assert choice is not None - tc = choice.get("delta", {}).get("tool_calls", []) + tc = _merge_tool_call_entries(choice.get("delta", {}).get("tool_calls", [])) assert len(tc) == 2 names = {t["function"]["name"] for t in tc} assert names == {"search_gutenberg_books", "get_weather"} @@ -584,5 +704,95 @@ def test_plain_text_skips_reparse(self, tokenizer): token_ids = tokenizer.encode("Hello, world!") choice = post.process_output({"token_ids": token_ids, "finish_reason": "stop"}) # No tool calls, plain content preserved, no crash. - tc = (choice or {}).get("delta", {}).get("tool_calls", []) + tc = _merge_tool_call_entries((choice or {}).get("delta", {}).get("tool_calls", [])) assert tc == [] + + +# --------------------------------------------------------------------------- +# Incremental tool-call streaming (TTFT) +# --------------------------------------------------------------------------- + + +class TestIncrementalToolStreaming: # FRONTEND.4 — tool_call deltas stream before finish + """TTFT contract: clients must see tool_call deltas as the parser + detects them instead of everything arriving with the finish chunk.""" + + TEXT = ( + '\n{"name": "get_weather", ' + '"arguments": {"city": "Paris"}}\n' + ) + + MULTI_TEXT = ( + '\n{"name": "search_gutenberg_books", ' + '"arguments": {"search_terms": ["Joyce"]}}\n\n' + '\n{"name": "get_weather", ' + '"arguments": {"city": "London"}}\n' + ) + + @staticmethod + def _all_entries(results): + return [ + e + for r in results + for e in (r.get("delta", {}).get("tool_calls") or []) + ] + + def test_tool_deltas_arrive_before_finish(self, tokenizer): + results = _run_postprocessor(tokenizer, self.TEXT, 3, use_reasoning=False) + assert results[-1]["finish_reason"] is not None + first_tool_idx = next( + i for i, r in enumerate(results) if r.get("delta", {}).get("tool_calls") + ) + assert first_tool_idx < len(results) - 1, ( + "tool_call deltas must stream before the finish chunk " + "(buffering them is what inflated client-side TTFT)" + ) + + def test_first_entry_carries_id_and_name(self, tokenizer): + results = _run_postprocessor(tokenizer, self.TEXT, 3, use_reasoning=False) + entries = self._all_entries(results) + assert entries, "expected streamed tool_call entries" + first = entries[0] + assert first["id"].startswith("call_") + assert first["type"] == "function" + assert first["function"]["name"] == "get_weather" + assert first["function"]["arguments"] == "" + + def test_reassembled_arguments_complete(self, tokenizer): + tc = _extract_tool_calls( + _run_postprocessor(tokenizer, self.TEXT, 3, use_reasoning=False) + ) + assert len(tc) == 1 + assert tc[0]["function"]["name"] == "get_weather" + assert json.loads(tc[0]["function"]["arguments"]) == {"city": "Paris"} + + def test_multiple_calls_stream_with_distinct_indices(self, tokenizer): + results = _run_postprocessor(tokenizer, self.MULTI_TEXT, 5, use_reasoning=False) + entries = self._all_entries(results) + assert entries, "expected streamed tool_call entries" + tc = _merge_tool_call_entries(entries) + assert {t["function"]["name"] for t in tc} == { + "search_gutenberg_books", + "get_weather", + } + assert len({t["index"] for t in tc}) == 2 + ids = [t["id"] for t in tc] + assert len(set(ids)) == len(ids), "tool call ids must be unique" + + def test_finish_reason_rewritten_to_tool_calls(self, tokenizer): + results = _run_postprocessor(tokenizer, self.TEXT, 3, use_reasoning=False) + assert results[-1]["finish_reason"] == "tool_calls" + + def test_pure_tool_call_first_choice_is_not_delayed(self, tokenizer): + """Pure tool-call responses (no reasoning preface) must produce a + tool_call-bearing choice well before the final one — this is the + scenario where buffered emission made TTFT equal to the whole + generation time.""" + results = _run_postprocessor(tokenizer, self.TEXT, 3, use_reasoning=False) + first_tool_idx = next( + i for i, r in enumerate(results) if r.get("delta", {}).get("tool_calls") + ) + # The name completes after roughly the first third of the tokens; + # allow generous slack but require strict precedence over finish. + assert first_tool_idx <= len(results) // 2 + From dd44d0a1c8cbe961b190d4507e9f1fb3f3003f36 Mon Sep 17 00:00:00 2001 From: githubgxll <1094462054@qq.com> Date: Wed, 23 Sep 2026 17:25:12 +0800 Subject: [PATCH 2/3] fix(frontend): preserve terminal tool deltas and stream confirmed names --- dingo/frontend/sglang_prepost.py | 55 ++--- .../frontend/tests/test_sglang_tool_calls.py | 198 +++++++++++++++++- 2 files changed, 222 insertions(+), 31 deletions(-) diff --git a/dingo/frontend/sglang_prepost.py b/dingo/frontend/sglang_prepost.py index 513df55d646e..94a9cbb10ab1 100644 --- a/dingo/frontend/sglang_prepost.py +++ b/dingo/frontend/sglang_prepost.py @@ -1241,18 +1241,21 @@ def __init__( # Incremental tool-call streaming state (TTFT fix). Parsed events # are emitted as OpenAI deltas as they arrive instead of being - # buffered until finish: the id+name delta goes out with the first - # argument fragment, argument fragments stream as they come, and + # buffered until finish: confirmed names go out immediately, + # argument fragments stream as they come, and # the finish-time re-parse only patches what streaming missed. # Guards retained from the buffered design: - # - names without any argument fragment are never emitted (a - # misidentified prompt word cannot leak as a dangling call); + # - without a request tool list, require an argument fragment + # before emitting a name that cannot be independently confirmed; # - names not present in the request's tool list are suppressed # entirely and purged at finish, exactly as before. self._known_tool_names = ( - {t.function.name for t in self._sglang_tools} if self._sglang_tools else set() + {t.function.name for t in self._sglang_tools} + if self._sglang_tools + else set() ) - self._emitted_tool_names: set[int] = set() # indices whose id+name delta was sent + # Indices whose id+name delta was sent. + self._emitted_tool_names: set[int] = set() self._suppressed_tool_indices: set[int] = set() # unknown-name indices withheld self._emitted_args_len: dict[int, int] = {} # index -> args length already sent @@ -1281,18 +1284,15 @@ def _tool_call_id(self, name: str, index: int) -> str: self.history_tool_calls_count, ) - def _streaming_tool_deltas( - self, tool_calls: list[Any] - ) -> list[dict[str, Any]]: + def _streaming_tool_deltas(self, tool_calls: list[Any]) -> list[dict[str, Any]]: """Build incremental OpenAI tool_call deltas from parser events. Emission rules (TTFT fix: stream tool calls instead of buffering them until finish): - - The id+name delta is withheld until the first argument fragment - arrives, so a detected-but-never-argued name never reaches the - client; the finish-time logic drops such calls exactly as in - the old buffered design. + - Names confirmed by the request's tool list are emitted immediately, + even when the parser needs to buffer arguments (e.g. GLM47 union + schemas). Without that confirmation, wait for an argument fragment. - Names not present in the request's tool list are suppressed entirely (mirrors the finish-time known-name purge). - Argument fragments stream as they arrive. The not-yet-emitted @@ -1306,7 +1306,11 @@ def _streaming_tool_deltas( if idx in self._suppressed_tool_indices: continue name = tc.name or self._tool_call_names.get(idx) - if tc.parameters and name and idx not in self._emitted_tool_names: + if ( + name + and (tc.parameters or name in self._known_tool_names) + and idx not in self._emitted_tool_names + ): if self._known_tool_names and name not in self._known_tool_names: # Unknown tool name: withhold the whole call. The # accumulated state is purged at finish as before. @@ -1763,7 +1767,7 @@ def process_output(self, engine_response: dict[str, Any]) -> dict[str, Any] | No and len(final_args) > len(accumulated) and final_args.startswith(accumulated) ): - arg_patches[idx] = final_args[len(accumulated):] + arg_patches[idx] = final_args[len(accumulated) :] self._tool_call_args[idx] = [final_args] elif final_args and final_args != accumulated: logger.warning( @@ -1784,16 +1788,16 @@ def process_output(self, engine_response: dict[str, Any]) -> dict[str, Any] | No self._tool_call_args[next_idx] = [tc.parameters] next_idx += 1 - # Do not emit partial tool calls. A streaming parser can detect a - # tool name before the model finishes malformed JSON; if the - # finish-time re-parse cannot recover valid arguments, treat the - # response as plain text instead of surfacing name + empty args. - # (Incrementally streamed names always carried at least one - # argument fragment, so this drop only affects calls that were - # never emitted to the client.) + # Drop incomplete calls only when their identity was never sent. + # A confirmed name may already be visible while its arguments are + # buffered; it cannot be retracted if final recovery fails. Keep + # its identity and tool_calls finish reason, without fabricating + # arguments for an incomplete model output. dropped_names = [] for idx in list(self._tool_call_names): - if not "".join(self._tool_call_args.get(idx, [])): + if idx not in self._emitted_tool_names and not "".join( + self._tool_call_args.get(idx, []) + ): dropped_names.append(self._tool_call_names[idx]) del self._tool_call_names[idx] self._tool_call_ids.pop(idx, None) @@ -1843,7 +1847,10 @@ def process_output(self, engine_response: dict[str, Any]) -> dict[str, Any] | No {"index": idx, "function": {"arguments": arg_patches[idx]}} ) if tool_calls_out: - delta["tool_calls"] = tool_calls_out + # This terminal batch may already contain argument deltas or + # identities. Preserve them before appending recovery output: + # _emitted_* tracks constructed deltas, not yet-yielded bytes. + delta.setdefault("tool_calls", []).extend(tool_calls_out) has_content = True # Rewrite finish_reason "stop" → "tool_calls" when tool calls were diff --git a/dingo/frontend/tests/test_sglang_tool_calls.py b/dingo/frontend/tests/test_sglang_tool_calls.py index 8da8f94d3649..132bad6d1bfd 100644 --- a/dingo/frontend/tests/test_sglang_tool_calls.py +++ b/dingo/frontend/tests/test_sglang_tool_calls.py @@ -11,10 +11,12 @@ """ import json +from typing import Any import pytest from sglang.srt.entrypoints.openai.protocol import Function as SglangFunction from sglang.srt.entrypoints.openai.protocol import Tool as SglangTool +from sglang.srt.function_call.core_types import ToolCallItem from sglang.srt.function_call.function_call_parser import FunctionCallParser from sglang.srt.function_call.json_array_parser import JsonArrayParser from sglang.srt.parser.reasoning_parser import ReasoningParser @@ -500,9 +502,8 @@ def test_finish_reason_rewritten_to_tool_calls(self, tokenizer): class TestMalformedToolCalls: # FRONTEND.4 — malformed model output → graceful degradation """Contract under incremental streaming: - - A name detected without any argument fragment is never emitted - (neither mid-stream nor at finish) and does not rewrite the - finish_reason. + - Without tools-list confirmation, a name detected without any argument + fragment is never emitted and does not rewrite the finish_reason. - An unknown tool name is suppressed mid-stream and purged at finish; nothing reaches the client. - A known name with malformed (non-JSON) arguments IS streamed @@ -704,7 +705,9 @@ def test_plain_text_skips_reparse(self, tokenizer): token_ids = tokenizer.encode("Hello, world!") choice = post.process_output({"token_ids": token_ids, "finish_reason": "stop"}) # No tool calls, plain content preserved, no crash. - tc = _merge_tool_call_entries((choice or {}).get("delta", {}).get("tool_calls", [])) + tc = _merge_tool_call_entries( + (choice or {}).get("delta", {}).get("tool_calls", []) + ) assert tc == [] @@ -732,9 +735,7 @@ class TestIncrementalToolStreaming: # FRONTEND.4 — tool_call deltas stream be @staticmethod def _all_entries(results): return [ - e - for r in results - for e in (r.get("delta", {}).get("tool_calls") or []) + e for r in results for e in (r.get("delta", {}).get("tool_calls") or []) ] def test_tool_deltas_arrive_before_finish(self, tokenizer): @@ -796,3 +797,186 @@ def test_pure_tool_call_first_choice_is_not_delayed(self, tokenizer): # allow generous slack but require strict precedence over finish. assert first_tool_idx <= len(results) // 2 + +class TestToolStreamingRecoveryRegression: + """Exercise terminal recovery and early identity without model downloads.""" + + class Tokenizer: + def decode(self, token_ids: list[int], skip_special_tokens: bool = True) -> str: + return "".join(chr(token) for token in token_ids) + + class Parser: + def __init__( + self, events: list[list[ToolCallItem]], recovered: list[ToolCallItem] + ) -> None: + self.events = iter(events) + self.recovered = recovered + + def parse_stream_chunk(self, text: str) -> tuple[str, list[ToolCallItem]]: + return "", next(self.events) + + def has_tool_call(self, text: str) -> bool: + return True + + def parse_non_stream(self, text: str) -> tuple[str, list[ToolCallItem]]: + return "", self.recovered + + @staticmethod + def call(index: int, name: str | None, arguments: str) -> ToolCallItem: + return ToolCallItem(tool_index=index, name=name, parameters=arguments) + + def make_post( + self, + events: list[list[ToolCallItem]], + recovered: list[ToolCallItem], + *, + confirm: bool = True, + ) -> SglangStreamingPostProcessor: + return SglangStreamingPostProcessor( + tokenizer=self.Tokenizer(), + tool_call_parser=self.Parser(events, recovered), + reasoning_parser=None, + sglang_tools=TOOLS if confirm else None, + ) + + @staticmethod + def feed( + post: SglangStreamingPostProcessor, text: str = "x", *, finish: bool = False + ) -> dict[str, Any] | None: + return post.process_output( + { + "token_ids": [ord(c) for c in text], + "finish_reason": "stop" if finish else None, + } + ) + + @pytest.mark.parametrize("confirm", [False, True]) + @pytest.mark.parametrize("all_at_finish", [False, True]) + def test_terminal_recovery_preserves_current_deltas( + self, confirm: bool, all_at_finish: bool + ) -> None: + call = self.call + final_args = '{"city":"Paris"}' + tail = final_args if all_at_finish else '"Paris"}' + events = [] if all_at_finish else [[call(0, "get_weather", '{"city":')]] + events.append( + [ + call(0, "get_weather" if all_at_finish else None, tail), + call(1, "search_gutenberg_books", ""), + ] + ) + post = self.make_post( + events, + [ + call(0, "get_weather", final_args), + call(1, "search_gutenberg_books", '{"search_terms":["Joyce"]}'), + ], + confirm=confirm, + ) + choices = [] if all_at_finish else [self.feed(post)] + choices.append(self.feed(post, finish=True)) + merged = _extract_tool_calls([c for c in choices if c]) + assert len(merged) == 2 + by_name = {c["function"]["name"]: c for c in merged} + assert by_name["get_weather"]["function"]["arguments"] == final_args + assert json.loads( + by_name["search_gutenberg_books"]["function"]["arguments"] + ) == {"search_terms": ["Joyce"]} + assert len({c["id"] for c in merged}) == 2 + assert choices[-1]["finish_reason"] == "tool_calls" + + def test_confirmed_name_precedes_arguments_and_keeps_identity(self) -> None: + call = self.call + post = self.make_post( + [[call(0, "get_weather", "")], [], [call(0, None, '{"city":"Paris"}')]], + [], + ) + first = self.feed(post) + assert first["finish_reason"] is None + identity = first["delta"]["tool_calls"][0] + assert identity["function"] == {"name": "get_weather", "arguments": ""} + assert self.feed(post) is None + final = self.feed(post, finish=True) + entries = first["delta"]["tool_calls"] + final["delta"]["tool_calls"] + assert sum(bool(e.get("id")) for e in entries) == 1 + merged = _merge_tool_call_entries(entries) + assert merged[0]["id"] == identity["id"] + assert merged[0]["function"]["arguments"] == '{"city":"Paris"}' + + def test_name_only_call_gets_recovered_arguments_on_same_index(self) -> None: + call = self.call + post = self.make_post( + [[call(0, "get_weather", "")]], + [call(0, "get_weather", '{"city":"Paris"}')], + ) + first = self.feed(post) + final = self.feed(post, text="", finish=True) + merged = _extract_tool_calls([first, final]) + assert len(merged) == 1 + assert merged[0]["index"] == 0 + assert merged[0]["function"]["arguments"] == '{"city":"Paris"}' + assert "id" not in final["delta"]["tool_calls"][0] + + def test_unrecoverable_name_is_not_retracted_or_given_fake_arguments(self) -> None: + post = self.make_post([[self.call(0, "get_weather", "")]], []) + first = self.feed(post) + final = self.feed(post, text="", finish=True) + assert first["delta"]["tool_calls"][0]["function"]["name"] == "get_weather" + assert final["finish_reason"] == "tool_calls" + assert not final["delta"].get("tool_calls") + assert _extract_tool_calls([first, final])[0]["function"]["arguments"] == "" + + def test_unknown_name_only_is_suppressed(self) -> None: + post = self.make_post([[self.call(0, "unknown_tool", "")]], []) + assert self.feed(post) is None + final = self.feed(post, text="", finish=True) + assert final["finish_reason"] == "stop" + assert not final["delta"].get("tool_calls") + + def test_glm47_union_schema_emits_identity_before_tool_closes(self) -> None: + tools = [ + SglangTool( + type="function", + function=SglangFunction( + name="lookup", + parameters={ + "oneOf": [ + { + "type": "object", + "properties": {"value": {"type": "string"}}, + }, + { + "type": "object", + "properties": {"value": {"type": "integer"}}, + }, + ] + }, + ), + ) + ] + post = SglangStreamingPostProcessor( + tokenizer=self.Tokenizer(), + tool_call_parser=FunctionCallParser(tools=tools, tool_call_parser="glm47"), + reasoning_parser=None, + sglang_tools=tools, + ) + chunks = [ + "lookupvalue", + "long value part 1", + "long value part 2", + "", + ] + first = self.feed(post, chunks[0]) + assert first is not None + assert first["finish_reason"] is None + assert first["delta"]["tool_calls"][0]["function"]["name"] == "lookup" + choices = [first] + for i, chunk in enumerate(chunks[1:], 1): + choice = self.feed(post, chunk, finish=i == len(chunks) - 1) + if choice: + choices.append(choice) + merged = _extract_tool_calls(choices) + assert len(merged) == 1 + assert json.loads(merged[0]["function"]["arguments"]) == { + "value": "long value part 1long value part 2" + } From d22fac41b85a46c049733b1f15df44b921fe3122 Mon Sep 17 00:00:00 2001 From: githubgxll <1094462054@qq.com> Date: Wed, 23 Sep 2026 18:17:24 +0800 Subject: [PATCH 3/3] fix(frontend): recover tool calls entirely missed during streaming --- dingo/frontend/sglang_prepost.py | 28 ++++------- .../frontend/tests/test_sglang_tool_calls.py | 48 +++++++++++++++++++ 2 files changed, 58 insertions(+), 18 deletions(-) diff --git a/dingo/frontend/sglang_prepost.py b/dingo/frontend/sglang_prepost.py index 94a9cbb10ab1..c253b7cfcf5a 100644 --- a/dingo/frontend/sglang_prepost.py +++ b/dingo/frontend/sglang_prepost.py @@ -1646,9 +1646,8 @@ def process_output(self, engine_response: dict[str, Any]) -> dict[str, Any] | No # Discard malformed (non-JSON) argument fragments that the # streaming parser accumulated from mixed content. For calls # whose arguments were already streamed the fragments cannot - # be retracted; instead flag them so the re-parse below can + # be retracted; log them and let the re-parse below attempt to # recover the authoritative arguments and patch the suffix. - streamed_malformed = False for idx in list(self._tool_call_args): combined = "".join(self._tool_call_args[idx]) if combined: @@ -1661,26 +1660,19 @@ def process_output(self, engine_response: dict[str, Any]) -> dict[str, Any] | No "arguments; attempting finish-time recovery", self._tool_call_ids.get(idx), ) - streamed_malformed = True continue del self._tool_call_args[idx] - missing_names = not self._tool_call_names - missing_args = any( - idx not in self._tool_call_args for idx in self._tool_call_names + # Complete arguments for observed calls do not prove that every + # call was observed: a streaming detector can miss an entire later + # call in a token batch. Reconcile all marked tool-call responses + # at finish, including when every streamed call is already valid. + full_text = "".join(self._tool_text_parts) + # Plain text still skips recovery, avoiding unnecessary parsing + # and detectors that reject input without tool-call markers. + should_reparse = bool(full_text) and self.tool_call_parser.has_tool_call( + full_text ) - should_reparse = False - full_text = "" - if missing_names or missing_args or streamed_malformed: - full_text = "".join(self._tool_text_parts) - # Skip the re-parse when the accumulated text has no - # tool-call markers. Avoids wasted `parse_non_stream` - # work on plain-text responses (common when tools are - # offered but the model replies without calling any) and - # guards against detectors that raise on arbitrary input. - should_reparse = bool( - full_text - ) and self.tool_call_parser.has_tool_call(full_text) if should_reparse: if self._is_json_array_parser: diff --git a/dingo/frontend/tests/test_sglang_tool_calls.py b/dingo/frontend/tests/test_sglang_tool_calls.py index 132bad6d1bfd..b282c7e290ac 100644 --- a/dingo/frontend/tests/test_sglang_tool_calls.py +++ b/dingo/frontend/tests/test_sglang_tool_calls.py @@ -885,6 +885,54 @@ def test_terminal_recovery_preserves_current_deltas( assert len({c["id"] for c in merged}) == 2 assert choices[-1]["finish_reason"] == "tool_calls" + @pytest.mark.parametrize("same_name", [False, True]) + @pytest.mark.parametrize("empty_finish", [False, True]) + def test_completely_missed_call_is_recovered( + self, same_name: bool, empty_finish: bool + ) -> None: + first_args = '{"city":"Paris"}' + second_name = "get_weather" if same_name else "search_gutenberg_books" + second_args = '{"city":"Rome"}' if same_name else '{"search_terms":["Joyce"]}' + post = self.make_post( + [[self.call(0, "get_weather", first_args)], []], + [ + self.call(0, "get_weather", first_args), + # Non-stream indices can be tool-definition indices, including + # the same index for two calls of the same tool. + self.call(0 if same_name else 1, second_name, second_args), + ], + ) + first = self.feed(post) + assert first is not None + assert set(post._tool_call_names.values()) == {"get_weather"} + assert post._tool_call_args == {0: [first_args]} + final = self.feed(post, text="" if empty_finish else "x", finish=True) + assert final is not None + entries = first["delta"]["tool_calls"] + final["delta"]["tool_calls"] + merged = _merge_tool_call_entries(entries) + assert [c["index"] for c in merged] == [0, 1] + assert [c["function"]["name"] for c in merged] == ["get_weather", second_name] + assert [c["function"]["arguments"] for c in merged] == [first_args, second_args] + assert len({c["id"] for c in merged}) == 2 + assert sum(bool(e.get("id")) for e in entries) == 2 + assert all(e["index"] == 1 for e in final["delta"]["tool_calls"]) + assert final["finish_reason"] == "tool_calls" + + def test_plain_text_without_tool_markers_skips_reparse(self) -> None: + class PlainTextParser(self.Parser): + def has_tool_call(self, text: str) -> bool: + return False + + def parse_non_stream(self, text: str) -> tuple[str, list[ToolCallItem]]: + pytest.fail("Plain text should not require tool-call recovery") + + post = self.make_post([], []) + post.tool_call_parser = PlainTextParser([[]], []) + final = self.feed(post, "hello", finish=True) + assert final is not None + assert final["finish_reason"] == "stop" + assert not final["delta"].get("tool_calls") + def test_confirmed_name_precedes_arguments_and_keeps_identity(self) -> None: call = self.call post = self.make_post(