fix(python): huggingface TEI auth, procedural-memory content handling, and proxy pip auto-install (#6947)
This commit is contained in:
@@ -1,4 +1,5 @@
|
||||
import logging
|
||||
import os
|
||||
from typing import Literal, Optional
|
||||
|
||||
from openai import OpenAI
|
||||
@@ -17,7 +18,8 @@ class HuggingFaceEmbedding(EmbeddingBase):
|
||||
super().__init__(config)
|
||||
|
||||
if self.config.huggingface_base_url:
|
||||
self.client = OpenAI(base_url=self.config.huggingface_base_url, api_key=self.config.api_key)
|
||||
api_key = self.config.api_key or os.getenv("HUGGINGFACE_API_KEY") or "hf"
|
||||
self.client = OpenAI(base_url=self.config.huggingface_base_url, api_key=api_key)
|
||||
self.config.model = self.config.model or "tei"
|
||||
else:
|
||||
self.config.model = self.config.model or "multi-qa-MiniLM-L6-cos-v1"
|
||||
|
||||
+13
-1
@@ -2017,6 +2017,12 @@ class Memory(MemoryBase):
|
||||
logger.error(f"Error generating procedural memory summary: {e}")
|
||||
raise
|
||||
|
||||
if not procedural_memory:
|
||||
raise ValueError(
|
||||
"The LLM returned no content for the procedural memory summary. "
|
||||
"The model may have declined the request or returned an empty response."
|
||||
)
|
||||
|
||||
if metadata is None:
|
||||
raise ValueError("Metadata cannot be done for procedural memory.")
|
||||
|
||||
@@ -3701,11 +3707,17 @@ class AsyncMemory(MemoryBase):
|
||||
else:
|
||||
procedural_memory = await asyncio.to_thread(self.llm.generate_response, messages=parsed_messages)
|
||||
procedural_memory = remove_code_blocks(procedural_memory)
|
||||
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error generating procedural memory summary: {e}")
|
||||
raise
|
||||
|
||||
if not procedural_memory:
|
||||
raise ValueError(
|
||||
"The LLM returned no content for the procedural memory summary. "
|
||||
"The model may have declined the request or returned an empty response."
|
||||
)
|
||||
|
||||
if metadata is None:
|
||||
raise ValueError("Metadata cannot be done for procedural memory.")
|
||||
|
||||
|
||||
@@ -121,7 +121,15 @@ def remove_code_blocks(content: str) -> str:
|
||||
- If a code block is detected, it returns only the inner content, stripping out the markers.
|
||||
- If no code block markers are found, the original content is returned as-is.
|
||||
"""
|
||||
if content is None:
|
||||
if isinstance(content, list):
|
||||
parts = []
|
||||
for block in content:
|
||||
if isinstance(block, str):
|
||||
parts.append(block)
|
||||
elif isinstance(block, dict):
|
||||
parts.append(block.get("text", ""))
|
||||
content = "".join(parts)
|
||||
if not isinstance(content, str):
|
||||
return ""
|
||||
pattern = r"^```[a-zA-Z0-9]*\n([\s\S]*?)\n```$"
|
||||
match = re.match(pattern, content.strip())
|
||||
|
||||
+1
-8
@@ -1,6 +1,4 @@
|
||||
import logging
|
||||
import subprocess
|
||||
import sys
|
||||
import threading
|
||||
from typing import List, Optional, Union
|
||||
|
||||
@@ -11,12 +9,7 @@ import mem0
|
||||
try:
|
||||
import litellm
|
||||
except ImportError:
|
||||
try:
|
||||
subprocess.check_call([sys.executable, "-m", "pip", "install", "litellm"])
|
||||
import litellm
|
||||
except subprocess.CalledProcessError:
|
||||
print("Failed to install 'litellm'. Please install it manually using 'pip install litellm'.")
|
||||
sys.exit(1)
|
||||
raise ImportError("The 'litellm' library is required. Please install it using 'pip install litellm'.")
|
||||
|
||||
from mem0 import Memory, MemoryClient
|
||||
from mem0.configs.prompts import MEMORY_ANSWER_PROMPT
|
||||
|
||||
@@ -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