diff --git a/src/converter/anti_truncation.py b/src/converter/anti_truncation.py index 26702811b..fcc421b41 100644 --- a/src/converter/anti_truncation.py +++ b/src/converter/anti_truncation.py @@ -56,7 +56,6 @@ 现在请继续输出:""" - # ==================== 请求注入 ==================== @@ -126,8 +125,6 @@ def apply_anti_truncation(payload: Dict[str, Any]) -> Dict[str, Any]: return modified_payload - - # ==================== 响应提取 ==================== @@ -221,7 +218,6 @@ def _extract_content_from_json_str(args_str: str) -> Optional[str]: return None - def build_text_chunk_from_synthetic( original_data: Dict[str, Any], synthetic_content: str, @@ -288,8 +284,6 @@ def build_text_chunk_from_synthetic( return modified_inner - - # ==================== 流式处理器 ==================== @@ -354,7 +348,8 @@ async def process_stream(self) -> AsyncGenerator[bytes, None]: # 每轮流内状态:初始化放在 try 之外,流中断的异常路径需要访问 found_synthetic = False has_real_tool_calls = False - side_buffer = io.StringIO() # 暂存普通文本(防拼接) + emitted_text = False + side_buffer = io.StringIO() # 保存已透传正文,供异常中断后续传 last_finish_reason: Optional[str] = None try: @@ -444,14 +439,14 @@ async def process_stream(self) -> AsyncGenerator[bytes, None]: synthetic_content, real_calls, chunk_has_synthetic = ( extract_synthetic_content_from_response(data) ) + has_real_tool_calls = has_real_tool_calls or bool(real_calls) if chunk_has_synthetic: found_synthetic = True - # 防拼接:丢弃之前暂存的普通文本 - if side_buffer.getvalue(): + if emitted_text: print( - "Anti-truncation: Discarding side-buffered text " - "(content conflict with synthetic tool)", + "Anti-truncation: Plain text was streamed before " + "synthetic tool call (possible duplication)", flush=True, ) side_buffer.close() @@ -471,7 +466,6 @@ async def process_stream(self) -> AsyncGenerator[bytes, None]: elif real_calls: # 真实工具调用,原样透传 - has_real_tool_calls = True yield line else: @@ -487,16 +481,18 @@ async def process_stream(self) -> AsyncGenerator[bytes, None]: stripped, separators=(",", ":"), ensure_ascii=False ) yield f"data: {json_str}\n\n".encode("utf-8") + elif not self._chunk_has_plain_text(data): + yield line continue else: - # 暂存到 side buffer,等待看是否有合成工具调用 - text = self._extract_text_from_chunk(data) - if text: - side_buffer.write(text) + # 正文与思考实时透传;正文另存一份用于异常续传历史。 + if self._chunk_has_plain_text(data): + emitted_text = True + side_buffer.write(self._extract_text_from_chunk(data)) chunk_finish = self._get_finish_reason(data) if chunk_finish: last_finish_reason = chunk_finish - # 暂时不透传,等流结束时决定 + yield line continue else: @@ -504,7 +500,6 @@ async def process_stream(self) -> AsyncGenerator[bytes, None]: yield line # 流结束(break 或正常结束) - side_text = side_buffer.getvalue() side_buffer.close() if found_synthetic: @@ -514,41 +509,17 @@ async def process_stream(self) -> AsyncGenerator[bytes, None]: yield b"data: [DONE]\n\n" return - # 模型自然结束(STOP)且未调用合成工具:视为完整回答,不再续传 - # (续传会让模型重发一遍已有内容,浪费上游请求且可能产生重复文本) - if last_finish_reason == "STOP": + # 真实工具或正文已交给客户端,不能续传重放;保留上游 STOP 终止语义。 + if has_real_tool_calls or emitted_text or last_finish_reason == "STOP": print( - f"Anti-truncation: Stream ended with STOP without synthetic tool call " - f"(text length: {len(side_text)}), treating as complete", + "Anti-truncation: Real output or STOP received, treating turn as complete", flush=True, ) - if side_text: - # 输出暂存文本(正常情况下模型守规矩时不会走到这里) - self._append_content(side_text) - fallback_chunk = self._build_fallback_text_chunk(side_text) - if fallback_chunk: - yield fallback_chunk self._clear_content() - # 补发 finishReason 收尾 chunk(Gemini 客户端靠它判断流正常结束) - yield self._build_finish_reason_chunk() yield b"data: [DONE]\n\n" return - # 未收到合成工具调用 - if side_text: - # 有普通文本作为 fallback,输出它 - print( - f"Anti-truncation: No synthetic tool call, " - f"using side-buffered text as fallback (length: {len(side_text)})", - flush=True, - ) - self._append_content(side_text) - # 构建一个包含 side buffer 文本的 chunk 输出 - fallback_chunk = self._build_fallback_text_chunk(side_text) - if fallback_chunk: - yield fallback_chunk - - # 触发续传 + # 正常结束但没有正文或工具产出(如纯思考)才触发续传。 if self.current_attempt < self.max_attempts: accumulated_text = self._get_collected_text() total_length = len(accumulated_text) @@ -580,20 +551,12 @@ async def process_stream(self) -> AsyncGenerator[bytes, None]: ) self._append_content(interrupted_text) - if self.current_attempt >= self.max_attempts: - # 重试额度用尽:把已收到的内容(含残文)作为 fallback 输出,尽量不浪费 + if has_real_tool_calls or self.current_attempt >= self.max_attempts: + # 工具已经发出时必须报错终止,续传可能导致客户端重复执行。 + # 正文已实时发送;重试额度用尽时不能再次输出整个历史。 salvaged = self._get_collected_text() self._clear_content() - if salvaged: - print( - f"Anti-truncation: Max attempts reached after error, " - f"yielding salvaged text (length: {len(salvaged)})", - flush=True, - ) - fallback_chunk = self._build_fallback_text_chunk(salvaged) - if fallback_chunk: - yield fallback_chunk - else: + if has_real_tool_calls or not salvaged: error_chunk = { "error": { "message": f"Anti-truncation failed: {str(e)}", @@ -633,8 +596,8 @@ def _build_current_payload(self) -> Dict[str, Any]: if accumulated_text: new_contents.append({"role": "model", "parts": [{"text": accumulated_text}]}) - # 预填充模式:直接用拼接内容作为末尾 model 预填充 - if self.enable_prefill_mode: + # 没有正文可续写时使用明确的续传指令,避免空预填充原样重发。 + if self.enable_prefill_mode and accumulated_text: request_data["contents"] = new_contents continuation_payload["request"] = request_data return continuation_payload @@ -669,10 +632,21 @@ def _extract_text_from_chunk(self, data: Dict[str, Any]) -> str: for candidate in data.get("candidates", []): content = candidate.get("content", {}) for part in content.get("parts", []): - if isinstance(part, dict) and "text" in part: + if isinstance(part, dict) and "text" in part and not part.get("thought", False): text += part["text"] return text + @staticmethod + def _chunk_has_plain_text(data: Dict[str, Any]) -> bool: + """非空正文才算产出;思考与空文本收尾块不算。""" + if "response" in data: + data = data["response"] + return any( + isinstance(part, dict) and part.get("text") and not part.get("thought", False) + for candidate in data.get("candidates", []) + for part in candidate.get("content", {}).get("parts", []) + ) + @staticmethod def _has_finish_reason(data: Dict[str, Any]) -> bool: """判断 chunk 是否携带 finishReason(控制信号 chunk)。""" @@ -822,6 +796,11 @@ async def _handle_non_streaming_response(self, response) -> bytes: extract_synthetic_content_from_response(response_data) ) + if not found_synthetic and ( + real_calls or self._chunk_has_plain_text(response_data) + ): + return content.encode() if isinstance(content, str) else content + if found_synthetic or self.current_attempt >= self.max_attempts: if found_synthetic: # 替换响应中的合成工具调用为普通文本 @@ -856,5 +835,3 @@ async def _handle_non_streaming_response(self, response) -> bytes: } } ).encode() - - diff --git a/tests/test_anti_truncation.py b/tests/test_anti_truncation.py new file mode 100644 index 000000000..4d7be14bf --- /dev/null +++ b/tests/test_anti_truncation.py @@ -0,0 +1,211 @@ +import asyncio +from copy import deepcopy +import json + +import pytest +from fastapi.responses import JSONResponse, StreamingResponse + +from src.converter.anti_truncation import AntiTruncationStreamProcessor + + +DONE = b"data: [DONE]\n\n" + + +def chunk(parts=(), finish=None, **metadata): + candidate = {"content": {"role": "model", "parts": list(parts)}} + if finish: + candidate["finishReason"] = finish + return f"data: {json.dumps({'candidates': [candidate], **metadata})}\n\n".encode() + + +def payload(): + return {"request": {"contents": [{"role": "user", "parts": [{"text": "hello"}]}]}} + + +def collect(attempts, **options): + requests = [] + + async def request(body): + requests.append(deepcopy(body)) + current = attempts[min(len(requests) - 1, len(attempts) - 1)] + if isinstance(current, JSONResponse): + return current + + async def generate(): + for item in current: + if isinstance(item, Exception): + raise item + yield item + + return StreamingResponse(generate()) + + processor = AntiTruncationStreamProcessor(request, payload(), **options) + + async def run(): + return [item async for item in processor.process_stream()] + + return asyncio.run(run()), requests + + +def texts(output): + result = [] + for item in output: + if item == DONE: + continue + data = json.loads(item.decode()[6:]) + for candidate in data.get("candidates", []): + result.extend( + part["text"] + for part in candidate.get("content", {}).get("parts", []) + if "text" in part and not part.get("thought") + ) + return "".join(result) + + +def test_plain_text_is_yielded_before_upstream_advances(): + first = chunk([{"text": "hello"}]) + + async def request(_): + async def generate(): + yield first + raise AssertionError("upstream advanced before caller received its first chunk") + + return StreamingResponse(generate()) + + async def run(): + stream = AntiTruncationStreamProcessor(request, payload()).process_stream() + try: + assert await anext(stream) == first + finally: + await stream.aclose() + + asyncio.run(run()) + + +@pytest.mark.parametrize( + "parts", [[{"text": "hello"}], [{"functionCall": {"name": "lookup", "args": {}}}]] +) +def test_real_output_without_synthetic_call_is_not_retried(parts): + first = chunk(parts) + output, requests = collect([[first, DONE]]) + assert len(requests) == 1 + assert output == [first, DONE] + + +def test_tool_finish_and_usage_are_preserved_once(): + first = chunk([{"functionCall": {"name": "lookup", "args": {}}}]) + finish = chunk([{"text": ""}], "STOP", usageMetadata={"totalTokenCount": 5}) + output, requests = collect([[first, finish, DONE]]) + assert len(requests) == 1 + assert output == [first, finish, DONE] + + +@pytest.mark.parametrize("with_synthetic", [False, True]) +def test_interrupted_tool_turn_is_not_replayed(with_synthetic): + parts = [{"functionCall": {"name": "lookup", "args": {}}}] + if with_synthetic: + parts.insert(0, {"functionCall": {"name": "emit_answer", "args": {"content": "answer"}}}) + tool_chunk = chunk(parts) + output, requests = collect( + [[tool_chunk, RuntimeError("offline interruption")], [tool_chunk, DONE]] + ) + assert len(requests) == 1 + data = [json.loads(item.decode()[6:]) for item in output if item != DONE] + calls = [ + part["functionCall"] + for item in data + for candidate in item.get("candidates", []) + for part in candidate.get("content", {}).get("parts", []) + if "functionCall" in part + ] + assert calls == [{"name": "lookup", "args": {}}] + assert data[-1]["error"]["code"] == 500 + assert "offline interruption" in data[-1]["error"]["message"] + assert output[-1] == DONE + + +@pytest.mark.parametrize("trailing_text", ["", "duplicate"]) +def test_synthetic_finish_strips_text_and_keeps_usage(trailing_text): + first = chunk([{"functionCall": {"name": "emit_answer", "args": {"content": "answer"}}}]) + finish = chunk([{"text": trailing_text}], "STOP", usageMetadata={"totalTokenCount": 5}) + output, requests = collect([[first, finish, DONE]]) + assert len(requests) == 1 + assert texts(output) == "answer" + final = json.loads(output[-2].decode()[6:]) + assert final["candidates"][0]["finishReason"] == "STOP" + assert final["usageMetadata"]["totalTokenCount"] == 5 + + +def test_separate_usage_chunk_survives_synthetic_answer(): + first = chunk([{"functionCall": {"name": "emit_answer", "args": {"content": "answer"}}}]) + usage = b'data: {"usageMetadata":{"totalTokenCount":5}}\n\n' + output, _ = collect([[first, usage, DONE]]) + assert usage in output + + +def test_thought_only_continuation_uses_prompt_without_empty_prefill(): + thought = chunk([{"text": "thinking", "thought": True}]) + output, requests = collect( + [[thought, DONE], [chunk([{"text": "answer"}]), DONE]], enable_prefill_mode=True + ) + assert len(requests) == 2 + assert thought in output + assert texts(output) == "answer" + assert requests[1]["request"]["contents"][-1]["role"] == "user" + assert len(requests[1]["request"]["contents"]) == 2 + + +def test_interrupted_text_enters_continuation_history_without_replay(): + output, requests = collect( + [ + [chunk([{"text": "partial"}]), RuntimeError("offline interruption")], + [chunk([{"text": " remainder"}]), DONE], + ] + ) + assert len(requests) == 2 + assert texts(output) == "partial remainder" + assert requests[1]["request"]["contents"][-2] == { + "role": "model", + "parts": [{"text": "partial"}], + } + + +def test_exhausted_interruption_does_not_repeat_already_streamed_text(): + output, requests = collect( + [[chunk([{"text": "partial"}]), RuntimeError("offline interruption")]], max_attempts=1 + ) + assert len(requests) == 1 + assert texts(output) == "partial" + assert output[-1] == DONE + + +@pytest.mark.parametrize( + "parts", [[{"text": "answer"}], [{"functionCall": {"name": "lookup", "args": {}}}]] +) +def test_non_streaming_real_output_is_not_retried(parts): + response = JSONResponse({"candidates": [{"content": {"parts": parts}}]}) + output, requests = collect([response]) + assert len(requests) == 1 + assert json.loads(output[0].decode()[6:])["candidates"][0]["content"]["parts"] == parts + + +def test_upstream_error_is_not_retried(): + response = JSONResponse({"error": {"message": "offline failure"}}, status_code=429) + output, requests = collect([response]) + assert len(requests) == 1 + assert json.loads(output[0].decode()[6:])["error"]["message"] == "offline failure" + + +def test_continuation_does_not_mutate_or_accumulate_original_history(): + original = payload() + saved = deepcopy(original) + processor = AntiTruncationStreamProcessor(None, original) + processor._append_content("partial") + processor.current_attempt = 2 + second = processor._build_current_payload() + processor.current_attempt = 3 + third = processor._build_current_payload() + assert second == third + assert len(second["request"]["contents"]) == 3 + assert original == saved + assert processor.base_payload == saved