fix(python): huggingface TEI auth, procedural-memory content handling, and proxy pip auto-install (#6947)
This commit is contained in:
@@ -72,7 +72,8 @@ def test_embed_with_custom_embedding_dims(mock_sentence_transformer):
|
||||
assert result == [1.0, 1.1, 1.2]
|
||||
|
||||
|
||||
def test_embed_with_huggingface_base_url():
|
||||
def test_embed_with_huggingface_base_url(monkeypatch):
|
||||
monkeypatch.delenv("HUGGINGFACE_API_KEY", raising=False)
|
||||
config = BaseEmbedderConfig(
|
||||
huggingface_base_url="http://localhost:8080",
|
||||
model="my-custom-model",
|
||||
@@ -81,20 +82,20 @@ def test_embed_with_huggingface_base_url():
|
||||
with patch("mem0.embeddings.huggingface.OpenAI") as mock_openai:
|
||||
mock_client = Mock()
|
||||
mock_openai.return_value = mock_client
|
||||
|
||||
|
||||
# Create a mock for the response object and its attributes
|
||||
mock_embedding_response = Mock()
|
||||
mock_embedding_response.embedding = [0.1, 0.2, 0.3]
|
||||
|
||||
|
||||
mock_create_response = Mock()
|
||||
mock_create_response.data = [mock_embedding_response]
|
||||
|
||||
|
||||
mock_client.embeddings.create.return_value = mock_create_response
|
||||
|
||||
embedder = HuggingFaceEmbedding(config)
|
||||
result = embedder.embed("Hello from custom endpoint")
|
||||
|
||||
mock_openai.assert_called_once_with(base_url="http://localhost:8080", api_key=None)
|
||||
mock_openai.assert_called_once_with(base_url="http://localhost:8080", api_key="hf")
|
||||
mock_client.embeddings.create.assert_called_once_with(
|
||||
input="Hello from custom endpoint",
|
||||
model="my-custom-model",
|
||||
@@ -115,6 +116,26 @@ def test_embed_with_huggingface_base_url_forwards_api_key():
|
||||
mock_openai.assert_called_once_with(base_url="http://localhost:8080", api_key="tei-token")
|
||||
|
||||
|
||||
def test_embed_with_huggingface_base_url_falls_back_to_env_var(monkeypatch):
|
||||
monkeypatch.setenv("HUGGINGFACE_API_KEY", "env-token")
|
||||
config = BaseEmbedderConfig(huggingface_base_url="http://localhost:8080", model="my-custom-model")
|
||||
with patch("mem0.embeddings.huggingface.OpenAI") as mock_openai:
|
||||
HuggingFaceEmbedding(config)
|
||||
|
||||
mock_openai.assert_called_once_with(base_url="http://localhost:8080", api_key="env-token")
|
||||
|
||||
|
||||
def test_embed_with_huggingface_base_url_config_key_beats_env_var(monkeypatch):
|
||||
monkeypatch.setenv("HUGGINGFACE_API_KEY", "env-token")
|
||||
config = BaseEmbedderConfig(
|
||||
huggingface_base_url="http://localhost:8080", model="my-custom-model", api_key="config-token"
|
||||
)
|
||||
with patch("mem0.embeddings.huggingface.OpenAI") as mock_openai:
|
||||
HuggingFaceEmbedding(config)
|
||||
|
||||
mock_openai.assert_called_once_with(base_url="http://localhost:8080", api_key="config-token")
|
||||
|
||||
|
||||
def test_embed_batch_sentence_transformer(mock_sentence_transformer):
|
||||
config = BaseEmbedderConfig()
|
||||
embedder = HuggingFaceEmbedding(config)
|
||||
|
||||
@@ -0,0 +1,61 @@
|
||||
import pytest
|
||||
|
||||
pytest.importorskip("langchain", reason="langchain is an optional extra")
|
||||
|
||||
from langchain.embeddings.base import Embeddings # noqa: E402
|
||||
|
||||
from mem0.configs.embeddings.base import BaseEmbedderConfig # noqa: E402
|
||||
from mem0.embeddings.langchain import LangchainEmbedding # noqa: E402
|
||||
|
||||
|
||||
class DummyEmbeddings(Embeddings):
|
||||
def __init__(self):
|
||||
self.queries = []
|
||||
|
||||
def embed_documents(self, texts):
|
||||
return [[0.1, 0.2, 0.3] for _ in texts]
|
||||
|
||||
def embed_query(self, text):
|
||||
self.queries.append(text)
|
||||
return [0.1, 0.2, 0.3]
|
||||
|
||||
|
||||
def test_missing_model_raises():
|
||||
with pytest.raises(ValueError, match="`model` parameter is required"):
|
||||
LangchainEmbedding(BaseEmbedderConfig())
|
||||
|
||||
|
||||
def test_model_that_is_not_an_embeddings_instance_raises():
|
||||
with pytest.raises(ValueError, match="`model` must be an instance of Embeddings"):
|
||||
LangchainEmbedding(BaseEmbedderConfig(model="text-embedding-3-small"))
|
||||
|
||||
|
||||
def test_configured_model_instance_is_kept():
|
||||
model = DummyEmbeddings()
|
||||
|
||||
assert LangchainEmbedding(BaseEmbedderConfig(model=model)).langchain_model is model
|
||||
|
||||
|
||||
def test_embed_delegates_to_embed_query():
|
||||
model = DummyEmbeddings()
|
||||
embedder = LangchainEmbedding(BaseEmbedderConfig(model=model))
|
||||
|
||||
assert embedder.embed("hello", "add") == [0.1, 0.2, 0.3]
|
||||
assert model.queries == ["hello"]
|
||||
|
||||
|
||||
def test_memory_action_never_reaches_the_model():
|
||||
model = DummyEmbeddings()
|
||||
embedder = LangchainEmbedding(BaseEmbedderConfig(model=model))
|
||||
|
||||
embedder.embed("hello", memory_action="add")
|
||||
|
||||
assert model.queries == ["hello"]
|
||||
|
||||
|
||||
def test_embed_batch_falls_back_to_sequential_embed():
|
||||
model = DummyEmbeddings()
|
||||
embedder = LangchainEmbedding(BaseEmbedderConfig(model=model))
|
||||
|
||||
assert embedder.embed_batch(["a", "b"]) == [[0.1, 0.2, 0.3], [0.1, 0.2, 0.3]]
|
||||
assert model.queries == ["a", "b"]
|
||||
@@ -1,3 +1,6 @@
|
||||
import builtins
|
||||
import importlib
|
||||
import sys
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
@@ -94,3 +97,25 @@ def test_embed_batch_count_mismatch_raises(mock_ollama_client):
|
||||
|
||||
with pytest.raises(ValueError, match="returned 1 embeddings for 2 texts"):
|
||||
embedder.embed_batch(["first text", "second text"])
|
||||
|
||||
|
||||
def test_missing_ollama_raises_actionable_import_error(monkeypatch):
|
||||
"""Missing ollama raises a catchable ImportError naming the install command, never input()/sys.exit()."""
|
||||
monkeypatch.delitem(sys.modules, "mem0.embeddings.ollama", raising=False)
|
||||
monkeypatch.delitem(sys.modules, "ollama", raising=False)
|
||||
|
||||
real_import = builtins.__import__
|
||||
|
||||
def fake_import(name, *args, **kwargs):
|
||||
if name == "ollama":
|
||||
raise ImportError("No module named 'ollama'")
|
||||
return real_import(name, *args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", fake_import)
|
||||
|
||||
with patch("builtins.input") as mock_input, patch("sys.exit") as mock_exit:
|
||||
with pytest.raises(ImportError, match="pip install ollama"):
|
||||
importlib.import_module("mem0.embeddings.ollama")
|
||||
|
||||
mock_input.assert_not_called()
|
||||
mock_exit.assert_not_called()
|
||||
|
||||
@@ -192,3 +192,16 @@ class TestProcessTelemetryFilters:
|
||||
class TestRemoveCodeBlocks:
|
||||
def test_none_content_returns_empty_string(self):
|
||||
assert remove_code_blocks(None) == ""
|
||||
|
||||
def test_list_content_from_langchain_aimessage_is_joined(self):
|
||||
assert remove_code_blocks([{"type": "text", "text": "hello"}]) == "hello"
|
||||
|
||||
def test_list_content_joins_multiple_blocks_and_bare_strings(self):
|
||||
content = [{"type": "text", "text": "hello "}, "world", {"type": "thinking", "thinking": "ignored"}]
|
||||
assert remove_code_blocks(content) == "hello world"
|
||||
|
||||
def test_list_content_still_strips_code_fences(self):
|
||||
assert remove_code_blocks([{"type": "text", "text": "```json\n{}\n```"}]) == "{}"
|
||||
|
||||
def test_unsupported_content_type_returns_empty_string(self):
|
||||
assert remove_code_blocks(42) == ""
|
||||
|
||||
+93
-2
@@ -1235,9 +1235,10 @@ class TestHybridSearchWarning:
|
||||
def test_warning_for_store_without_keyword_search(
|
||||
self, mock_vs_factory, mock_emb, mock_llm, mock_sqlite, _cap, caplog
|
||||
):
|
||||
from mem0.vector_stores.base import VectorStoreBase
|
||||
import logging
|
||||
|
||||
from mem0.vector_stores.base import VectorStoreBase
|
||||
|
||||
class StoreWithoutKeywordSearch(VectorStoreBase):
|
||||
def create_col(self, *a, **kw): pass
|
||||
def insert(self, *a, **kw): pass
|
||||
@@ -1271,9 +1272,10 @@ class TestHybridSearchWarning:
|
||||
def test_no_warning_for_store_with_keyword_search(
|
||||
self, mock_vs_factory, mock_emb, mock_llm, mock_sqlite, _cap, caplog
|
||||
):
|
||||
from mem0.vector_stores.base import VectorStoreBase
|
||||
import logging
|
||||
|
||||
from mem0.vector_stores.base import VectorStoreBase
|
||||
|
||||
class StoreWithKeywordSearch(VectorStoreBase):
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
return []
|
||||
@@ -1574,3 +1576,92 @@ async def test_async_procedural_memory_default_path_without_langchain(mock_llm_f
|
||||
|
||||
assert result["results"][0]["event"] == "ADD"
|
||||
memory.llm.generate_response.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("mem0.memory.main.VectorStoreFactory")
|
||||
@patch("mem0.memory.main.EmbedderFactory")
|
||||
@patch("mem0.memory.main.LlmFactory")
|
||||
async def test_async_procedural_memory_langchain_list_content(mock_llm_factory, mock_emb, mock_vs):
|
||||
"""Regression #6150: a LangChain block list must be extracted, not dropped."""
|
||||
mock_vs.return_value = MagicMock()
|
||||
mock_emb.return_value = MagicMock()
|
||||
mock_emb.return_value.embed.return_value = [0.1] * 1536
|
||||
mock_llm_factory.return_value = MagicMock()
|
||||
|
||||
from mem0.memory.main import AsyncMemory
|
||||
|
||||
memory = AsyncMemory(MemoryConfig())
|
||||
memory.vector_store = MagicMock()
|
||||
memory.vector_store.insert = MagicMock()
|
||||
|
||||
mock_langchain_llm = MagicMock()
|
||||
mock_response = MagicMock()
|
||||
mock_response.content = [{"type": "text", "text": "- deploy with the release script"}]
|
||||
mock_langchain_llm.invoke.return_value = mock_response
|
||||
|
||||
await memory._create_procedural_memory(
|
||||
[{"role": "user", "content": "how do we deploy"}],
|
||||
metadata={"user_id": "test_user"},
|
||||
llm=mock_langchain_llm,
|
||||
)
|
||||
|
||||
stored_data = memory.vector_store.insert.call_args[1]["payloads"][0]["data"]
|
||||
assert stored_data == "- deploy with the release script"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("mem0.memory.main.VectorStoreFactory")
|
||||
@patch("mem0.memory.main.EmbedderFactory")
|
||||
@patch("mem0.memory.main.LlmFactory")
|
||||
async def test_async_procedural_memory_empty_content_raises(mock_llm_factory, mock_emb, mock_vs):
|
||||
"""Regression #6150: a refusal (content=None) must raise, not store an empty memory."""
|
||||
mock_vs.return_value = MagicMock()
|
||||
mock_emb.return_value = MagicMock()
|
||||
mock_llm_factory.return_value = MagicMock()
|
||||
|
||||
from mem0.memory.main import AsyncMemory
|
||||
|
||||
memory = AsyncMemory(MemoryConfig())
|
||||
memory.vector_store = MagicMock()
|
||||
memory.vector_store.insert = MagicMock()
|
||||
|
||||
mock_langchain_llm = MagicMock()
|
||||
mock_response = MagicMock()
|
||||
mock_response.content = None
|
||||
mock_langchain_llm.invoke.return_value = mock_response
|
||||
|
||||
with pytest.raises(ValueError, match="no content for the procedural memory summary"):
|
||||
await memory._create_procedural_memory(
|
||||
[{"role": "user", "content": "how do we deploy"}],
|
||||
metadata={"user_id": "test_user"},
|
||||
llm=mock_langchain_llm,
|
||||
)
|
||||
|
||||
memory.vector_store.insert.assert_not_called()
|
||||
|
||||
|
||||
@patch("mem0.utils.factory.EmbedderFactory.create")
|
||||
@patch("mem0.utils.factory.VectorStoreFactory.create")
|
||||
@patch("mem0.utils.factory.LlmFactory.create")
|
||||
@patch("mem0.memory.storage.SQLiteManager")
|
||||
def test_sync_procedural_memory_empty_content_raises(
|
||||
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
|
||||
):
|
||||
"""Regression #6150: the sync path must raise on empty LLM content too."""
|
||||
mock_embedder_factory.return_value = MagicMock()
|
||||
mock_vector_store = MagicMock()
|
||||
mock_vector_factory.return_value = mock_vector_store
|
||||
mock_llm_factory.return_value = MagicMock()
|
||||
mock_sqlite.return_value = MagicMock()
|
||||
|
||||
memory = Memory(MemoryConfig())
|
||||
memory.llm.generate_response = Mock(return_value=None)
|
||||
|
||||
with pytest.raises(ValueError, match="no content for the procedural memory summary"):
|
||||
memory._create_procedural_memory(
|
||||
[{"role": "user", "content": "how do we deploy"}],
|
||||
metadata={"user_id": "test_user"},
|
||||
)
|
||||
|
||||
mock_vector_store.insert.assert_not_called()
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
import builtins
|
||||
import importlib
|
||||
import inspect
|
||||
import sys
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
@@ -138,3 +141,20 @@ def test_completions_create_messages_default_does_not_leak_between_calls(mock_me
|
||||
f"Completions.create(messages=...) must default to None to avoid the "
|
||||
f"B006 shared-default-list bug; got {messages_default!r}."
|
||||
)
|
||||
|
||||
|
||||
def test_missing_litellm_raises_actionable_import_error(monkeypatch):
|
||||
monkeypatch.delitem(sys.modules, "mem0.proxy.main", raising=False)
|
||||
monkeypatch.delitem(sys.modules, "litellm", raising=False)
|
||||
|
||||
real_import = builtins.__import__
|
||||
|
||||
def fake_import(name, *args, **kwargs):
|
||||
if name == "litellm":
|
||||
raise ImportError("No module named 'litellm'")
|
||||
return real_import(name, *args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", fake_import)
|
||||
|
||||
with pytest.raises(ImportError, match="pip install litellm"):
|
||||
importlib.import_module("mem0.proxy.main")
|
||||
|
||||
Reference in New Issue
Block a user