fix(memory): handle block-list content in remove_code_blocks

A custom LangChain LLM passed to AsyncMemory._create_procedural_memory
returns AIMessage.content, which is str | list[str | dict]. Providers
that emit content blocks hit content.strip() on a list and raised
AttributeError.

remove_code_blocks already normalized None, so the list case is handled
in the same guard rather than at the single call site.

Closes #6150
This commit is contained in:
kartik-mem0
2026-08-13 17:23:37 +05:30
parent 96d45b78c7
commit d5b041e942
3 changed files with 63 additions and 4 deletions
+8 -2
View File
@@ -1,7 +1,7 @@
import hashlib
import logging
import re
from typing import Any, Dict, List
from typing import Any, Dict, List, Union
from mem0.configs.prompts import (
AGENT_MEMORY_EXTRACTION_PROMPT,
@@ -112,7 +112,7 @@ def normalize_facts(raw_facts):
return normalized
def remove_code_blocks(content: str) -> str:
def remove_code_blocks(content: Union[str, List[Union[str, Dict[str, Any]]], None]) -> str:
"""
Removes enclosing code block markers ```[language] and ``` from a given string.
@@ -120,9 +120,15 @@ def remove_code_blocks(content: str) -> str:
- The function uses a regex pattern to match code blocks that may start with ``` followed by an optional language tag (letters or numbers) and end with ```.
- If a code block is detected, it returns only the inner content, stripping out the markers.
- If no code block markers are found, the original content is returned as-is.
- Block-style content (a LangChain message whose `.content` is a list of text
blocks) is flattened to its concatenated text before matching.
"""
if content is None:
return ""
if isinstance(content, list):
content = "".join(
part if isinstance(part, str) else part.get("text", "") for part in content if isinstance(part, (str, dict))
)
pattern = r"^```[a-zA-Z0-9]*\n([\s\S]*?)\n```$"
match = re.match(pattern, content.strip())
match_res=match.group(1).strip() if match else content.strip()
+18
View File
@@ -192,3 +192,21 @@ class TestProcessTelemetryFilters:
class TestRemoveCodeBlocks:
def test_none_content_returns_empty_string(self):
assert remove_code_blocks(None) == ""
def test_block_list_content_is_flattened(self):
assert remove_code_blocks([{"type": "text", "text": "step one"}]) == "step one"
def test_block_list_content_strips_code_fences(self):
content = [{"type": "text", "text": '```json\n{"a": 1}\n```'}]
assert remove_code_blocks(content) == '{"a": 1}'
def test_block_list_content_joins_multiple_blocks(self):
content = [{"type": "text", "text": "step one. "}, {"type": "text", "text": "step two"}]
assert remove_code_blocks(content) == "step one. step two"
def test_block_list_content_ignores_non_text_blocks(self):
content = [{"type": "thinking", "thinking": "hmm"}, {"type": "text", "text": "answer"}]
assert remove_code_blocks(content) == "answer"
def test_plain_string_blocks_are_supported(self):
assert remove_code_blocks(["step one. ", "step two"]) == "step one. step two"
+37 -2
View File
@@ -1235,9 +1235,10 @@ class TestHybridSearchWarning:
def test_warning_for_store_without_keyword_search(
self, mock_vs_factory, mock_emb, mock_llm, mock_sqlite, _cap, caplog
):
from mem0.vector_stores.base import VectorStoreBase
import logging
from mem0.vector_stores.base import VectorStoreBase
class StoreWithoutKeywordSearch(VectorStoreBase):
def create_col(self, *a, **kw): pass
def insert(self, *a, **kw): pass
@@ -1271,9 +1272,10 @@ class TestHybridSearchWarning:
def test_no_warning_for_store_with_keyword_search(
self, mock_vs_factory, mock_emb, mock_llm, mock_sqlite, _cap, caplog
):
from mem0.vector_stores.base import VectorStoreBase
import logging
from mem0.vector_stores.base import VectorStoreBase
class StoreWithKeywordSearch(VectorStoreBase):
def keyword_search(self, query, top_k=5, filters=None):
return []
@@ -1536,6 +1538,39 @@ async def test_async_procedural_memory_langchain_strips_code_blocks(mock_llm_fac
assert "```" not in stored_data
@pytest.mark.asyncio
@patch("mem0.memory.main.VectorStoreFactory")
@patch("mem0.memory.main.EmbedderFactory")
@patch("mem0.memory.main.LlmFactory")
async def test_async_procedural_memory_langchain_block_list_content(mock_llm_factory, mock_emb, mock_vs):
"""Regression #6150: block-list message content must not raise AttributeError."""
mock_vs.return_value = MagicMock()
mock_emb.return_value = MagicMock()
mock_emb.return_value.embed.return_value = [0.1] * 1536
mock_llm_factory.return_value = MagicMock()
from mem0.memory.main import AsyncMemory
config = MemoryConfig()
memory = AsyncMemory(config)
memory.vector_store = MagicMock()
memory.vector_store.insert = MagicMock()
mock_langchain_llm = MagicMock()
mock_response = MagicMock()
mock_response.content = [{"type": "text", "text": '```json\n{"key": "value"}\n```'}]
mock_langchain_llm.invoke.return_value = mock_response
messages = [{"role": "user", "content": "test"}]
metadata = {"user_id": "test_user"}
await memory._create_procedural_memory(messages, metadata=metadata, llm=mock_langchain_llm)
insert_call = memory.vector_store.insert.call_args
stored_data = insert_call[1]["payloads"][0]["data"]
assert stored_data == '{"key": "value"}'
@pytest.mark.asyncio
@patch("mem0.memory.main.VectorStoreFactory")
@patch("mem0.memory.main.EmbedderFactory")