Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 53 additions & 34 deletions src/memos/embedders/universal_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,56 +56,75 @@ 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(
"Embedding request failed error_type=%s; trying backup client",
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:
Expand Down
182 changes: 133 additions & 49 deletions tests/embedders/test_universal_api.py
Original file line number Diff line number Diff line change
@@ -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]]
Loading