From 60b841bbfb633a00d4ce17596ea83ce679b70bf1 Mon Sep 17 00:00:00 2001 From: luyifan Date: Tue, 30 Jun 2026 04:51:32 +0800 Subject: [PATCH] fix(tts): bound upstream response bodies --- tests/tools/test_tts_response_body_cap.py | 85 ++++++++++++++ tests/tools/test_tts_xai_speech_tags.py | 30 +++-- tools/tts_tool.py | 128 +++++++++++++++++++--- 3 files changed, 220 insertions(+), 23 deletions(-) create mode 100644 tests/tools/test_tts_response_body_cap.py diff --git a/tests/tools/test_tts_response_body_cap.py b/tests/tools/test_tts_response_body_cap.py new file mode 100644 index 000000000000..973bec9e537e --- /dev/null +++ b/tests/tools/test_tts_response_body_cap.py @@ -0,0 +1,85 @@ +"""Regression tests for bounded upstream TTS response reads.""" + +from __future__ import annotations + +from unittest.mock import patch + +import pytest + +from tools import tts_tool + + +class StreamingResponse: + def __init__(self, chunks, *, status_code=200, headers=None): + self._chunks = list(chunks) + self.status_code = status_code + self.headers = headers or {} + self.closed = False + + def iter_content(self, chunk_size=65536): + del chunk_size + yield from self._chunks + + def close(self): + self.closed = True + + def raise_for_status(self): + if self.status_code >= 400: + raise RuntimeError(f"HTTP {self.status_code}") + + +@pytest.fixture(autouse=True) +def small_tts_body_cap(monkeypatch): + monkeypatch.setattr(tts_tool, "TTS_RESPONSE_BODY_LIMIT_BYTES", 8) + + +def test_xai_tts_rejects_oversized_audio_response(tmp_path, monkeypatch): + monkeypatch.setenv("XAI_API_KEY", "test-xai-key") + response = StreamingResponse([b"12345", b"6789"], headers={"Content-Type": "audio/mpeg"}) + output_path = tmp_path / "out.mp3" + + with patch("requests.post", return_value=response) as post: + with pytest.raises(RuntimeError, match="xAI TTS response exceeds 8 bytes"): + tts_tool._generate_xai_tts("hello", str(output_path), {}) + + assert post.call_args.kwargs["stream"] is True + assert response.closed is True + assert not output_path.exists() + + +def test_minimax_t2a_rejects_oversized_json_response(tmp_path, monkeypatch): + monkeypatch.setenv("MINIMAX_API_KEY", "test-minimax-key") + response = StreamingResponse([b'{"data":', b'"too large"}'], headers={"Content-Type": "application/json"}) + + with patch("requests.post", return_value=response) as post: + with pytest.raises(RuntimeError, match="MiniMax TTS response exceeds 8 bytes"): + tts_tool._generate_minimax_tts("hello", str(tmp_path / "out.mp3"), {}) + + assert post.call_args.kwargs["stream"] is True + assert response.closed is True + + +def test_minimax_legacy_rejects_oversized_audio_response(tmp_path, monkeypatch): + monkeypatch.setenv("MINIMAX_API_KEY", "test-minimax-key") + response = StreamingResponse([b"12345", b"6789"], headers={"Content-Type": "audio/mpeg"}) + config = {"minimax": {"base_url": "https://api.minimax.chat/v1/text_to_speech"}} + output_path = tmp_path / "out.mp3" + + with patch("requests.post", return_value=response): + with pytest.raises(RuntimeError, match="MiniMax TTS response exceeds 8 bytes"): + tts_tool._generate_minimax_tts("hello", str(output_path), config) + + assert response.closed is True + assert not output_path.exists() + + +def test_gemini_tts_rejects_oversized_json_response(tmp_path, monkeypatch): + monkeypatch.setenv("GEMINI_API_KEY", "test-gemini-key") + response = StreamingResponse([b'{"candidates":', b"[{}]}"], headers={"Content-Type": "application/json"}) + + with patch("requests.post", return_value=response) as post: + with pytest.raises(RuntimeError, match="Gemini TTS response exceeds 8 bytes"): + tts_tool._generate_gemini_tts("hello", str(tmp_path / "out.wav"), {}) + + assert post.call_args.kwargs["stream"] is True + assert response.closed is True diff --git a/tests/tools/test_tts_xai_speech_tags.py b/tests/tools/test_tts_xai_speech_tags.py index bae1bbe91e61..98efe8eac9af 100644 --- a/tests/tools/test_tts_xai_speech_tags.py +++ b/tests/tools/test_tts_xai_speech_tags.py @@ -101,11 +101,12 @@ def test_generate_xai_tts_sends_auxiliary_rewriter_output_to_api( def raise_for_status(self): pass - def fake_post(url, headers, json, timeout): + def fake_post(url, headers, json, timeout, stream=False): captured["url"] = url captured["headers"] = headers captured["json"] = json captured["timeout"] = timeout + captured["stream"] = stream return FakeResponse() fake_response = SimpleNamespace( @@ -352,8 +353,9 @@ def test_generate_xai_tts_leaves_text_plain_by_default(tmp_path, monkeypatch): fake_response.content = b"mp3" fake_response.raise_for_status.return_value = None - def fake_post(url, headers, json, timeout): + def fake_post(url, headers, json, timeout, stream=False): captured["json"] = json + captured["stream"] = stream return fake_response monkeypatch.setenv("XAI_API_KEY", "test-xai-key") @@ -379,8 +381,9 @@ def test_generate_xai_tts_omits_speed_and_latency_by_default(tmp_path, monkeypat fake_response.content = b"mp3" fake_response.raise_for_status.return_value = None - def fake_post(url, headers, json, timeout): + def fake_post(url, headers, json, timeout, stream=False): captured["json"] = json + captured["stream"] = stream return fake_response monkeypatch.setenv("XAI_API_KEY", "test-xai-key") @@ -404,8 +407,9 @@ def test_generate_xai_tts_sends_speed_when_set(tmp_path, monkeypatch): fake_response.content = b"mp3" fake_response.raise_for_status.return_value = None - def fake_post(url, headers, json, timeout): + def fake_post(url, headers, json, timeout, stream=False): captured["json"] = json + captured["stream"] = stream return fake_response monkeypatch.setenv("XAI_API_KEY", "test-xai-key") @@ -428,8 +432,9 @@ def test_generate_xai_tts_speed_clamped_to_valid_range(tmp_path, monkeypatch): fake_response.content = b"mp3" fake_response.raise_for_status.return_value = None - def fake_post(url, headers, json, timeout): + def fake_post(url, headers, json, timeout, stream=False): captured["json"] = json + captured["stream"] = stream return fake_response monkeypatch.setenv("XAI_API_KEY", "test-xai-key") @@ -460,8 +465,9 @@ def test_generate_xai_tts_omits_speed_when_exactly_default(tmp_path, monkeypatch fake_response.content = b"mp3" fake_response.raise_for_status.return_value = None - def fake_post(url, headers, json, timeout): + def fake_post(url, headers, json, timeout, stream=False): captured["json"] = json + captured["stream"] = stream return fake_response monkeypatch.setenv("XAI_API_KEY", "test-xai-key") @@ -484,8 +490,9 @@ def test_generate_xai_tts_sends_optimize_streaming_latency_when_set(tmp_path, mo fake_response.content = b"mp3" fake_response.raise_for_status.return_value = None - def fake_post(url, headers, json, timeout): + def fake_post(url, headers, json, timeout, stream=False): captured["json"] = json + captured["stream"] = stream return fake_response monkeypatch.setenv("XAI_API_KEY", "test-xai-key") @@ -508,8 +515,9 @@ def test_generate_xai_tts_optimize_streaming_latency_omitted_at_default(tmp_path fake_response.content = b"mp3" fake_response.raise_for_status.return_value = None - def fake_post(url, headers, json, timeout): + def fake_post(url, headers, json, timeout, stream=False): captured["json"] = json + captured["stream"] = stream return fake_response monkeypatch.setenv("XAI_API_KEY", "test-xai-key") @@ -532,8 +540,9 @@ def test_generate_xai_tts_global_speed_used_as_fallback(tmp_path, monkeypatch): fake_response.content = b"mp3" fake_response.raise_for_status.return_value = None - def fake_post(url, headers, json, timeout): + def fake_post(url, headers, json, timeout, stream=False): captured["json"] = json + captured["stream"] = stream return fake_response monkeypatch.setenv("XAI_API_KEY", "test-xai-key") @@ -556,8 +565,9 @@ def test_generate_xai_tts_provider_speed_overrides_global(tmp_path, monkeypatch) fake_response.content = b"mp3" fake_response.raise_for_status.return_value = None - def fake_post(url, headers, json, timeout): + def fake_post(url, headers, json, timeout, stream=False): captured["json"] = json + captured["stream"] = stream return fake_response monkeypatch.setenv("XAI_API_KEY", "test-xai-key") diff --git a/tools/tts_tool.py b/tools/tts_tool.py index dadf5ef13a40..42e446cf9239 100644 --- a/tools/tts_tool.py +++ b/tools/tts_tool.py @@ -246,6 +246,8 @@ DEFAULT_DEEPINFRA_TTS_VOICE = "default" GEMINI_TTS_SAMPLE_RATE = 24000 GEMINI_TTS_CHANNELS = 1 GEMINI_TTS_SAMPLE_WIDTH = 2 # 16-bit PCM (L16) +TTS_RESPONSE_BODY_LIMIT_BYTES = 16 * 1024 * 1024 +TTS_RESPONSE_BODY_CHUNK_BYTES = 64 * 1024 def _get_default_output_dir() -> str: from hermes_constants import get_hermes_dir @@ -302,6 +304,93 @@ def _config_bool(value: Any, default: bool = False) -> bool: return False return default + +def _response_has_explicit_stream(response: Any) -> bool: + iter_content = getattr(response, "iter_content", None) + if not callable(iter_content): + return False + response_type = type(response) + if response_type.__module__.startswith("requests."): + return True + return "iter_content" in vars(response_type) + + +def _close_response(response: Any) -> None: + close = getattr(response, "close", None) + if callable(close): + try: + close() + except Exception: + pass + + +def _read_tts_response_bytes( + response: Any, + *, + label: str, + limit: Optional[int] = None, +) -> bytes: + """Read an upstream TTS response with a hard byte cap.""" + limit = TTS_RESPONSE_BODY_LIMIT_BYTES if limit is None else limit + chunks: list[bytes] = [] + total = 0 + try: + if _response_has_explicit_stream(response): + iterator = response.iter_content(chunk_size=TTS_RESPONSE_BODY_CHUNK_BYTES) + else: + content = vars(response).get("content", getattr(type(response), "content", b"")) + if isinstance(content, str): + content = content.encode("utf-8", errors="replace") + iterator = (content,) if isinstance(content, (bytes, bytearray)) else () + + for chunk in iterator: + if not chunk: + continue + if isinstance(chunk, str): + chunk = chunk.encode("utf-8", errors="replace") + chunk = bytes(chunk) + total += len(chunk) + if total > limit: + _close_response(response) + raise RuntimeError(f"{label} response exceeds {limit} bytes") + chunks.append(chunk) + return b"".join(chunks) + finally: + _close_response(response) + + +def _read_tts_response_json( + response: Any, + *, + label: str, + limit: Optional[int] = None, +) -> Dict[str, Any]: + raw = _read_tts_response_bytes(response, label=label, limit=limit) + if raw: + return json.loads(raw.decode("utf-8")) + + # Unit-test doubles often only provide `.json()`. Real requests.Response + # objects use the streaming path above, so this fallback does not re-open + # the production eager-buffering behavior. + if not _response_has_explicit_stream(response): + json_reader = getattr(response, "json", None) + if callable(json_reader): + parsed = json_reader() + return parsed if isinstance(parsed, dict) else {} + return {} + + +def _write_tts_response_to_file( + response: Any, + output_path: str, + *, + label: str, + limit: Optional[int] = None, +) -> None: + audio_bytes = _read_tts_response_bytes(response, label=label, limit=limit) + with open(output_path, "wb") as f: + f.write(audio_bytes) + # Final fallback when provider isn't recognised at all. FALLBACK_MAX_TEXT_LENGTH = 4000 @@ -1550,11 +1639,11 @@ def _generate_xai_tts(text: str, output_path: str, tts_config: Dict[str, Any]) - }, json=payload, timeout=60, + stream=True, ) response.raise_for_status() - with open(output_path, "wb") as f: - f.write(response.content) + _write_tts_response_to_file(response, output_path, label="xAI TTS") return output_path @@ -1642,12 +1731,18 @@ def _generate_minimax_tts(text: str, output_path: str, tts_config: Dict[str, Any "voice_id": voice_id, } - response = requests.post(base_url, json=payload, headers=headers, timeout=60) + response = requests.post( + base_url, + json=payload, + headers=headers, + timeout=60, + stream=True, + ) if is_t2a_v2: # t2a_v2 returns JSON with hex-encoded audio response.raise_for_status() - result = response.json() + result = _read_tts_response_json(response, label="MiniMax TTS") base_resp = result.get("base_resp", {}) status_code = base_resp.get("status_code", -1) @@ -1669,23 +1764,23 @@ def _generate_minimax_tts(text: str, output_path: str, tts_config: Dict[str, Any content_type = response.headers.get("Content-Type", "") if "audio/" in content_type: - with open(output_path, "wb") as f: - f.write(response.content) + _write_tts_response_to_file(response, output_path, label="MiniMax TTS") return output_path # Fallback: try parsing as JSON try: - result = response.json() + raw_body = _read_tts_response_bytes(response, label="MiniMax TTS") + result = json.loads(raw_body.decode("utf-8")) if raw_body else {} base_resp = result.get("base_resp", {}) status_code = base_resp.get("status_code", -1) if status_code != 0: status_msg = base_resp.get("status_msg", "unknown error") raise RuntimeError(f"MiniMax TTS API error (code {status_code}): {status_msg}") - except Exception: + except (json.JSONDecodeError, UnicodeDecodeError, TypeError): response.raise_for_status() raise RuntimeError( f"MiniMax TTS returned unexpected Content-Type '{content_type}' " - f"({len(response.content)} bytes)" + f"({len(raw_body) if 'raw_body' in locals() else 0} bytes)" ) raise RuntimeError("MiniMax TTS returned no audio data") @@ -2023,20 +2118,27 @@ def _generate_gemini_tts(text: str, output_path: str, tts_config: Dict[str, Any] headers=headers, json=payload, timeout=60, + stream=True, ) if response.status_code != 200: # Surface the API error message when present + raw_body = _read_tts_response_bytes(response, label="Gemini TTS") try: - err = response.json().get("error", {}) - detail = err.get("message") or response.text[:300] + if raw_body: + err = json.loads(raw_body.decode("utf-8")).get("error", {}) + elif not _response_has_explicit_stream(response) and callable(getattr(response, "json", None)): + err = response.json().get("error", {}) + else: + err = {} + detail = err.get("message") or raw_body.decode("utf-8", errors="replace")[:300] except Exception: - detail = response.text[:300] + detail = raw_body.decode("utf-8", errors="replace")[:300] raise RuntimeError( f"Gemini TTS API error (HTTP {response.status_code}): {detail}" ) try: - data = response.json() + data = _read_tts_response_json(response, label="Gemini TTS") parts = data["candidates"][0]["content"]["parts"] audio_part = next((p for p in parts if "inlineData" in p or "inline_data" in p), None) if audio_part is None: