Revert "fix(memory): accept llm kwarg in sync Memory.add()/_create_procedural_memory (#5911)" (#5990)

This commit is contained in:
Kartik
2026-06-29 21:27:16 +05:30
committed by GitHub
parent bedf862d64
commit 754034abbc
2 changed files with 4 additions and 103 deletions
+2 -19
View File
@@ -726,7 +726,6 @@ class Memory(MemoryBase):
infer: bool = True,
memory_type: Optional[str] = None,
prompt: Optional[str] = None,
llm=None,
):
"""
Create a new memory.
@@ -803,7 +802,7 @@ class Memory(MemoryBase):
)
if agent_id is not None and memory_type == MemoryType.PROCEDURAL.value:
results = self._create_procedural_memory(messages, metadata=processed_metadata, prompt=prompt, llm=llm)
results = self._create_procedural_memory(messages, metadata=processed_metadata, prompt=prompt)
scale_threshold_notice = detect_scale_threshold_from_add_result(self, results)
if temporal_usage_notice:
display_temporal_usage_notice(self, "sync", "add", *temporal_usage_notice)
@@ -1911,27 +1910,15 @@ class Memory(MemoryBase):
)
return memory_id
def _create_procedural_memory(self, messages, metadata=None, llm=None, prompt=None):
def _create_procedural_memory(self, messages, metadata=None, prompt=None):
"""
Create a procedural memory
Args:
messages (list): List of messages to create a procedural memory from.
metadata (dict): Metadata to create a procedural memory from.
llm (llm, optional): LLM to use for the procedural memory creation. Defaults to None.
prompt (str, optional): Prompt to use for the procedural memory creation. Defaults to None.
"""
if llm is not None:
try:
from langchain_core.messages.utils import (
convert_to_messages, # type: ignore
)
except Exception:
logger.error(
"Import error while loading langchain-core. Please install 'langchain-core' to use procedural memory."
)
raise
logger.info("Creating procedural memory")
parsed_messages = [
@@ -1944,10 +1931,6 @@ class Memory(MemoryBase):
]
try:
if llm is not None:
response = llm.invoke(input=convert_to_messages(parsed_messages))
procedural_memory = remove_code_blocks(response.content)
else:
procedural_memory = self.llm.generate_response(messages=parsed_messages)
procedural_memory = remove_code_blocks(procedural_memory)
except Exception as e:
-82
View File
@@ -1529,85 +1529,3 @@ async def test_async_procedural_memory_langchain_strips_code_blocks(mock_llm_fac
insert_call = memory.vector_store.insert.call_args
stored_data = insert_call[1]["payloads"][0]["data"]
assert "```" not in stored_data
@patch("mem0.memory.main.VectorStoreFactory")
@patch("mem0.memory.main.EmbedderFactory")
@patch("mem0.memory.main.LlmFactory")
def test_sync_procedural_memory_accepts_llm_kwarg_and_strips_code_blocks(
mock_llm_factory, mock_emb, mock_vs
):
"""Regression #5911: sync Memory.add()/Memory._create_procedural_memory must
accept the same `llm` kwarg as the async path and route through llm.invoke()."""
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 Memory
config = MemoryConfig()
memory = Memory(config)
memory.vector_store = MagicMock()
memory.vector_store.insert = MagicMock()
mock_langchain_llm = MagicMock()
mock_response = MagicMock()
mock_response.content = '```json\n{"key": "value"}\n```'
mock_langchain_llm.invoke.return_value = mock_response
messages = [{"role": "user", "content": "test"}]
metadata = {"user_id": "test_user"}
memory._create_procedural_memory(messages, metadata=metadata, llm=mock_langchain_llm)
mock_langchain_llm.invoke.assert_called_once()
insert_call = memory.vector_store.insert.call_args
stored_data = insert_call[1]["payloads"][0]["data"]
assert "```" not in stored_data
@patch("mem0.memory.main.VectorStoreFactory")
@patch("mem0.memory.main.EmbedderFactory")
@patch("mem0.memory.main.LlmFactory")
def test_sync_memory_add_forwards_llm_to_procedural_memory(
mock_llm_factory, mock_emb, mock_vs
):
"""Regression #5911: Memory.add() must forward `llm` to _create_procedural_memory."""
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 Memory
config = MemoryConfig()
memory = Memory(config)
memory.vector_store = MagicMock()
memory.vector_store.insert = MagicMock()
mock_langchain_llm = MagicMock()
mock_response = MagicMock()
mock_response.content = "procedural summary"
mock_langchain_llm.invoke.return_value = mock_response
forwarded = {"called": False}
original_create = memory._create_procedural_memory
def spy(*args, **kwargs):
forwarded["called"] = True
forwarded["llm"] = kwargs.get("llm")
return original_create(*args, **kwargs)
memory._create_procedural_memory = spy
memory.add(
[{"role": "user", "content": "Always verify inputs before processing."}],
agent_id="agent-1",
memory_type="procedural_memory",
llm=mock_langchain_llm,
)
assert forwarded["called"] is True
assert forwarded["llm"] is mock_langchain_llm
mock_langchain_llm.invoke.assert_called_once()