diff --git a/services/providers/deepseek_provider.py b/services/providers/deepseek_provider.py index 1f13554..5d9e118 100644 --- a/services/providers/deepseek_provider.py +++ b/services/providers/deepseek_provider.py @@ -44,6 +44,24 @@ def _get_language_name(code: str) -> str: return language_name(code) + +def _build_system_prompt( + source_lang: str, target_lang: str, custom_prompt: Optional[str] = None +) -> str: + """Build system prompt for translation. + + The base translation instructions are ALWAYS present — a custom prompt + (glossary, tone, context) is appended as additional directives, never a + replacement. Same contract as the OpenAI provider. + """ + base = DEFAULT_TRANSLATION_PROMPT.format( + source_lang=source_lang, target_lang=target_lang + ) + if custom_prompt and custom_prompt.strip(): + return f"{base}\n\nADDITIONAL CONTEXT AND INSTRUCTIONS:\n{custom_prompt.strip()}" + return base + + class DeepSeekProviderError(Exception): def __init__(self, code: str, message: str, details: Optional[Dict[str, Any]] = None): self.code = code @@ -155,19 +173,11 @@ class DeepSeekTranslationProvider(TranslationProvider): source_lang_name = _get_language_name(source_language) target_lang_name = _get_language_name(target_language) custom_prompt = request.metadata.get("custom_prompt") if request.metadata else None - _base_prompt = DEFAULT_TRANSLATION_PROMPT.format( - source_lang=source_lang_name, target_lang=target_lang_name - ) # Base translation instructions always present; the custom prompt # (glossary/tone/context) is appended, never a replacement. - if custom_prompt and custom_prompt.strip(): - system_prompt = ( - _base_prompt - + "\n\nADDITIONAL CONTEXT AND INSTRUCTIONS:\n" - + custom_prompt.strip() - ) - else: - system_prompt = _base_prompt + system_prompt = _build_system_prompt( + source_lang_name, target_lang_name, custom_prompt + ) last_error = None for attempt in range(self._max_retries + 1): @@ -196,7 +206,120 @@ class DeepSeekTranslationProvider(TranslationProvider): error=last_error.message if last_error else "Unknown error", error_code=last_error.code if last_error else DEEPSEEK_SERVICE_ERROR) - def translate_batch(self, requests: List[TranslationRequest]) -> List[TranslationResponse]: + def _make_batch_api_request( + self, texts: List[str], system_prompt: str + ) -> Optional[List[str]]: + """Translate a whole chunk in ONE request via a numbered JSON list. + + The user message is a JSON array; the model must answer with a JSON + array of the same length. Returns None when the answer cannot be + parsed confidently — callers then fall back to per-item calls + (correctness over latency). Mirrors the OpenAI provider batch mode. + """ + import json as _json + + numbered = _json.dumps( + [{"id": i, "text": t} for i, t in enumerate(texts)], + ensure_ascii=False, + ) + batch_system = ( + system_prompt + + "\n\nBATCH MODE: the user message is a JSON array of items with " + "unique ids. Answer with ONLY a JSON array of objects " + '[{"id": , "translation": ""}], same ' + "length and same ids, in the same order. Translate every item; " + "keep ids unchanged; no comments, no markdown fence." + ) + + try: + content, _usage = self._make_api_request(numbered, batch_system) + raw = content.strip() + # Strip an optional markdown fence + if raw.startswith("```"): + raw = raw.strip("`") + if raw.lower().startswith("json"): + raw = raw[4:] + raw = raw.strip() + parsed = _json.loads(raw) + if not isinstance(parsed, list) or len(parsed) != len(texts): + return None + out: List[str] = [""] * len(texts) + for item in parsed: + if not isinstance(item, dict): + return None + idx = item.get("id") + translation = item.get("translation") + if not isinstance(idx, int) or not 0 <= idx < len(texts): + return None + if not isinstance(translation, str) or not translation.strip(): + return None + out[idx] = translation.strip() + return out + except DeepSeekProviderError: + raise + except Exception: + return None + + def translate_batch( + self, requests: List[TranslationRequest] + ) -> List[TranslationResponse]: + """ + Translate multiple texts. + + Chunks arrive from the translators as ~15 texts. When every request + shares the same language pair and metadata, they are sent in ONE + call (numbered JSON list — ~15x fewer requests, better contextual + consistency across neighbouring segments). Any parse/API doubt falls + back to the per-item path so a batch failure never corrupts output. + """ + if not requests: + return [] + + same_pair = len({(r.source_language, r.target_language) for r in requests}) == 1 + same_meta = len( + {tuple(sorted((r.metadata or {}).items())) for r in requests} + ) == 1 + + if same_pair and same_meta and len(requests) > 1: + try: + source_lang_name = _get_language_name( + requests[0].source_language or "auto" + ) or "the source language (auto-detect)" + target_lang_name = _get_language_name(requests[0].target_language) + custom_prompt = None + if requests[0].metadata: + custom_prompt = requests[0].metadata.get("custom_prompt") + system_prompt = _build_system_prompt( + source_lang_name, target_lang_name, custom_prompt + ) + texts = [r.text for r in requests] + translations = self._make_batch_api_request(texts, system_prompt) + if translations is not None: + logger.info( + "deepseek_batch_translation_success", + items=len(requests), + model=self._model, + ) + return [ + TranslationResponse( + translated_text=t, + provider_name=self._provider_name, + from_cache=False, + ) + for t in translations + ] + logger.warning( + "deepseek_batch_translation_fallback", + reason="unparseable_response", + items=len(requests), + ) + except Exception as e: + logger.warning( + "deepseek_batch_translation_fallback", + reason=type(e).__name__, + items=len(requests), + ) + return [self.translate_text(req) for req in requests] def health_check(self) -> ProviderHealthStatus: diff --git a/services/providers/minimax_provider.py b/services/providers/minimax_provider.py index 7e36e98..a548ed1 100644 --- a/services/providers/minimax_provider.py +++ b/services/providers/minimax_provider.py @@ -46,6 +46,24 @@ def _get_language_name(code: str) -> str: return language_name(code) + +def _build_system_prompt( + source_lang: str, target_lang: str, custom_prompt: Optional[str] = None +) -> str: + """Build system prompt for translation. + + The base translation instructions are ALWAYS present — a custom prompt + (glossary, tone, context) is appended as additional directives, never a + replacement. Same contract as the OpenAI provider. + """ + base = DEFAULT_TRANSLATION_PROMPT.format( + source_lang=source_lang, target_lang=target_lang + ) + if custom_prompt and custom_prompt.strip(): + return f"{base}\n\nADDITIONAL CONTEXT AND INSTRUCTIONS:\n{custom_prompt.strip()}" + return base + + class MinimaxProviderError(Exception): def __init__(self, code: str, message: str, details: Optional[Dict[str, Any]] = None): self.code = code @@ -172,19 +190,11 @@ class MinimaxTranslationProvider(TranslationProvider): source_lang_name = _get_language_name(source_language) target_lang_name = _get_language_name(target_language) custom_prompt = request.metadata.get("custom_prompt") if request.metadata else None - _base_prompt = DEFAULT_TRANSLATION_PROMPT.format( - source_lang=source_lang_name, target_lang=target_lang_name - ) # Base translation instructions always present; the custom prompt # (glossary/tone/context) is appended, never a replacement. - if custom_prompt and custom_prompt.strip(): - system_prompt = ( - _base_prompt - + "\n\nADDITIONAL CONTEXT AND INSTRUCTIONS:\n" - + custom_prompt.strip() - ) - else: - system_prompt = _base_prompt + system_prompt = _build_system_prompt( + source_lang_name, target_lang_name, custom_prompt + ) last_error = None for attempt in range(self._max_retries + 1): @@ -213,7 +223,120 @@ class MinimaxTranslationProvider(TranslationProvider): error=last_error.message if last_error else "Unknown error", error_code=last_error.code if last_error else MINIMAX_SERVICE_ERROR) - def translate_batch(self, requests: List[TranslationRequest]) -> List[TranslationResponse]: + def _make_batch_api_request( + self, texts: List[str], system_prompt: str + ) -> Optional[List[str]]: + """Translate a whole chunk in ONE request via a numbered JSON list. + + The user message is a JSON array; the model must answer with a JSON + array of the same length. Returns None when the answer cannot be + parsed confidently — callers then fall back to per-item calls + (correctness over latency). Mirrors the OpenAI provider batch mode. + """ + import json as _json + + numbered = _json.dumps( + [{"id": i, "text": t} for i, t in enumerate(texts)], + ensure_ascii=False, + ) + batch_system = ( + system_prompt + + "\n\nBATCH MODE: the user message is a JSON array of items with " + "unique ids. Answer with ONLY a JSON array of objects " + '[{"id": , "translation": ""}], same ' + "length and same ids, in the same order. Translate every item; " + "keep ids unchanged; no comments, no markdown fence." + ) + + try: + content, _usage = self._make_api_request(numbered, batch_system) + raw = content.strip() + # Strip an optional markdown fence + if raw.startswith("```"): + raw = raw.strip("`") + if raw.lower().startswith("json"): + raw = raw[4:] + raw = raw.strip() + parsed = _json.loads(raw) + if not isinstance(parsed, list) or len(parsed) != len(texts): + return None + out: List[str] = [""] * len(texts) + for item in parsed: + if not isinstance(item, dict): + return None + idx = item.get("id") + translation = item.get("translation") + if not isinstance(idx, int) or not 0 <= idx < len(texts): + return None + if not isinstance(translation, str) or not translation.strip(): + return None + out[idx] = translation.strip() + return out + except MinimaxProviderError: + raise + except Exception: + return None + + def translate_batch( + self, requests: List[TranslationRequest] + ) -> List[TranslationResponse]: + """ + Translate multiple texts. + + Chunks arrive from the translators as ~15 texts. When every request + shares the same language pair and metadata, they are sent in ONE + call (numbered JSON list — ~15x fewer requests, better contextual + consistency across neighbouring segments). Any parse/API doubt falls + back to the per-item path so a batch failure never corrupts output. + """ + if not requests: + return [] + + same_pair = len({(r.source_language, r.target_language) for r in requests}) == 1 + same_meta = len( + {tuple(sorted((r.metadata or {}).items())) for r in requests} + ) == 1 + + if same_pair and same_meta and len(requests) > 1: + try: + source_lang_name = _get_language_name( + requests[0].source_language or "auto" + ) or "the source language (auto-detect)" + target_lang_name = _get_language_name(requests[0].target_language) + custom_prompt = None + if requests[0].metadata: + custom_prompt = requests[0].metadata.get("custom_prompt") + system_prompt = _build_system_prompt( + source_lang_name, target_lang_name, custom_prompt + ) + texts = [r.text for r in requests] + translations = self._make_batch_api_request(texts, system_prompt) + if translations is not None: + logger.info( + "minimax_batch_translation_success", + items=len(requests), + model=self._model, + ) + return [ + TranslationResponse( + translated_text=t, + provider_name=self._provider_name, + from_cache=False, + ) + for t in translations + ] + logger.warning( + "minimax_batch_translation_fallback", + reason="unparseable_response", + items=len(requests), + ) + except Exception as e: + logger.warning( + "minimax_batch_translation_fallback", + reason=type(e).__name__, + items=len(requests), + ) + return [self.translate_text(req) for req in requests] def health_check(self) -> ProviderHealthStatus: diff --git a/tests/test_providers/test_deepseek_provider.py b/tests/test_providers/test_deepseek_provider.py new file mode 100644 index 0000000..2d749bd --- /dev/null +++ b/tests/test_providers/test_deepseek_provider.py @@ -0,0 +1,312 @@ +""" +Tests for the DeepSeekTranslationProvider. + +Covers: +- defaults (public endpoint ``api.deepseek.com/v1``, model ``deepseek-chat``) +- translation success, 401/429 handling, empty-text short-circuit + +Batch mode (ported from the OpenAI provider): +- same language pair + metadata chunk → ONE numbered-JSON request for all texts +- unparseable/failed batch answer → per-item fallback, every text still translated +- mixed metadata or a single text → no batch (individual requests) +""" + +import json + +import pytest +from unittest.mock import patch, MagicMock + +from services.providers.deepseek_provider import ( + DeepSeekTranslationProvider, + DeepSeekProviderError, + DEEPSEEK_RATE_LIMITED, + DEEPSEEK_INVALID_KEY, + DEEPSEEK_TIMEOUT, + DEEPSEEK_SERVICE_ERROR, +) +from services.providers.schemas import TranslationRequest + + +class TestDeepSeekProviderConfig: + """Defaults must point at the real public endpoint.""" + + def test_default_base_url_is_public_host(self): + provider = DeepSeekTranslationProvider(api_key="k", max_retries=0) + assert provider._base_url == "https://api.deepseek.com/v1" + + def test_default_model_is_deepseek_chat(self): + provider = DeepSeekTranslationProvider(api_key="k", max_retries=0) + assert provider._model == "deepseek-chat" + + def test_custom_base_url_respected(self): + provider = DeepSeekTranslationProvider( + api_key="k", base_url="https://proxy.example.com/v1", max_retries=0 + ) + assert provider._base_url == "https://proxy.example.com/v1" + + def test_get_name(self): + provider = DeepSeekTranslationProvider(api_key="k", max_retries=0) + assert provider.get_name() == "deepseek" + + def test_empty_api_key_raises(self): + with pytest.raises(ValueError, match="API key cannot be empty"): + DeepSeekTranslationProvider(api_key=" ") + + +class TestDeepSeekTranslateText: + @pytest.fixture + def provider(self): + return DeepSeekTranslationProvider(api_key="k", max_retries=0) + + def _mock_post(self, payload, status_code=200): + mock_response = MagicMock() + mock_response.status_code = status_code + mock_response.json.return_value = payload + mock_response.text = "" + return mock_response + + @patch("requests.post") + def test_success(self, mock_post, provider): + mock_post.return_value = self._mock_post( + {"choices": [{"message": {"content": "Bonjour"}}], "usage": {}} + ) + resp = provider.translate_text(TranslationRequest(text="Hello", target_language="fr")) + assert resp.translated_text == "Bonjour" + assert resp.provider_name == "deepseek" + assert mock_post.call_args[0][0] == "https://api.deepseek.com/v1/chat/completions" + + def test_empty_text_short_circuits(self, provider): + resp = provider.translate_text(TranslationRequest(text="", target_language="fr")) + assert resp.translated_text == "" + + @patch("requests.post") + def test_invalid_key_returns_error(self, mock_post, provider): + mock_post.return_value = self._mock_post({"error": "bad key"}, status_code=401) + resp = provider.translate_text(TranslationRequest(text="Hello", target_language="fr")) + assert resp.error_code == DEEPSEEK_INVALID_KEY + # Original text returned on failure. + assert resp.translated_text == "Hello" + + @patch("time.sleep") + @patch("requests.post") + def test_rate_limit_then_success(self, mock_post, mock_sleep): + provider = DeepSeekTranslationProvider(api_key="k", max_retries=2, retry_delay=0.01) + mock_post.side_effect = [ + self._mock_post({"error": "slow down"}, status_code=429), + self._mock_post({"choices": [{"message": {"content": "Hola"}}], "usage": {}}), + ] + resp = provider.translate_text(TranslationRequest(text="Hello", target_language="es")) + assert resp.translated_text == "Hola" + assert mock_sleep.called # backoff happened + + @patch("requests.post") + def test_service_error_when_empty_choices(self, mock_post, provider): + mock_post.return_value = self._mock_post({"choices": []}) + resp = provider.translate_text(TranslationRequest(text="Hello", target_language="fr")) + assert resp.error_code == DEEPSEEK_SERVICE_ERROR + + +class TestDeepSeekBatchTranslation: + """Batch mode: one numbered-JSON request per same pair/metadata chunk.""" + + @pytest.fixture + def provider(self): + return DeepSeekTranslationProvider(api_key="k", max_retries=0) + + def _mock_post(self, payload, status_code=200): + mock_response = MagicMock() + mock_response.status_code = status_code + mock_response.json.return_value = payload + mock_response.text = "" + return mock_response + + def _batch_response(self, translations): + """Mock a 200 response whose content is a numbered JSON array.""" + content = json.dumps( + [{"id": i, "translation": t} for i, t in enumerate(translations)] + ) + return self._mock_post({"choices": [{"message": {"content": content}}], "usage": {}}) + + @patch("requests.post") + def test_batch_of_three_uses_single_request(self, mock_post, provider): + mock_post.return_value = self._batch_response(["Bonjour", "Au revoir", "Merci"]) + requests = [ + TranslationRequest(text="Hello", target_language="fr"), + TranslationRequest(text="Goodbye", target_language="fr"), + TranslationRequest(text="Thank you", target_language="fr"), + ] + + responses = provider.translate_batch(requests) + + assert mock_post.call_count == 1 + assert [r.translated_text for r in responses] == ["Bonjour", "Au revoir", "Merci"] + assert all(r.provider_name == "deepseek" for r in responses) + assert all(r.error is None for r in responses) + assert mock_post.call_args[0][0] == "https://api.deepseek.com/v1/chat/completions" + + @patch("requests.post") + def test_batch_request_payload_is_numbered_json_with_instructions(self, mock_post, provider): + mock_post.return_value = self._batch_response(["Bonjour", "Au revoir"]) + requests = [ + TranslationRequest( + text="Hello", + target_language="fr", + source_language="en", + metadata={"custom_prompt": "use a formal tone"}, + ), + TranslationRequest( + text="Goodbye", + target_language="fr", + source_language="en", + metadata={"custom_prompt": "use a formal tone"}, + ), + ] + + provider.translate_batch(requests) + + payload = mock_post.call_args[1]["json"] + system = payload["messages"][0]["content"] + user = payload["messages"][1]["content"] + # System prompt = base instructions + custom prompt + strict JSON rule. + assert "professional translator" in system + assert "ADDITIONAL CONTEXT AND INSTRUCTIONS" in system + assert "use a formal tone" in system + assert "BATCH MODE" in system + assert "ONLY a JSON array" in system + # User message = the numbered JSON list of the texts. + parsed = json.loads(user) + assert [item["text"] for item in parsed] == ["Hello", "Goodbye"] + assert [item["id"] for item in parsed] == [0, 1] + + @patch("requests.post") + def test_invalid_batch_json_falls_back_to_individual(self, mock_post, provider): + broken_batch = self._mock_post( + {"choices": [{"message": {"content": "désolé, pas de JSON ici"}}], "usage": {}} + ) + individuals = [ + self._mock_post({"choices": [{"message": {"content": "Bonjour"}}], "usage": {}}), + self._mock_post({"choices": [{"message": {"content": "Au revoir"}}], "usage": {}}), + self._mock_post({"choices": [{"message": {"content": "Merci"}}], "usage": {}}), + ] + mock_post.side_effect = [broken_batch] + individuals + requests = [ + TranslationRequest(text="Hello", target_language="fr"), + TranslationRequest(text="Goodbye", target_language="fr"), + TranslationRequest(text="Thank you", target_language="fr"), + ] + + responses = provider.translate_batch(requests) + + # 1 failed batch attempt + 3 individual calls. + assert mock_post.call_count == 4 + assert [r.translated_text for r in responses] == ["Bonjour", "Au revoir", "Merci"] + assert all(r.error is None for r in responses) + + @patch("requests.post") + def test_batch_api_error_falls_back_to_individual(self, mock_post, provider): + # The batch call itself fails (500) — per-item calls still translate everything. + failed_batch = self._mock_post({"error": "boom"}, status_code=500) + individuals = [ + self._mock_post({"choices": [{"message": {"content": "Bonjour"}}], "usage": {}}), + self._mock_post({"choices": [{"message": {"content": "Au revoir"}}], "usage": {}}), + ] + mock_post.side_effect = [failed_batch] + individuals + requests = [ + TranslationRequest(text="Hello", target_language="fr"), + TranslationRequest(text="Goodbye", target_language="fr"), + ] + + responses = provider.translate_batch(requests) + + assert mock_post.call_count == 3 + assert [r.translated_text for r in responses] == ["Bonjour", "Au revoir"] + + @patch("requests.post") + def test_wrong_batch_length_falls_back_to_individual(self, mock_post, provider): + # A JSON answer with a missing item is NOT trusted → fallback. + short_batch = self._batch_response(["Bonjour", "Au revoir"]) + individuals = [ + self._mock_post({"choices": [{"message": {"content": "Bonjour"}}], "usage": {}}), + self._mock_post({"choices": [{"message": {"content": "Au revoir"}}], "usage": {}}), + self._mock_post({"choices": [{"message": {"content": "Merci"}}], "usage": {}}), + ] + mock_post.side_effect = [short_batch] + individuals + requests = [ + TranslationRequest(text="Hello", target_language="fr"), + TranslationRequest(text="Goodbye", target_language="fr"), + TranslationRequest(text="Thank you", target_language="fr"), + ] + + responses = provider.translate_batch(requests) + + assert mock_post.call_count == 4 + assert [r.translated_text for r in responses] == ["Bonjour", "Au revoir", "Merci"] + + @patch("requests.post") + def test_mixed_metadata_skips_batch(self, mock_post, provider): + mock_post.side_effect = [ + self._mock_post({"choices": [{"message": {"content": "Bonjour"}}], "usage": {}}), + self._mock_post({"choices": [{"message": {"content": "Au revoir"}}], "usage": {}}), + ] + requests = [ + TranslationRequest( + text="Hello", target_language="fr", metadata={"custom_prompt": "formal"} + ), + TranslationRequest( + text="Goodbye", target_language="fr", metadata={"custom_prompt": "casual"} + ), + ] + + responses = provider.translate_batch(requests) + + # Two individual calls, no numbered-JSON batch attempt. + assert mock_post.call_count == 2 + assert [r.translated_text for r in responses] == ["Bonjour", "Au revoir"] + for call in mock_post.call_args_list: + assert call[1]["json"]["messages"][1]["content"] in ("Hello", "Goodbye") + + @patch("requests.post") + def test_mixed_language_pair_skips_batch(self, mock_post, provider): + mock_post.side_effect = [ + self._mock_post({"choices": [{"message": {"content": "Bonjour"}}], "usage": {}}), + self._mock_post({"choices": [{"message": {"content": "Hola"}}], "usage": {}}), + ] + requests = [ + TranslationRequest(text="Hello", target_language="fr"), + TranslationRequest(text="Hello", target_language="es"), + ] + + responses = provider.translate_batch(requests) + + assert mock_post.call_count == 2 + assert [r.translated_text for r in responses] == ["Bonjour", "Hola"] + + @patch("requests.post") + def test_single_request_is_individual(self, mock_post, provider): + mock_post.return_value = self._mock_post( + {"choices": [{"message": {"content": "Bonjour"}}], "usage": {}} + ) + requests = [TranslationRequest(text="Hello", target_language="fr")] + + responses = provider.translate_batch(requests) + + assert mock_post.call_count == 1 + assert responses[0].translated_text == "Bonjour" + # The user message is the raw text, not a numbered JSON list. + user_message = mock_post.call_args[1]["json"]["messages"][1]["content"] + assert user_message == "Hello" + + def test_empty_batch_returns_empty(self, provider): + assert provider.translate_batch([]) == [] + + +class TestDeepSeekProviderError: + def test_error_carries_code_and_message(self): + err = DeepSeekProviderError(DEEPSEEK_TIMEOUT, "timed out", details={"wait": 1}) + assert err.code == DEEPSEEK_TIMEOUT + assert err.message == "timed out" + assert err.details == {"wait": 1} + + def test_rate_limited_code_exists(self): + # Sanity: the retryable codes used by the provider are importable. + assert DEEPSEEK_RATE_LIMITED == "DEEPSEEK_RATE_LIMITED" diff --git a/tests/test_providers/test_minimax_provider.py b/tests/test_providers/test_minimax_provider.py index 3ab563a..9569134 100644 --- a/tests/test_providers/test_minimax_provider.py +++ b/tests/test_providers/test_minimax_provider.py @@ -7,8 +7,15 @@ Validates Bug 1 fix: - ``is_available()`` / ``health_check()`` tolerate a missing ``/models`` path (Minimax does not document it) and only mark the provider down on 401/network error - translation success, 429 retry, 401 handling + +Batch mode (ported from the OpenAI provider): +- same language pair + metadata chunk → ONE numbered-JSON request for all texts +- unparseable/failed batch answer → per-item fallback, every text still translated +- mixed metadata or a single text → no batch (individual requests) """ +import json + import pytest from unittest.mock import patch, MagicMock from requests.exceptions import Timeout @@ -159,3 +166,198 @@ class TestMinimaxProviderError: assert err.code == MINIMAX_TIMEOUT assert err.message == "timed out" assert err.details == {"wait": 1} + + +class TestMinimaxBatchTranslation: + """Batch mode: one numbered-JSON request per same pair/metadata chunk.""" + + @pytest.fixture + def provider(self): + return MinimaxTranslationProvider(api_key="k", max_retries=0) + + def _mock_post(self, payload, status_code=200): + mock_response = MagicMock() + mock_response.status_code = status_code + mock_response.json.return_value = payload + mock_response.text = "" + return mock_response + + def _batch_response(self, translations): + """Mock a 200 response whose content is a numbered JSON array.""" + content = json.dumps( + [{"id": i, "translation": t} for i, t in enumerate(translations)] + ) + return self._mock_post({"choices": [{"message": {"content": content}}], "usage": {}}) + + @patch("requests.post") + def test_batch_of_three_uses_single_request(self, mock_post, provider): + mock_post.return_value = self._batch_response(["Bonjour", "Au revoir", "Merci"]) + requests = [ + TranslationRequest(text="Hello", target_language="fr"), + TranslationRequest(text="Goodbye", target_language="fr"), + TranslationRequest(text="Thank you", target_language="fr"), + ] + + responses = provider.translate_batch(requests) + + assert mock_post.call_count == 1 + assert [r.translated_text for r in responses] == ["Bonjour", "Au revoir", "Merci"] + assert all(r.provider_name == "minimax" for r in responses) + assert all(r.error is None for r in responses) + # The single call went to the public chat completions endpoint. + assert mock_post.call_args[0][0] == "https://api.minimax.io/v1/chat/completions" + + @patch("requests.post") + def test_batch_request_payload_is_numbered_json_with_instructions(self, mock_post, provider): + mock_post.return_value = self._batch_response(["Bonjour", "Au revoir"]) + requests = [ + TranslationRequest( + text="Hello", + target_language="fr", + source_language="en", + metadata={"custom_prompt": "use a formal tone"}, + ), + TranslationRequest( + text="Goodbye", + target_language="fr", + source_language="en", + metadata={"custom_prompt": "use a formal tone"}, + ), + ] + + provider.translate_batch(requests) + + payload = mock_post.call_args[1]["json"] + system = payload["messages"][0]["content"] + user = payload["messages"][1]["content"] + # System prompt = base instructions + custom prompt + strict JSON rule. + assert "professional translator" in system + assert "ADDITIONAL CONTEXT AND INSTRUCTIONS" in system + assert "use a formal tone" in system + assert "BATCH MODE" in system + assert "ONLY a JSON array" in system + # User message = the numbered JSON list of the texts. + parsed = json.loads(user) + assert [item["text"] for item in parsed] == ["Hello", "Goodbye"] + assert [item["id"] for item in parsed] == [0, 1] + + @patch("requests.post") + def test_invalid_batch_json_falls_back_to_individual(self, mock_post, provider): + broken_batch = self._mock_post( + {"choices": [{"message": {"content": "désolé, pas de JSON ici"}}], "usage": {}} + ) + individuals = [ + self._mock_post({"choices": [{"message": {"content": "Bonjour"}}], "usage": {}}), + self._mock_post({"choices": [{"message": {"content": "Au revoir"}}], "usage": {}}), + self._mock_post({"choices": [{"message": {"content": "Merci"}}], "usage": {}}), + ] + mock_post.side_effect = [broken_batch] + individuals + requests = [ + TranslationRequest(text="Hello", target_language="fr"), + TranslationRequest(text="Goodbye", target_language="fr"), + TranslationRequest(text="Thank you", target_language="fr"), + ] + + responses = provider.translate_batch(requests) + + # 1 failed batch attempt + 3 individual calls. + assert mock_post.call_count == 4 + assert [r.translated_text for r in responses] == ["Bonjour", "Au revoir", "Merci"] + assert all(r.error is None for r in responses) + + @patch("requests.post") + def test_batch_api_error_falls_back_to_individual(self, mock_post, provider): + # The batch call itself fails (500) — per-item calls still translate everything. + failed_batch = self._mock_post({"error": "boom"}, status_code=500) + individuals = [ + self._mock_post({"choices": [{"message": {"content": "Bonjour"}}], "usage": {}}), + self._mock_post({"choices": [{"message": {"content": "Au revoir"}}], "usage": {}}), + ] + mock_post.side_effect = [failed_batch] + individuals + requests = [ + TranslationRequest(text="Hello", target_language="fr"), + TranslationRequest(text="Goodbye", target_language="fr"), + ] + + responses = provider.translate_batch(requests) + + assert mock_post.call_count == 3 + assert [r.translated_text for r in responses] == ["Bonjour", "Au revoir"] + + @patch("requests.post") + def test_wrong_batch_length_falls_back_to_individual(self, mock_post, provider): + # A JSON answer with a missing item is NOT trusted → fallback. + short_batch = self._batch_response(["Bonjour", "Au revoir"]) + individuals = [ + self._mock_post({"choices": [{"message": {"content": "Bonjour"}}], "usage": {}}), + self._mock_post({"choices": [{"message": {"content": "Au revoir"}}], "usage": {}}), + self._mock_post({"choices": [{"message": {"content": "Merci"}}], "usage": {}}), + ] + mock_post.side_effect = [short_batch] + individuals + requests = [ + TranslationRequest(text="Hello", target_language="fr"), + TranslationRequest(text="Goodbye", target_language="fr"), + TranslationRequest(text="Thank you", target_language="fr"), + ] + + responses = provider.translate_batch(requests) + + assert mock_post.call_count == 4 + assert [r.translated_text for r in responses] == ["Bonjour", "Au revoir", "Merci"] + + @patch("requests.post") + def test_mixed_metadata_skips_batch(self, mock_post, provider): + mock_post.side_effect = [ + self._mock_post({"choices": [{"message": {"content": "Bonjour"}}], "usage": {}}), + self._mock_post({"choices": [{"message": {"content": "Au revoir"}}], "usage": {}}), + ] + requests = [ + TranslationRequest( + text="Hello", target_language="fr", metadata={"custom_prompt": "formal"} + ), + TranslationRequest( + text="Goodbye", target_language="fr", metadata={"custom_prompt": "casual"} + ), + ] + + responses = provider.translate_batch(requests) + + # Two individual calls, no numbered-JSON batch attempt. + assert mock_post.call_count == 2 + assert [r.translated_text for r in responses] == ["Bonjour", "Au revoir"] + for call in mock_post.call_args_list: + assert call[1]["json"]["messages"][1]["content"] in ("Hello", "Goodbye") + + @patch("requests.post") + def test_mixed_language_pair_skips_batch(self, mock_post, provider): + mock_post.side_effect = [ + self._mock_post({"choices": [{"message": {"content": "Bonjour"}}], "usage": {}}), + self._mock_post({"choices": [{"message": {"content": "Hola"}}], "usage": {}}), + ] + requests = [ + TranslationRequest(text="Hello", target_language="fr"), + TranslationRequest(text="Hello", target_language="es"), + ] + + responses = provider.translate_batch(requests) + + assert mock_post.call_count == 2 + assert [r.translated_text for r in responses] == ["Bonjour", "Hola"] + + @patch("requests.post") + def test_single_request_is_individual(self, mock_post, provider): + mock_post.return_value = self._mock_post( + {"choices": [{"message": {"content": "Bonjour"}}], "usage": {}} + ) + requests = [TranslationRequest(text="Hello", target_language="fr")] + + responses = provider.translate_batch(requests) + + assert mock_post.call_count == 1 + assert responses[0].translated_text == "Bonjour" + # The user message is the raw text, not a numbered JSON list. + user_message = mock_post.call_args[1]["json"]["messages"][1]["content"] + assert user_message == "Hello" + + def test_empty_batch_returns_empty(self, provider): + assert provider.translate_batch([]) == []