diff --git a/mem0/memory/memgraph_memory.py b/mem0/memory/memgraph_memory.py index 3ad1c4198..14a2e5119 100644 --- a/mem0/memory/memgraph_memory.py +++ b/mem0/memory/memgraph_memory.py @@ -219,7 +219,8 @@ class MemoryGraph: if tool_call["name"] != "extract_entities": continue for item in tool_call["arguments"]["entities"]: - entity_type_map[item["entity"]] = item["entity_type"] + if "entity" in item and "entity_type" in item: + entity_type_map[item["entity"]] = item["entity_type"] except Exception as e: logger.exception( f"Error in search tool: {e}, llm_provider={self.llm_provider}, search_results={search_results}" @@ -263,8 +264,8 @@ class MemoryGraph: ) entities = [] - if extracted_entities["tool_calls"]: - entities = extracted_entities["tool_calls"][0]["arguments"]["entities"] + if extracted_entities and extracted_entities.get("tool_calls"): + entities = extracted_entities["tool_calls"][0].get("arguments", {}).get("entities", []) entities = self._remove_spaces_from_entities(entities) logger.debug(f"Extracted entities: {entities}") diff --git a/tests/memory/test_memgraph_memory.py b/tests/memory/test_memgraph_memory.py new file mode 100644 index 000000000..4bf51a400 --- /dev/null +++ b/tests/memory/test_memgraph_memory.py @@ -0,0 +1,96 @@ +from unittest.mock import MagicMock, Mock, patch + +# langchain_memgraph and rank_bm25 are optional deps — mock them so tests run without install +_memgraph_mock = Mock() +patch.dict("sys.modules", { + "langchain_memgraph": _memgraph_mock, + "langchain_memgraph.graphs": _memgraph_mock, + "langchain_memgraph.graphs.memgraph": _memgraph_mock, + "rank_bm25": Mock(), +}).start() + +from mem0.memory.memgraph_memory import MemoryGraph as MemgraphMemoryGraph # noqa: E402 +MemoryGraph = MemgraphMemoryGraph + + +def _make_instance(): + with patch.object(MemoryGraph, "__init__", return_value=None): + instance = MemoryGraph.__new__(MemoryGraph) + instance.llm_provider = "openai" + instance.llm = MagicMock() + instance.embedding_model = MagicMock() + instance.config = MagicMock() + instance.config.graph_store.custom_prompt = None + return instance + + +class TestRetrieveNodesFromData: + """Tests for _retrieve_nodes_from_data in MemoryGraph.""" + + def test_normal_entities_extracted(self): + instance = _make_instance() + instance.llm.generate_response.return_value = { + "tool_calls": [{"name": "extract_entities", "arguments": {"entities": [ + {"entity": "Alice", "entity_type": "person"}, + {"entity": "hiking", "entity_type": "activity"}, + ]}}] + } + result = instance._retrieve_nodes_from_data("Alice loves hiking", {"user_id": "u1"}) + assert result == {"alice": "person", "hiking": "activity"} + + def test_malformed_entity_missing_entity_type_is_skipped(self): + """LLM returns entity dict without entity_type — should skip it, keep valid ones. + Reproduces the exact data from issue #4055.""" + instance = _make_instance() + instance.llm.generate_response.return_value = { + "tool_calls": [{"name": "extract_entities", "arguments": {"entities": [ + {"entity": "matrix multiplication", "entity_type": "task"}, + {"entity": "task"}, + {"entity": "ReLU", "entity_type": "task"}, + ]}}] + } + result = instance._retrieve_nodes_from_data("some text", {"user_id": "u1"}) + assert "matrix_multiplication" in result + assert "relu" in result + assert "task" not in result + + def test_none_tool_calls_returns_empty(self): + instance = _make_instance() + instance.llm.generate_response.return_value = {"tool_calls": None} + result = instance._retrieve_nodes_from_data("hello world", {"user_id": "u1"}) + assert result == {} + + +class TestEstablishNodesRelationsFromData: + """Tests for _establish_nodes_relations_from_data in MemoryGraph.""" + + def test_none_response_does_not_crash(self): + """openai_structured returns None when no relations found — must not crash. + Exact crash from issue #4055: TypeError: 'NoneType' object is not subscriptable.""" + instance = _make_instance() + instance.llm.generate_response.return_value = None + result = instance._establish_nodes_relations_from_data( + "Hello world", {"user_id": "u1"}, {} + ) + assert result == [] + + def test_empty_tool_calls_returns_empty(self): + instance = _make_instance() + instance.llm.generate_response.return_value = {"tool_calls": []} + result = instance._establish_nodes_relations_from_data( + "Hello world", {"user_id": "u1"}, {} + ) + assert result == [] + + def test_valid_entities_returned(self): + instance = _make_instance() + instance.llm.generate_response.return_value = { + "tool_calls": [{"name": "add_entities", "arguments": {"entities": [ + {"source": "alice", "relationship": "loves", "destination": "hiking"} + ]}}] + } + result = instance._establish_nodes_relations_from_data( + "Alice loves hiking", {"user_id": "u1"}, {"alice": "person"} + ) + assert len(result) == 1 + assert result[0]["source"] == "alice"