2ef6291307
The Python Together LLM and embedder ignored TOGETHER_API_BASE, unlike their TypeScript siblings, so users behind a gateway or proxy could not redirect Together traffic. Add together_base_url to a new TogetherConfig (LLM side) and to BaseEmbedderConfig (embedder side, matching the existing openai_base_url/huggingface_base_url pattern), and resolve base_url with explicit config value > TOGETHER_API_BASE env var > hardcoded default in both mem0/llms/together.py and mem0/embeddings/together.py. Also register TogetherConfig in LlmFactory.provider_to_class so config dicts routed through Memory pick it up. Closes #6587
114 lines
4.5 KiB
Python
114 lines
4.5 KiB
Python
from unittest.mock import Mock, patch
|
|
|
|
import pytest
|
|
|
|
from mem0.configs.embeddings.base import BaseEmbedderConfig
|
|
from mem0.embeddings.together import TogetherEmbedding
|
|
|
|
DEFAULT_MODEL = "intfloat/multilingual-e5-large-instruct"
|
|
DEFAULT_EMBEDDING_DIMS = 1024
|
|
|
|
|
|
@pytest.fixture
|
|
def mock_together_client():
|
|
with patch("mem0.embeddings.together.Together") as mock_together:
|
|
mock_client = Mock()
|
|
mock_together.return_value = mock_client
|
|
yield mock_client
|
|
|
|
|
|
def test_embed_text(mock_together_client):
|
|
config = BaseEmbedderConfig(model=DEFAULT_MODEL, embedding_dims=DEFAULT_EMBEDDING_DIMS)
|
|
embedder = TogetherEmbedding(config)
|
|
|
|
mock_together_client.embeddings.create.return_value = Mock(data=[Mock(embedding=[0.1, 0.2, 0.3, 0.4, 0.5])])
|
|
|
|
text = "Sample text to embed."
|
|
embedding = embedder.embed(text)
|
|
|
|
mock_together_client.embeddings.create.assert_called_once_with(model=DEFAULT_MODEL, input=text)
|
|
assert embedding == [0.1, 0.2, 0.3, 0.4, 0.5]
|
|
|
|
|
|
def test_embed_batch_single_call(mock_together_client):
|
|
config = BaseEmbedderConfig(model=DEFAULT_MODEL, embedding_dims=DEFAULT_EMBEDDING_DIMS)
|
|
embedder = TogetherEmbedding(config)
|
|
|
|
mock_item0 = Mock(index=0, embedding=[0.1, 0.2, 0.3])
|
|
mock_item1 = Mock(index=1, embedding=[0.4, 0.5, 0.6])
|
|
mock_together_client.embeddings.create.return_value = Mock(data=[mock_item0, mock_item1])
|
|
|
|
texts = ["First text.", "Second text."]
|
|
embeddings = embedder.embed_batch(texts)
|
|
|
|
mock_together_client.embeddings.create.assert_called_once_with(model=DEFAULT_MODEL, input=texts)
|
|
assert embeddings == [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]
|
|
|
|
|
|
def test_embed_batch_empty_list(mock_together_client):
|
|
config = BaseEmbedderConfig(model=DEFAULT_MODEL, embedding_dims=DEFAULT_EMBEDDING_DIMS)
|
|
embedder = TogetherEmbedding(config)
|
|
|
|
result = embedder.embed_batch([])
|
|
|
|
assert result == []
|
|
mock_together_client.embeddings.create.assert_not_called()
|
|
|
|
|
|
def test_embed_batch_count_mismatch_raises(mock_together_client):
|
|
config = BaseEmbedderConfig(model=DEFAULT_MODEL, embedding_dims=DEFAULT_EMBEDDING_DIMS)
|
|
embedder = TogetherEmbedding(config)
|
|
|
|
mock_item0 = Mock(index=0, embedding=[0.1, 0.2, 0.3])
|
|
mock_together_client.embeddings.create.return_value = Mock(data=[mock_item0])
|
|
|
|
with pytest.raises(ValueError, match="returned 1 embeddings for 2 texts"):
|
|
embedder.embed_batch(["first text", "second text"])
|
|
|
|
|
|
def test_default_config_applies_together_defaults(mock_together_client):
|
|
embedder = TogetherEmbedding(BaseEmbedderConfig())
|
|
|
|
assert embedder.config.model == DEFAULT_MODEL
|
|
assert embedder.config.embedding_dims == DEFAULT_EMBEDDING_DIMS
|
|
|
|
|
|
def test_explicit_config_overrides_defaults(mock_together_client):
|
|
# The `config.x or default` wiring must honor user-provided values, not clobber them.
|
|
config = BaseEmbedderConfig(model="BAAI/bge-base-en-v1.5", embedding_dims=768)
|
|
embedder = TogetherEmbedding(config)
|
|
|
|
assert embedder.config.model == "BAAI/bge-base-en-v1.5"
|
|
assert embedder.config.embedding_dims == 768
|
|
|
|
# ...and the chosen model actually reaches the Together API call.
|
|
mock_together_client.embeddings.create.return_value = Mock(data=[Mock(embedding=[0.0] * 768)])
|
|
embedder.embed("hello")
|
|
mock_together_client.embeddings.create.assert_called_once_with(model="BAAI/bge-base-en-v1.5", input="hello")
|
|
|
|
|
|
def test_together_embedding_base_url(monkeypatch):
|
|
# case1: default config uses Together's official base url
|
|
monkeypatch.delenv("TOGETHER_API_BASE", raising=False)
|
|
config = BaseEmbedderConfig(model=DEFAULT_MODEL, embedding_dims=DEFAULT_EMBEDDING_DIMS, api_key="api_key")
|
|
embedder = TogetherEmbedding(config)
|
|
assert str(embedder.client.base_url) == "https://api.together.ai/v1/"
|
|
|
|
# case2: with env variable TOGETHER_API_BASE
|
|
provider_base_url = "https://api.provider.com/v1/"
|
|
monkeypatch.setenv("TOGETHER_API_BASE", provider_base_url)
|
|
config = BaseEmbedderConfig(model=DEFAULT_MODEL, embedding_dims=DEFAULT_EMBEDDING_DIMS, api_key="api_key")
|
|
embedder = TogetherEmbedding(config)
|
|
assert str(embedder.client.base_url) == provider_base_url
|
|
|
|
# case3: with config.together_base_url (explicit config beats env var)
|
|
config_base_url = "https://api.config.com/v1/"
|
|
config = BaseEmbedderConfig(
|
|
model=DEFAULT_MODEL,
|
|
embedding_dims=DEFAULT_EMBEDDING_DIMS,
|
|
api_key="api_key",
|
|
together_base_url=config_base_url,
|
|
)
|
|
embedder = TogetherEmbedding(config)
|
|
assert str(embedder.client.base_url) == config_base_url
|