Fix handle malformed entity dicts and None LLM response in memgraph_memory (#4238)
This commit is contained in:
@@ -219,6 +219,7 @@ class MemoryGraph:
|
||||
if tool_call["name"] != "extract_entities":
|
||||
continue
|
||||
for item in tool_call["arguments"]["entities"]:
|
||||
if "entity" in item and "entity_type" in item:
|
||||
entity_type_map[item["entity"]] = item["entity_type"]
|
||||
except Exception as e:
|
||||
logger.exception(
|
||||
@@ -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}")
|
||||
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user