diff --git a/src/memos/embedders/universal_api.py b/src/memos/embedders/universal_api.py index 24a022eae..579b9bc93 100644 --- a/src/memos/embedders/universal_api.py +++ b/src/memos/embedders/universal_api.py @@ -56,29 +56,56 @@ def __init__(self, config: UniversalAPIEmbedderConfig): else None, ) + @staticmethod + def _build_embedding_kwargs(model: str, texts: list[str], embedding_dims: int | None) -> dict: + kwargs = {"model": model, "input": texts} + if embedding_dims is not None: + kwargs["dimensions"] = embedding_dims + return kwargs + + def _call_embeddings_api( + self, client, model: str, texts: list[str], timeout: int + ) -> list[list[float]]: + embedding_dims = getattr(self.config, "embedding_dims", None) + kwargs = self._build_embedding_kwargs(model, texts, embedding_dims) + + try: + response = asyncio.run( + asyncio.wait_for( + client.embeddings.create(**kwargs), + timeout=timeout, + ) + ) + return [r.embedding for r in response.data] + except Exception as e: + if embedding_dims is not None: + logger.warning( + "Embeddings request with dimensions=%d failed error_type=%s; " + "retrying without dimensions", + embedding_dims, + type(e).__name__, + ) + fallback_kwargs = self._build_embedding_kwargs(model, texts, None) + response = asyncio.run( + asyncio.wait_for( + client.embeddings.create(**fallback_kwargs), + timeout=timeout, + ) + ) + return [r.embedding for r in response.data] + raise + @log_embedding_call def embed(self, texts: list[str]) -> list[list[float]]: if isinstance(texts, str): texts = [texts] - # Sanitize Unicode to prevent encoding errors with emoji/surrogates texts = [_sanitize_unicode(t) for t in texts] - # Truncate texts if max_tokens is configured texts = self._truncate_texts(texts) if self.provider == "openai" or self.provider == "azure": + timeout = int(os.getenv("MOS_EMBEDDER_TIMEOUT", 5)) try: - - async def _create_embeddings(): - return self.client.embeddings.create( - model=getattr(self.config, "model_name_or_path", "text-embedding-3-large"), - input=texts, - ) - - response = asyncio.run( - asyncio.wait_for( - _create_embeddings(), timeout=int(os.getenv("MOS_EMBEDDER_TIMEOUT", 5)) - ) - ) - return [r.embedding for r in response.data] + model = getattr(self.config, "model_name_or_path", "text-embedding-3-large") + return self._call_embeddings_api(self.client, model, texts, timeout) except Exception as e: if self.use_backup_client: logger.warning( @@ -86,26 +113,18 @@ async def _create_embeddings(): type(e).__name__, ) try: - - async def _create_embeddings_backup(): - return self.backup_client.embeddings.create( - model=getattr( - self.config, - "backup_model_name_or_path", - "text-embedding-3-large", - ), - input=texts, - ) - - response = asyncio.run( - asyncio.wait_for( - _create_embeddings_backup(), - timeout=int(os.getenv("MOS_EMBEDDER_TIMEOUT", 5)), - ) + backup_model = getattr( + self.config, + "backup_model_name_or_path", + "text-embedding-3-large", + ) + return self._call_embeddings_api( + self.backup_client, backup_model, texts, timeout ) - return [r.embedding for r in response.data] - except Exception as e: - raise ValueError(f"Backup embeddings request ended with error: {e}") from e + except Exception as e_backup: + raise ValueError( + f"Backup embeddings request ended with error: {e_backup}" + ) from e_backup else: raise ValueError(f"Embeddings request ended with error: {e}") from e else: diff --git a/tests/embedders/test_universal_api.py b/tests/embedders/test_universal_api.py index fd61b3e9a..18cab8dd2 100644 --- a/tests/embedders/test_universal_api.py +++ b/tests/embedders/test_universal_api.py @@ -1,75 +1,159 @@ -import unittest +"""Tests for UniversalAPIEmbedder.""" +import asyncio + +from types import SimpleNamespace from unittest.mock import MagicMock, patch +import pytest + from memos.configs.embedder import UniversalAPIEmbedderConfig from memos.embedders.universal_api import UniversalAPIEmbedder -class TestUniversalAPIEmbedder(unittest.TestCase): - @patch("memos.embedders.universal_api.OpenAIClient") - def test_embed_single_text(self, mock_openai_client): - """Test embedding a single text with OpenAI provider.""" - # Mock the embeddings.create return value - mock_response = MagicMock() - mock_response.data = [MagicMock(embedding=[0.1, 0.2, 0.3, 0.4])] - mock_openai_client.return_value.embeddings.create.return_value = mock_response +class _DimensionsUnsupportedError(Exception): + """Raised by the mock backend when dimensions are rejected.""" - config = UniversalAPIEmbedderConfig( - provider="openai", - api_key="fake-api-key", - base_url="https://api.openai.com/v1", - model_name_or_path="text-embedding-3-large", - ) - embedder = UniversalAPIEmbedder(config) - text = ["Test input for embedding."] - result = embedder.embed(text) +def _make_config(**overrides): + defaults = { + "provider": "openai", + "api_key": "test-key", + "model_name_or_path": "text-embedding-3-large", + "embedding_dims": None, + } + defaults.update(overrides) + return UniversalAPIEmbedderConfig(**defaults) - # Assert OpenAIClient was created with proper args - mock_openai_client.assert_called_once_with( - api_key="fake-api-key", base_url="https://api.openai.com/v1", default_headers=None + +class TestUniversalAPIEmbedderDimensions: + def test_build_embedding_kwargs_no_dims(self): + kwargs = UniversalAPIEmbedder._build_embedding_kwargs( + "text-embedding-3-large", ["hello"], None ) + assert kwargs == {"model": "text-embedding-3-large", "input": ["hello"]} + assert "dimensions" not in kwargs - # Assert embeddings.create called with correct params - embedder.client.embeddings.create.assert_called_once_with( - model="text-embedding-3-large", - input=text, + def test_build_embedding_kwargs_with_dims(self): + kwargs = UniversalAPIEmbedder._build_embedding_kwargs( + "text-embedding-3-large", ["hello"], 256 + ) + assert kwargs == { + "model": "text-embedding-3-large", + "input": ["hello"], + "dimensions": 256, + } + + def test_build_embedding_kwargs_zero_dims(self): + kwargs = UniversalAPIEmbedder._build_embedding_kwargs( + "text-embedding-3-large", ["hello"], 0 ) + assert kwargs["dimensions"] == 0 - self.assertEqual(len(result[0]), 4) + @patch("memos.embedders.universal_api.OpenAIClient") + def test_embed_passes_embedding_dims_to_api(self, mock_openai_client): + mock_response = MagicMock() + mock_response.data = [MagicMock(embedding=[0.1, 0.2])] + mock_openai_client.return_value.embeddings.create.return_value = mock_response + + config = _make_config(embedding_dims=256) + embedder = UniversalAPIEmbedder(config) + embedder.embed(["hello"]) + + _, kwargs = mock_openai_client.return_value.embeddings.create.call_args + assert kwargs.get("dimensions") == 256 @patch("memos.embedders.universal_api.OpenAIClient") - def test_embed_batch_text(self, mock_openai_client): - """Test embedding multiple texts at once with OpenAI provider.""" - # Mock response for multiple texts + def test_embed_without_dims_does_not_pass_dimensions(self, mock_openai_client): mock_response = MagicMock() - mock_response.data = [ - MagicMock(embedding=[0.1, 0.2]), - MagicMock(embedding=[0.3, 0.4]), - MagicMock(embedding=[0.5, 0.6]), - ] + mock_response.data = [MagicMock(embedding=[0.1, 0.2])] mock_openai_client.return_value.embeddings.create.return_value = mock_response - config = UniversalAPIEmbedderConfig( - provider="openai", - api_key="fake-api-key", - base_url="https://api.openai.com/v1", - model_name_or_path="text-embedding-3-large", + config = _make_config(embedding_dims=None) + embedder = UniversalAPIEmbedder(config) + embedder.embed(["hello"]) + + _, kwargs = mock_openai_client.return_value.embeddings.create.call_args + assert "dimensions" not in kwargs + + @patch("memos.embedders.universal_api.OpenAIClient") + def test_embed_with_backup_client(self, mock_openai_client): + primary_client = MagicMock() + primary_client.embeddings.create.side_effect = ValueError("down") + backup_response = MagicMock() + backup_response.data = [MagicMock(embedding=[0.1, 0.2])] + backup_client = MagicMock() + backup_client.embeddings.create.return_value = backup_response + + def client_factory(api_key, **kwargs): + if api_key == "primary": + return primary_client + return backup_client + + mock_openai_client.side_effect = client_factory + + config = _make_config( + api_key="primary", + embedding_dims=256, + backup_client=True, + backup_api_key="backup-key", + backup_base_url="https://api.example.com", + backup_model_name_or_path="text-embedding-3-small", ) + embedder = UniversalAPIEmbedder(config) + result = embedder.embed(["hello"]) + assert result == [[0.1, 0.2]] + assert backup_client.embeddings.create.call_count == 1 + + @patch("memos.embedders.universal_api.OpenAIClient") + def test_embed_raises_when_no_backup(self, mock_openai_client): + mock_openai_client.return_value.embeddings.create.side_effect = ValueError("primary failed") + config = _make_config(embedding_dims=256) + embedder = UniversalAPIEmbedder(config) + with pytest.raises(ValueError, match="Embeddings request ended with error"): + embedder.embed(["hello"]) + +class TestUniversalAPIEmbedderFallback: + def test_call_embeddings_api_falls_back_when_dimensions_not_supported(self): + config = _make_config(embedding_dims=256) embedder = UniversalAPIEmbedder(config) - texts = ["First text.", "Second text.", "Third text."] - result = embedder.embed(texts) - embedder.client.embeddings.create.assert_called_once_with( - model="text-embedding-3-large", - input=texts, - ) + mock_client = SimpleNamespace() + call_count = [0] + + def mock_create(**kwargs): + call_count[0] += 1 + if kwargs.get("dimensions") is not None: + raise _DimensionsUnsupportedError("dimensions not supported") + return SimpleNamespace(data=[SimpleNamespace(embedding=[0.1, 0.2])]) + + mock_client.embeddings = SimpleNamespace(create=mock_create) + + with patch.object(asyncio, "wait_for", side_effect=lambda coro, timeout: coro): + result = embedder._call_embeddings_api( + mock_client, "text-embedding-3-large", ["hello"], 5 + ) + + assert call_count[0] == 2 + assert result == [[0.1, 0.2]] + + def test_call_embeddings_api_no_fallback_when_dims_not_set(self): + config = _make_config(embedding_dims=None) + embedder = UniversalAPIEmbedder(config) + + mock_client = SimpleNamespace() + + def mock_create(**kwargs): + if kwargs.get("dimensions") is not None: + raise AssertionError("dimensions should not be passed when embedding_dims is None") + return SimpleNamespace(data=[SimpleNamespace(embedding=[0.1])]) - self.assertEqual(len(result), 3) - self.assertEqual(result[0], [0.1, 0.2]) + mock_client.embeddings = SimpleNamespace(create=mock_create) + with patch.object(asyncio, "wait_for", side_effect=lambda coro, timeout: coro): + result = embedder._call_embeddings_api( + mock_client, "text-embedding-3-large", ["hello"], 5 + ) -if __name__ == "__main__": - unittest.main() + assert result == [[0.1]]