From 577a5a2feb89581fd72ee991127a6fe1f125b996 Mon Sep 17 00:00:00 2001 From: Anisha Mahuli <63206204+amahuli03@users.noreply.github.com> Date: Wed, 18 Mar 2026 10:49:29 -0400 Subject: [PATCH] fix(oss): normalize malformed LLM fact output before embedding (#4224) Co-authored-by: kartik-mem0 --- mem0/memory/main.py | 3 ++ mem0/memory/utils.py | 29 ++++++++++++++++- tests/test_memory.py | 77 +++++++++++++++++++++++++++++++++++++++++++- 3 files changed, 107 insertions(+), 2 deletions(-) diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 40fd3f2df..ac745e800 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -29,6 +29,7 @@ from mem0.memory.utils import ( ensure_json_instruction, extract_json, get_fact_retrieval_messages, + normalize_facts, parse_messages, parse_vision_messages, process_telemetry_filters, @@ -456,6 +457,7 @@ class Memory(MemoryBase): # Try extracting JSON from response using built-in function extracted_json = extract_json(response) new_retrieved_facts = json.loads(extracted_json)["facts"] + new_retrieved_facts = normalize_facts(new_retrieved_facts) except Exception as e: logger.error(f"Error in new_retrieved_facts: {e}") new_retrieved_facts = [] @@ -1484,6 +1486,7 @@ class AsyncMemory(MemoryBase): # Try extracting JSON from response using built-in function extracted_json = extract_json(response) new_retrieved_facts = json.loads(extracted_json)["facts"] + new_retrieved_facts = normalize_facts(new_retrieved_facts) except Exception as e: logger.error(f"Error in new_retrieved_facts: {e}") new_retrieved_facts = [] diff --git a/mem0/memory/utils.py b/mem0/memory/utils.py index b451c0f6e..0d2cb6749 100644 --- a/mem0/memory/utils.py +++ b/mem0/memory/utils.py @@ -1,12 +1,15 @@ import hashlib +import logging import re from mem0.configs.prompts import ( + AGENT_MEMORY_EXTRACTION_PROMPT, FACT_RETRIEVAL_PROMPT, USER_MEMORY_EXTRACTION_PROMPT, - AGENT_MEMORY_EXTRACTION_PROMPT, ) +logger = logging.getLogger(__name__) + def get_fact_retrieval_messages(message, is_agent_memory=False): """Get fact retrieval messages based on the memory type. @@ -77,6 +80,30 @@ def format_entities(entities): return "\n".join(formatted_lines) +def normalize_facts(raw_facts): + """Normalize LLM-extracted facts to a list of strings. + + Smaller LLMs (e.g. llama3.1:8b) sometimes return facts as objects + like {"fact": "..."} or {"text": "..."} instead of plain strings. + This mirrors the TypeScript FactRetrievalSchema validation. + """ + if not raw_facts: + return [] + normalized = [] + for item in raw_facts: + if isinstance(item, str): + fact = item + elif isinstance(item, dict): + fact = item.get("fact") or item.get("text") + if fact is None: + logger.warning("Unexpected fact shape from LLM, skipping: %s", item) + continue + else: + fact = str(item) + if fact: + normalized.append(fact) + return normalized + def remove_code_blocks(content: str) -> str: """ diff --git a/tests/test_memory.py b/tests/test_memory.py index 39af062ef..ed98e6ab8 100644 --- a/tests/test_memory.py +++ b/tests/test_memory.py @@ -1,9 +1,11 @@ +import json from unittest.mock import MagicMock, patch import pytest from mem0 import Memory from mem0.configs.base import MemoryConfig +from mem0.memory.utils import normalize_facts class MockVectorMemory: @@ -244,4 +246,77 @@ def test_get_all_handles_flat_list_from_postgres(mock_sqlite, mock_llm_factory, assert len(result) == 2 assert result[0]["memory"] == "Memory 1" - assert result[1]["memory"] == "Memory 2" + assert result[1]["memory"] == "Memory 2" + + +@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_add_infer_with_malformed_llm_facts(mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory): + """ + Repro for: 'list' object has no attribute 'replace' on infer=true. + + When an LLM (especially smaller models like llama3.1:8b) returns facts as + objects ({"fact": "..."} or {"text": "..."}) instead of plain strings, + the embedding model's .replace() call crashes with AttributeError. + """ + mock_embedder = MagicMock() + mock_embedder.embed.side_effect = lambda text, action: (_ for _ in ()).throw( + AttributeError("'dict' object has no attribute 'replace'") + ) if not isinstance(text, str) else [0.1, 0.2, 0.3] + mock_embedder_factory.return_value = mock_embedder + + mock_vector_store = MagicMock() + mock_vector_store.search.return_value = [] + mock_vector_factory.return_value = mock_vector_store + + # LLM returns malformed facts: dicts instead of strings + malformed_response = json.dumps({ + "facts": [ + {"fact": "User likes Python"}, + {"text": "User is a developer"}, + ] + }) + mock_llm = MagicMock() + mock_llm.generate_response.return_value = malformed_response + mock_llm_factory.return_value = mock_llm + + mock_sqlite.return_value = MagicMock() + + from mem0.memory.main import Memory as MemoryClass + config = MemoryConfig() + memory = MemoryClass(config) + + # This should NOT raise AttributeError + memory._add_to_vector_store( + messages=[{"role": "user", "content": "I like Python and I'm a developer"}], + metadata={"user_id": "test_user"}, + filters={"user_id": "test_user"}, + infer=True, + ) + + +def test_normalize_facts_plain_strings(): + assert normalize_facts(["fact one", "fact two"]) == ["fact one", "fact two"] + + +def test_normalize_facts_dict_with_fact_key(): + assert normalize_facts([{"fact": "User likes Python"}]) == ["User likes Python"] + + +def test_normalize_facts_dict_with_text_key(): + assert normalize_facts([{"text": "User is a developer"}]) == ["User is a developer"] + + +def test_normalize_facts_mixed(): + raw = [ + "plain string", + {"fact": "from fact key"}, + {"text": "from text key"}, + ] + assert normalize_facts(raw) == ["plain string", "from fact key", "from text key"] + + +def test_normalize_facts_filters_empty_strings(): + assert normalize_facts(["", "valid", ""]) == ["valid"]