From 0aacd9b580780f8b72b7453a2a67269468090c96 Mon Sep 17 00:00:00 2001 From: mia-fourier Date: Wed, 15 Jul 2026 03:51:25 -0500 Subject: [PATCH] Retry transient embedding API failures --- mnemosyne/core/embeddings.py | 18 +++++-- tests/test_embedding_api_retry.py | 80 +++++++++++++++++++++++++++++++ 2 files changed, 95 insertions(+), 3 deletions(-) create mode 100644 tests/test_embedding_api_retry.py diff --git a/mnemosyne/core/embeddings.py b/mnemosyne/core/embeddings.py index d66bdd6b..4761ea2e 100644 --- a/mnemosyne/core/embeddings.py +++ b/mnemosyne/core/embeddings.py @@ -7,7 +7,10 @@ import json import os +import random import ssl +import time +import urllib.error import urllib.request from typing import List, Optional from functools import lru_cache @@ -231,6 +234,16 @@ def _is_rate_limit_error(exc: BaseException) -> bool: return False +def _is_transient_embedding_error(exc: Exception) -> bool: + """Return whether an API embedding request is safe to retry.""" + if isinstance(exc, urllib.error.HTTPError): + return exc.code == 429 or 500 <= exc.code < 600 + if isinstance(exc, (urllib.error.URLError, TimeoutError, ConnectionError, OSError)): + return True + message = str(exc).lower() + return "429" in message or "rate limit" in message or "rate-limit" in message + + def _embed_api(texts: List[str]) -> Optional[np.ndarray]: """Embed texts via OpenAI-compatible API (OpenRouter or custom endpoint).""" global _API_CALL_COUNT @@ -269,9 +282,8 @@ def _embed_api(texts: List[str]) -> Optional[np.ndarray]: _API_CALL_COUNT += 1 return np.array(embeddings, dtype=np.float32) except Exception as e: - if "429" in str(e) or "rate" in str(e).lower(): - import time - time.sleep(2 ** attempt) + if _is_transient_embedding_error(e) and attempt < 2: + time.sleep((0.5 * (2 ** attempt)) + random.uniform(0.0, 0.25)) continue return None diff --git a/tests/test_embedding_api_retry.py b/tests/test_embedding_api_retry.py new file mode 100644 index 00000000..b94fa914 --- /dev/null +++ b/tests/test_embedding_api_retry.py @@ -0,0 +1,80 @@ +"""Regression coverage for transient embedding endpoint failures.""" + +import io +import json +import urllib.error +from unittest.mock import patch + +import pytest + +from mnemosyne.core import embeddings + + +class Response: + def __init__(self, payload): + self.body = io.BytesIO(json.dumps(payload).encode()) + + def __enter__(self): + return self + + def __exit__(self, *_args): + return False + + def read(self): + return self.body.read() + + +def test_embed_api_retries_transient_network_failures(monkeypatch): + monkeypatch.setenv("MNEMOSYNE_EMBEDDING_API_URL", "http://127.0.0.1:11435/v1") + result = Response({"data": [{"embedding": [0.25, 0.75]}]}) + failures = [urllib.error.URLError(OSError(65, "No route to host")), TimeoutError(), result] + + with patch("urllib.request.urlopen", side_effect=failures) as request, \ + patch("mnemosyne.core.embeddings.random.uniform", return_value=0.1), \ + patch("mnemosyne.core.embeddings.time.sleep") as sleep: + vectors = embeddings._embed_api(["retry me"]) + + assert vectors.tolist() == [[0.25, 0.75]] + assert request.call_count == 3 + assert [call.args[0] for call in sleep.call_args_list] == [0.6, 1.1] + + +@pytest.mark.parametrize("status", [429, 503]) +def test_embed_api_retries_transient_http_errors(monkeypatch, status): + monkeypatch.setenv("MNEMOSYNE_EMBEDDING_API_URL", "http://127.0.0.1:11435/v1") + error = urllib.error.HTTPError("http://example", status, "transient", {}, None) + result = Response({"data": [{"embedding": [0.25, 0.75]}]}) + + with patch("urllib.request.urlopen", side_effect=[error, result]) as request, \ + patch("mnemosyne.core.embeddings.random.uniform", return_value=0), \ + patch("mnemosyne.core.embeddings.time.sleep") as sleep: + vectors = embeddings._embed_api(["retry me"]) + + assert vectors.tolist() == [[0.25, 0.75]] + assert request.call_count == 2 + sleep.assert_called_once_with(0.5) + + +def test_embed_api_does_not_retry_nontransient_client_error(monkeypatch): + monkeypatch.setenv("MNEMOSYNE_EMBEDDING_API_URL", "http://127.0.0.1:11435/v1") + error = urllib.error.HTTPError("http://example", 400, "bad request", {}, None) + + with patch("urllib.request.urlopen", side_effect=error) as request, \ + patch("mnemosyne.core.embeddings.time.sleep") as sleep: + assert embeddings._embed_api(["bad request"]) is None + + assert request.call_count == 1 + sleep.assert_not_called() + + +def test_embed_api_stops_after_three_transient_attempts(monkeypatch): + monkeypatch.setenv("MNEMOSYNE_EMBEDDING_API_URL", "http://127.0.0.1:11435/v1") + error = urllib.error.URLError(OSError(65, "No route to host")) + + with patch("urllib.request.urlopen", side_effect=error) as request, \ + patch("mnemosyne.core.embeddings.random.uniform", return_value=0), \ + patch("mnemosyne.core.embeddings.time.sleep") as sleep: + assert embeddings._embed_api(["offline"]) is None + + assert request.call_count == 3 + assert sleep.call_count == 2