fix(python): huggingface TEI auth, procedural-memory content handling, and proxy pip auto-install (#6947)

This commit is contained in:
Kartik
2026-08-20 15:56:55 +05:30
committed by GitHub
parent 530d802b55
commit ed38ddf873
11 changed files with 264 additions and 18 deletions
@@ -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()
+13
View File
@@ -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
View File
@@ -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()
+20
View File
@@ -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")