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:
@@ -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()
|
||||
|
||||
@@ -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
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user