diff --git a/core/llm_transport.py b/core/llm_transport.py index f3859c3..327eda0 100644 --- a/core/llm_transport.py +++ b/core/llm_transport.py @@ -17,13 +17,22 @@ from __future__ import annotations import json -from typing import Any, Callable, Dict, List, Optional, Tuple +from collections.abc import Callable +from typing import Any, Dict, List, Optional, Tuple from .storage import record_empty_response, record_malformed_tool_call - # ─────────────────────── delta / usage 提取 ─────────────────────── + +class PreferNonstreamToolCall(Exception): + """流式首包已选中不适合分片传输的工具,请调用方改走非流式。""" + + def __init__(self, tool: str) -> None: + super().__init__(tool) + self.tool = tool + + def extract_delta_content(chunk: Any) -> Optional[str]: """从 stream chunk 提 delta.content(文本片段)。chunk 形态 litellm ModelResponseStream: choices[0].delta.content。usage-only 收尾 chunk(没 choices / delta)返 None。 @@ -62,6 +71,29 @@ def extract_delta_reasoning(chunk: Any) -> Optional[str]: return None +def extract_delta_tool_names(chunk: Any) -> List[str]: + """提取本 chunk 新出现的 tool call function name;兼容对象和 dict 形态。""" + try: + choices = getattr(chunk, "choices", None) + if not choices: + return [] + delta = getattr(choices[0], "delta", None) + if delta is None: + return [] + tool_calls = getattr(delta, "tool_calls", None) or [] + names: List[str] = [] + for tc in tool_calls: + fn = tc.get("function") if isinstance(tc, dict) else getattr(tc, "function", None) + if fn is None: + continue + name = fn.get("name") if isinstance(fn, dict) else getattr(fn, "name", None) + if name: + names.append(str(name)) + return names + except Exception: + return [] + + def usage_to_dict(usage: Any) -> dict: if not usage: return {} @@ -375,7 +407,22 @@ def robust_stream( # 非流式重试期间用户点了停止(线程级 poll,见 loop._nonstream_once) return None, True else: - response, cancelled = collect_stream(llm_messages) + cancelled = False + try: + response, cancelled = collect_stream(llm_messages) + except PreferNonstreamToolCall as reroute: + # function.name 通常早于大段 arguments 到达。此时立刻关流并由 provider + # 一次性拼好 JSON,可避开 DeepSeek 长 write/edit 参数的 delta 错位。 + emit({ + "type": "warn", + "level": "info", + "msg": f"检测到大参数工具 {reroute.tool},已切换稳定传输模式", + }) + emit(dict(llm_start_event)) + response = nonstream(llm_messages) + if response is None: + return None, True + use_nonstream = True if cancelled: return None, True diff --git a/core/loop.py b/core/loop.py index f2e2cc2..5e68f7f 100644 --- a/core/loop.py +++ b/core/loop.py @@ -16,38 +16,39 @@ import threading import time from pathlib import Path from typing import Any, Callable, Dict, List, Optional, Tuple - from uuid import UUID import litellm +from . import pptx_guard +from .artifacts import MAX_ARTIFACTS_PER_MESSAGE from .capabilities import ModelCapabilities from .context import ( CHARS_PER_TOKEN, COMPACT_CONTEXT_RATIO, calibrated_chars_per_token, - filter_reasoning_for_replay, clamp_ratio, + filter_reasoning_for_replay, prepare_messages_with_stats, ) from .context_fold import maybe_fold from .executor import ExecCtx, Executor -from .artifacts import MAX_ARTIFACTS_PER_MESSAGE from .llm import LLM from .llm_transport import ( + PreferNonstreamToolCall, extract_delta_content, extract_delta_reasoning, + extract_delta_tool_names, extract_usage_details, robust_stream, ) from .salvage import salvage_tool_arguments from .session import Session -from .task_actions import DeferredTaskActions from .storage import ( record_chat_usage, record_salvaged_tool_call, ) -from . import pptx_guard +from .task_actions import DeferredTaskActions # 产物机检只挂能落盘的执行类工具(fs 写工具不适合造 pptx,机检无意义) _PPTX_GUARD_TOOLS = ("shell", "run_python") @@ -431,6 +432,11 @@ class AgentLoop: # 流式 delta 错位),历史数据里非流式兜底从未再畸形。 _MAX_MALFORMED_ATTEMPTS = 3 + # DeepSeek 的长 write/edit arguments 在流式 delta 中偶发错位。function.name 首包 + # 到达时立即关流并非流式重发;正文和其他工具仍走流式。只限定已实证的模型族, + # 避免把其他 provider 的正常工具调用无端降级。 + _DEEPSEEK_NONSTREAM_TOOLS = frozenset({"write", "edit"}) + # 连续多少步「整步无净产出」(全是 [Error]/重复结果/被拦)就判定空转、主动停。 # 比 max_iterations 早得多掐死死循环(第 8 步 vs 第 120 步),同时放正经长任务自由跑。 # 保守取 8:几乎不误伤"连踩几个错再纠正"的正常波动,配 _RepeatGuard 逐指纹 HARD=4 双保险。 @@ -617,23 +623,44 @@ class AgentLoop: cancel_check=self._is_cancelled, ) cancelled = False + pending_events: List[dict] = [] + may_reroute = self.caps.family == "deepseek_v4" + output_route_known = not may_reroute try: for chunk in stream: if self._is_cancelled(): cancelled = True break chunks.append(chunk) + tool_names = extract_delta_tool_names(chunk) + if ( + may_reroute + and any(name in self._DEEPSEEK_NONSTREAM_TOOLS for name in tool_names) + ): + # 推理 delta 先暂存,避免切换后非流式完整 reasoning 再发一次造成重复。 + raise PreferNonstreamToolCall(next( + name for name in tool_names + if name in self._DEEPSEEK_NONSTREAM_TOOLS + )) # delta.content 即时 emit 给前端打字机渲染;tool_call delta 不实时发 # (拼接散在多 chunk 跨 frame 难看,等拼回后整条 tool_call 事件由 # _execute_tool_call 时机发更直观)。 delta_text = extract_delta_content(chunk) if delta_text: - self._emit({"type": "text", "delta": delta_text}) + pending_events.append({"type": "text", "delta": delta_text}) # thinking 模型的推理 delta 也实时流出(reasoning 事件):深度推理可达 # 分钟级,不发的话前端全程静止"思考中",用户以为卡死。 delta_reasoning = extract_delta_reasoning(chunk) if delta_reasoning: - self._emit({"type": "reasoning", "delta": delta_reasoning}) + pending_events.append({"type": "reasoning", "delta": delta_reasoning}) + # 首个正文或任意工具名已确定本轮不会再切换;释放之前暂存的推理片段, + # 后续 chunk 也继续即时释放。只有尚未看到输出类型时才短暂缓冲。 + if delta_text or tool_names: + output_route_known = True + if output_route_known and pending_events: + for event in pending_events: + self._emit(event) + pending_events.clear() # interruptible stream 会在无新 chunk 的等待期直接因 cancel 结束迭代; # 循环体没有机会执行上面的检查,故在正常耗尽处再判一次。 if self._is_cancelled(): @@ -647,6 +674,10 @@ class AgentLoop: if cancelled: return None, True + # 极少数 provider 只有 reasoning/空收尾、始终没有正文或工具名,不能吞掉已收内容。 + for event in pending_events: + self._emit(event) + # 用 litellm 官方 helper 拼回完整 response(包括 tool_calls 拼接 + usage)。 # messages 参数仅用于失败时回填 prompt token 估算,正常路径 stream_options.include_usage # 已让最后一个 chunk 带准确 usage。 diff --git a/tests/test_loop_malformed_retry.py b/tests/test_loop_malformed_retry.py index 9d03d00..960ac6f 100644 --- a/tests/test_loop_malformed_retry.py +++ b/tests/test_loop_malformed_retry.py @@ -9,7 +9,12 @@ from types import SimpleNamespace sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -from core.llm_transport import malformed_tool_calls as _malformed_tool_calls # noqa: E402 +from core.llm_transport import ( # noqa: E402 + PreferNonstreamToolCall, +) +from core.llm_transport import ( + malformed_tool_calls as _malformed_tool_calls, +) from core.loop import AgentLoop # noqa: E402 @@ -51,6 +56,57 @@ def _make_loop(stream_results, nonstream_results): class TestMalformedRetry(unittest.TestCase): + def test_deepseek_write_name_aborts_before_buffered_reasoning_is_emitted(self): + """真实 collect 路径见到 write 名即关流,之前的 reasoning 不重复上屏。""" + reasoning = SimpleNamespace( + choices=[SimpleNamespace(delta=SimpleNamespace( + reasoning_content="thinking", content=None, tool_calls=None, + ))] + ) + write = SimpleNamespace( + choices=[SimpleNamespace(delta=SimpleNamespace( + reasoning_content=None, + content=None, + tool_calls=[SimpleNamespace( + function=SimpleNamespace(name="write", arguments=""), + )], + ))] + ) + loop = object.__new__(AgentLoop) + loop.caps = SimpleNamespace( + family="deepseek_v4", default_reasoning_effort=None, + ) + loop.executor = SimpleNamespace(schemas=lambda: []) + loop.cancel_check = None + loop.events = [] + loop._emit = loop.events.append + loop.llm = SimpleNamespace( + chat_stream=lambda **_kwargs: iter([reasoning, write]), + ) + + with self.assertRaises(PreferNonstreamToolCall): + loop._collect_stream_once([]) + self.assertEqual(loop.events, []) + + def test_preferred_tool_reroutes_before_arguments_stream(self): + """首包识别 write/edit 后直接非流式,不把它算作一次畸形失败。""" + loop, calls = _make_loop([], [_resp(GOOD)]) + + def prefer_nonstream(_messages): + calls["stream"] += 1 + raise PreferNonstreamToolCall("write") + + loop._collect_stream_once = prefer_nonstream + resp, cancelled = loop._stream_llm() + self.assertFalse(cancelled) + self.assertEqual(_malformed_tool_calls(resp), []) + self.assertEqual(calls, {"stream": 1, "nonstream": 1}) + infos = [e for e in loop.events if e.get("level") == "info"] + self.assertEqual(len(infos), 1) + self.assertIn("write", infos[0]["msg"]) + starts = [e for e in loop.events if e.get("type") == "llm_start"] + self.assertEqual(len(starts), 2) + def test_clean_stream_no_retry(self): loop, calls = _make_loop([_resp(GOOD)], []) resp, cancelled = loop._stream_llm()