Files
mem0/tests/memory/test_memory_utils.py
T

113 lines
4.8 KiB
Python

import pytest
from unittest.mock import Mock
from mem0.memory.utils import parse_vision_messages, remove_spaces_from_entities, sanitize_relationship_for_cypher
class TestParseVisionMessages:
def test_multimodal_list_without_llm_extracts_text(self):
messages = [
{"role": "user", "content": [
{"type": "text", "text": "What is this?"},
{"type": "image_url", "image_url": {"url": "https://example.com/img.png"}},
]},
]
result = parse_vision_messages(messages, llm=None)
assert len(result) == 1
assert result[0]["role"] == "user"
assert result[0]["content"] == "What is this?"
def test_image_dict_without_llm_is_skipped(self):
messages = [
{"role": "user", "content": {"type": "image_url", "image_url": {"url": "https://example.com/img.png"}}},
{"role": "user", "content": "hello"},
]
result = parse_vision_messages(messages, llm=None)
assert len(result) == 1
assert result[0]["content"] == "hello"
def test_multimodal_with_llm_calls_generate_response(self):
mock_llm = Mock()
mock_llm.generate_response.return_value = "A photo of a cat"
messages = [
{"role": "user", "content": [
{"type": "text", "text": "Describe this"},
{"type": "image_url", "image_url": {"url": "https://example.com/cat.png"}},
]},
]
result = parse_vision_messages(messages, llm=mock_llm, vision_details="auto")
assert result[0]["content"] == "A photo of a cat"
mock_llm.generate_response.assert_called_once()
def test_image_only_list_without_llm_is_skipped(self):
messages = [
{"role": "user", "content": [
{"type": "image_url", "image_url": {"url": "https://example.com/img.png"}},
]},
]
result = parse_vision_messages(messages, llm=None)
assert result == []
def test_plain_text_messages_pass_through(self):
messages = [
{"role": "system", "content": "You are helpful."},
{"role": "user", "content": "Hello"},
]
result = parse_vision_messages(messages, llm=None)
assert result == messages
class TestRemoveSpacesFromEntities:
"""
Covers behavior used by Neo4j, Memgraph (sanitize_relationship=True),
Kuzu, and Neptune (sanitize_relationship=False). All backends delegate here.
"""
@pytest.mark.parametrize(
"sanitize",
[True, False],
ids=["cypher_sanitized", "plain"],
)
def test_filters_empty_and_incomplete_dicts(self, sanitize):
mixed = [
{},
{"source": "a"},
{"source": "a", "relationship": "r"},
{"source": "x", "relationship": "rel", "destination": "y"},
]
out = remove_spaces_from_entities(mixed, sanitize_relationship=sanitize)
assert len(out) == 1
assert out[0]["source"] == "x"
assert out[0]["destination"] == "y"
@pytest.mark.parametrize("sanitize", [True, False])
def test_all_empty_returns_empty(self, sanitize):
assert remove_spaces_from_entities([{}, {}, {}], sanitize_relationship=sanitize) == []
def test_skips_non_dict_entries(self):
assert remove_spaces_from_entities([None, "not-a-dict", 123, {"source": "a", "relationship": "r", "destination": "b"}]) == [
{"source": "a", "relationship": "r", "destination": "b"}
]
def test_sanitize_true_relationship_uses_sanitizer(self):
"""Neo4j / Memgraph path: special characters mapped via sanitize_relationship_for_cypher."""
entities = [{"source": "A", "relationship": "x/y", "destination": "B"}]
out = remove_spaces_from_entities(entities, sanitize_relationship=True)
assert out[0]["relationship"] == sanitize_relationship_for_cypher("x/y".lower().replace(" ", "_"))
def test_sanitize_false_relationship_plain_only(self):
"""Kuzu / Neptune path: only lowercase and spaces to underscores."""
entities = [{"source": "A", "relationship": "Works At", "destination": "B Co"}]
out = remove_spaces_from_entities(entities, sanitize_relationship=False)
assert out[0]["relationship"] == "works_at"
assert out[0]["source"] == "a"
assert out[0]["destination"] == "b_co"
def test_sanitize_true_vs_false_slash_in_relationship(self):
"""Slash is rewritten when sanitizing (Cypher path); kept as-is for plain path."""
base = {"source": "s", "relationship": "a/b", "destination": "d"}
t = remove_spaces_from_entities([dict(base)], sanitize_relationship=True)[0]["relationship"]
f = remove_spaces_from_entities([dict(base)], sanitize_relationship=False)[0]["relationship"]
assert t == sanitize_relationship_for_cypher("a/b")
assert f == "a/b"