diff --git a/mem0/memory/utils.py b/mem0/memory/utils.py index b22d60a27..eb7e5754a 100644 --- a/mem0/memory/utils.py +++ b/mem0/memory/utils.py @@ -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() diff --git a/tests/memory/test_memory_utils.py b/tests/memory/test_memory_utils.py index 3bab79429..cef93a188 100644 --- a/tests/memory/test_memory_utils.py +++ b/tests/memory/test_memory_utils.py @@ -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" diff --git a/tests/test_memory.py b/tests/test_memory.py index bc95c0cca..326c8be98 100644 --- a/tests/test_memory.py +++ b/tests/test_memory.py @@ -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")