fix: OSS Python SDK hygiene batch, 9 small bug fixes (#6770)
This commit is contained in:
@@ -94,7 +94,7 @@ def test_embed_with_huggingface_base_url():
|
||||
embedder = HuggingFaceEmbedding(config)
|
||||
result = embedder.embed("Hello from custom endpoint")
|
||||
|
||||
mock_openai.assert_called_once_with(base_url="http://localhost:8080")
|
||||
mock_openai.assert_called_once_with(base_url="http://localhost:8080", api_key=None)
|
||||
mock_client.embeddings.create.assert_called_once_with(
|
||||
input="Hello from custom endpoint",
|
||||
model="my-custom-model",
|
||||
@@ -103,6 +103,18 @@ def test_embed_with_huggingface_base_url():
|
||||
assert result == [0.1, 0.2, 0.3]
|
||||
|
||||
|
||||
def test_embed_with_huggingface_base_url_forwards_api_key():
|
||||
config = BaseEmbedderConfig(
|
||||
huggingface_base_url="http://localhost:8080",
|
||||
model="my-custom-model",
|
||||
api_key="tei-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="tei-token")
|
||||
|
||||
|
||||
def test_embed_batch_sentence_transformer(mock_sentence_transformer):
|
||||
config = BaseEmbedderConfig()
|
||||
embedder = HuggingFaceEmbedding(config)
|
||||
|
||||
@@ -1,9 +1,12 @@
|
||||
import pytest
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
|
||||
from mem0.memory.utils import (
|
||||
parse_messages,
|
||||
parse_vision_messages,
|
||||
process_telemetry_filters,
|
||||
remove_code_blocks,
|
||||
remove_spaces_from_entities,
|
||||
sanitize_relationship_for_cypher,
|
||||
)
|
||||
@@ -114,6 +117,17 @@ class TestParseVisionMessages:
|
||||
parse_vision_messages(messages, llm=mock_llm)
|
||||
mock_llm.generate_response.assert_not_called()
|
||||
|
||||
def test_download_failure_preserves_original_exception(self):
|
||||
mock_llm = Mock()
|
||||
mock_llm.generate_response.side_effect = ValueError("network down")
|
||||
messages = [
|
||||
{"role": "user", "content": {"type": "image_url", "image_url": {"url": "https://example.com/img.png"}}}
|
||||
]
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
parse_vision_messages(messages, llm=mock_llm)
|
||||
assert isinstance(exc_info.value.__cause__, ValueError)
|
||||
assert "network down" in str(exc_info.value.__cause__)
|
||||
|
||||
|
||||
class TestRemoveSpacesFromEntities:
|
||||
"""
|
||||
@@ -168,3 +182,13 @@ class TestRemoveSpacesFromEntities:
|
||||
f = remove_spaces_from_entities([dict(base)], sanitize_relationship=False)[0]["relationship"]
|
||||
assert t == sanitize_relationship_for_cypher("a/b")
|
||||
assert f == "a/b"
|
||||
|
||||
|
||||
class TestProcessTelemetryFilters:
|
||||
def test_none_filters_returns_empty_list_and_dict(self):
|
||||
assert process_telemetry_filters(None) == ([], {})
|
||||
|
||||
|
||||
class TestRemoveCodeBlocks:
|
||||
def test_none_content_returns_empty_string(self):
|
||||
assert remove_code_blocks(None) == ""
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import sys
|
||||
from types import ModuleType
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
@@ -9,3 +11,33 @@ def mock_llm():
|
||||
mock_llm_instance = MagicMock()
|
||||
mock_factory.create.return_value = mock_llm_instance
|
||||
yield mock_factory, mock_llm_instance
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_cohere(monkeypatch):
|
||||
"""Provide a fake ``cohere`` module so CohereReranker imports/constructs."""
|
||||
fake_cohere = ModuleType("cohere")
|
||||
fake_client = MagicMock()
|
||||
fake_cohere.Client = MagicMock(return_value=fake_client)
|
||||
monkeypatch.setitem(sys.modules, "cohere", fake_cohere)
|
||||
|
||||
import mem0.reranker.cohere_reranker as cohere_reranker
|
||||
|
||||
monkeypatch.setattr(cohere_reranker, "cohere", fake_cohere, raising=False)
|
||||
monkeypatch.setattr(cohere_reranker, "COHERE_AVAILABLE", True, raising=False)
|
||||
return cohere_reranker, fake_client
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_zero_entropy(monkeypatch):
|
||||
"""Provide a fake ``zeroentropy`` module so ZeroEntropyReranker imports."""
|
||||
fake_module = ModuleType("zeroentropy")
|
||||
fake_client = MagicMock()
|
||||
fake_module.ZeroEntropy = MagicMock(return_value=fake_client)
|
||||
monkeypatch.setitem(sys.modules, "zeroentropy", fake_module)
|
||||
|
||||
import mem0.reranker.zero_entropy_reranker as zero_entropy_reranker
|
||||
|
||||
monkeypatch.setattr(zero_entropy_reranker, "ZeroEntropy", fake_module.ZeroEntropy, raising=False)
|
||||
monkeypatch.setattr(zero_entropy_reranker, "ZERO_ENTROPY_AVAILABLE", True, raising=False)
|
||||
return zero_entropy_reranker, fake_client
|
||||
|
||||
@@ -0,0 +1,75 @@
|
||||
"""Regression tests pinning that the reranker fallback path never mutates caller-owned documents."""
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from mem0.configs.rerankers.cohere import CohereRerankerConfig
|
||||
from mem0.configs.rerankers.huggingface import HuggingFaceRerankerConfig
|
||||
from mem0.configs.rerankers.sentence_transformer import (
|
||||
SentenceTransformerRerankerConfig,
|
||||
)
|
||||
from mem0.configs.rerankers.zero_entropy import ZeroEntropyRerankerConfig
|
||||
from mem0.reranker.huggingface_reranker import HuggingFaceReranker
|
||||
from mem0.reranker.sentence_transformer_reranker import SentenceTransformerReranker
|
||||
|
||||
|
||||
def _docs(n):
|
||||
return [{"memory": f"doc{i}"} for i in range(n)]
|
||||
|
||||
|
||||
class TestCohereFallbackNoMutation:
|
||||
def test_fallback_does_not_mutate_original_documents(self, mock_cohere):
|
||||
module, fake_client = mock_cohere
|
||||
fake_client.rerank.side_effect = RuntimeError("API error")
|
||||
|
||||
reranker = module.CohereReranker(CohereRerankerConfig(api_key="test-key"))
|
||||
documents = _docs(3)
|
||||
result = reranker.rerank("query", documents)
|
||||
|
||||
assert all("rerank_score" not in doc for doc in documents)
|
||||
assert all(doc["rerank_score"] == 0.0 for doc in result)
|
||||
|
||||
|
||||
class TestZeroEntropyFallbackNoMutation:
|
||||
def test_fallback_does_not_mutate_original_documents(self, mock_zero_entropy):
|
||||
module, fake_client = mock_zero_entropy
|
||||
fake_client.models.rerank.side_effect = RuntimeError("API error")
|
||||
|
||||
reranker = module.ZeroEntropyReranker(ZeroEntropyRerankerConfig(api_key="test-key"))
|
||||
documents = _docs(3)
|
||||
result = reranker.rerank("query", documents)
|
||||
|
||||
assert all("rerank_score" not in doc for doc in documents)
|
||||
assert all(doc["rerank_score"] == 0.0 for doc in result)
|
||||
|
||||
|
||||
class TestHuggingFaceFallbackNoMutation:
|
||||
def test_fallback_does_not_mutate_original_documents(self):
|
||||
with (
|
||||
patch("mem0.reranker.huggingface_reranker.AutoTokenizer") as mock_tokenizer_cls,
|
||||
patch("mem0.reranker.huggingface_reranker.AutoModelForSequenceClassification") as mock_model_cls,
|
||||
):
|
||||
mock_tokenizer = MagicMock(side_effect=RuntimeError("tokenizer error"))
|
||||
mock_tokenizer_cls.from_pretrained.return_value = mock_tokenizer
|
||||
mock_model_cls.from_pretrained.return_value = MagicMock()
|
||||
|
||||
reranker = HuggingFaceReranker(HuggingFaceRerankerConfig())
|
||||
documents = _docs(3)
|
||||
result = reranker.rerank("query", documents)
|
||||
|
||||
assert all("rerank_score" not in doc for doc in documents)
|
||||
assert all(doc["rerank_score"] == 0.0 for doc in result)
|
||||
|
||||
|
||||
class TestSentenceTransformerFallbackNoMutation:
|
||||
def test_fallback_does_not_mutate_original_documents(self):
|
||||
with patch("mem0.reranker.sentence_transformer_reranker.CrossEncoder") as mock_cross_encoder_cls:
|
||||
mock_model = MagicMock()
|
||||
mock_model.predict.side_effect = RuntimeError("predict error")
|
||||
mock_cross_encoder_cls.return_value = mock_model
|
||||
|
||||
reranker = SentenceTransformerReranker(SentenceTransformerRerankerConfig())
|
||||
documents = _docs(3)
|
||||
result = reranker.rerank("query", documents)
|
||||
|
||||
assert all("rerank_score" not in doc for doc in documents)
|
||||
assert all(doc["rerank_score"] == 0.0 for doc in result)
|
||||
@@ -1,52 +1,9 @@
|
||||
"""Regression tests for the reranker fallback path honoring ``config.top_k``.
|
||||
|
||||
When the underlying rerank call fails, the reranker falls back to returning the
|
||||
documents in their original order. That fallback must still respect the
|
||||
configured ``top_k`` limit, exactly like the success path does. The HuggingFace
|
||||
and SentenceTransformer rerankers already behave this way; these tests pin the
|
||||
same contract for the Cohere and ZeroEntropy rerankers.
|
||||
"""
|
||||
|
||||
import sys
|
||||
from types import ModuleType
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
"""Regression tests pinning that the reranker fallback path still honors ``config.top_k``."""
|
||||
|
||||
from mem0.configs.rerankers.cohere import CohereRerankerConfig
|
||||
from mem0.configs.rerankers.zero_entropy import ZeroEntropyRerankerConfig
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_cohere(monkeypatch):
|
||||
"""Provide a fake ``cohere`` module so CohereReranker imports/constructs."""
|
||||
fake_cohere = ModuleType("cohere")
|
||||
fake_client = MagicMock()
|
||||
fake_cohere.Client = MagicMock(return_value=fake_client)
|
||||
monkeypatch.setitem(sys.modules, "cohere", fake_cohere)
|
||||
|
||||
import mem0.reranker.cohere_reranker as cohere_reranker
|
||||
|
||||
monkeypatch.setattr(cohere_reranker, "cohere", fake_cohere, raising=False)
|
||||
monkeypatch.setattr(cohere_reranker, "COHERE_AVAILABLE", True, raising=False)
|
||||
return cohere_reranker, fake_client
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_zero_entropy(monkeypatch):
|
||||
"""Provide a fake ``zeroentropy`` module so ZeroEntropyReranker imports."""
|
||||
fake_module = ModuleType("zeroentropy")
|
||||
fake_client = MagicMock()
|
||||
fake_module.ZeroEntropy = MagicMock(return_value=fake_client)
|
||||
monkeypatch.setitem(sys.modules, "zeroentropy", fake_module)
|
||||
|
||||
import mem0.reranker.zero_entropy_reranker as zero_entropy_reranker
|
||||
|
||||
monkeypatch.setattr(zero_entropy_reranker, "ZeroEntropy", fake_module.ZeroEntropy, raising=False)
|
||||
monkeypatch.setattr(zero_entropy_reranker, "ZERO_ENTROPY_AVAILABLE", True, raising=False)
|
||||
return zero_entropy_reranker, fake_client
|
||||
|
||||
|
||||
def _docs(n):
|
||||
return [{"memory": f"doc{i}"} for i in range(n)]
|
||||
|
||||
|
||||
@@ -51,3 +51,12 @@ def test_base_to_provider_without_reasoning_fields_still_builds():
|
||||
|
||||
assert isinstance(built, AnthropicConfig)
|
||||
assert built.model == "claude-3-5-sonnet-20240620"
|
||||
|
||||
|
||||
def test_dict_config_not_mutated_by_kwargs():
|
||||
config = {"model": "gpt-4o-mini"}
|
||||
|
||||
with patch("mem0.utils.factory.load_class", return_value=Mock(return_value=Mock())):
|
||||
LlmFactory.create("openai", config, api_key="secret")
|
||||
|
||||
assert config == {"model": "gpt-4o-mini"}
|
||||
|
||||
Reference in New Issue
Block a user