fix(mem0): make memory service non-blocking and use position parameter

Move blocking Mem0 API calls off the event loop using asyncio.to_thread().
Store messages as a fire-and-forget background task via create_task() since
the result is not needed. Insert memory messages at the configured position
in the context instead of always appending.

Closes #1741
This commit is contained in:
Mark Backman
2026-03-26 12:42:36 -04:00
parent 6a6ee8d563
commit 6a87d0e87d

View File

@@ -11,6 +11,7 @@ and retrieve conversational memories, enhancing LLM context with relevant
historical information. historical information.
""" """
import asyncio
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional
from loguru import logger from loguru import logger
@@ -112,9 +113,12 @@ class Mem0MemoryService(FrameProcessor):
self.last_query = None self.last_query = None
logger.info(f"Initialized Mem0MemoryService with {user_id=}, {agent_id=}, {run_id=}") logger.info(f"Initialized Mem0MemoryService with {user_id=}, {agent_id=}, {run_id=}")
def _store_messages(self, messages: List[Dict[str, Any]]): async def _store_messages(self, messages: List[Dict[str, Any]]):
"""Store messages in Mem0. """Store messages in Mem0.
Runs the blocking Mem0 API call in a background thread to avoid
blocking the event loop.
Args: Args:
messages: List of message dictionaries to store in memory. messages: List of message dictionaries to store in memory.
""" """
@@ -131,14 +135,16 @@ class Mem0MemoryService(FrameProcessor):
if isinstance(self.memory_client, Memory): if isinstance(self.memory_client, Memory):
del params["output_format"] del params["output_format"]
# Note: You can run this in background to avoid blocking the conversation await asyncio.to_thread(lambda: self.memory_client.add(**params))
self.memory_client.add(**params)
except Exception as e: except Exception as e:
logger.error(f"Error storing messages in Mem0: {e}") logger.error(f"Error storing messages in Mem0: {e}")
def _retrieve_memories(self, query: str) -> List[Dict[str, Any]]: async def _retrieve_memories(self, query: str) -> List[Dict[str, Any]]:
"""Retrieve relevant memories from Mem0. """Retrieve relevant memories from Mem0.
Runs the blocking Mem0 API call in a background thread to avoid
blocking the event loop.
Args: Args:
query: The query to search for relevant memories. query: The query to search for relevant memories.
@@ -156,7 +162,7 @@ class Mem0MemoryService(FrameProcessor):
"limit": self.search_limit, "limit": self.search_limit,
} }
params = {k: v for k, v in params.items() if v is not None} params = {k: v for k, v in params.items() if v is not None}
results = self.memory_client.search(**params) results = await asyncio.to_thread(lambda: self.memory_client.search(**params))
else: else:
id_pairs = [ id_pairs = [
("user_id", self.user_id), ("user_id", self.user_id),
@@ -165,13 +171,15 @@ class Mem0MemoryService(FrameProcessor):
] ]
clauses = [{name: value} for name, value in id_pairs if value is not None] clauses = [{name: value} for name, value in id_pairs if value is not None]
filters = {"OR": clauses} if clauses else {} filters = {"OR": clauses} if clauses else {}
results = self.memory_client.search( results = await asyncio.to_thread(
query=query, lambda: self.memory_client.search(
filters=filters, query=query,
version=self.api_version, filters=filters,
top_k=self.search_limit, version=self.api_version,
threshold=self.search_threshold, top_k=self.search_limit,
output_format="v1.1", threshold=self.search_threshold,
output_format="v1.1",
)
) )
logger.debug(f"Retrieved {len(results)} memories from Mem0") logger.debug(f"Retrieved {len(results)} memories from Mem0")
@@ -180,7 +188,9 @@ class Mem0MemoryService(FrameProcessor):
logger.error(f"Error retrieving memories from Mem0: {e}") logger.error(f"Error retrieving memories from Mem0: {e}")
return [] return []
def _enhance_context_with_memories(self, context: LLMContext | OpenAILLMContext, query: str): async def _enhance_context_with_memories(
self, context: LLMContext | OpenAILLMContext, query: str
):
"""Enhance the LLM context with relevant memories. """Enhance the LLM context with relevant memories.
Args: Args:
@@ -193,7 +203,7 @@ class Mem0MemoryService(FrameProcessor):
self.last_query = query self.last_query = query
memories = self._retrieve_memories(query) memories = await self._retrieve_memories(query)
if not memories: if not memories:
return return
@@ -203,11 +213,14 @@ class Mem0MemoryService(FrameProcessor):
memory_text += f"{i}. {memory.get('memory', '')}\n\n" memory_text += f"{i}. {memory.get('memory', '')}\n\n"
# Add memories as a system message or user message based on configuration # Add memories as a system message or user message based on configuration
if self.add_as_system_message: role = "system" if self.add_as_system_message else "user"
context.add_message({"role": "system", "content": memory_text}) memory_message = {"role": role, "content": memory_text}
else:
# Add as a user message that provides context messages = context.get_messages()
context.add_message({"role": "user", "content": memory_text}) position = max(0, min(self.position, len(messages)))
messages.insert(position, memory_message)
context.set_messages(messages)
logger.debug(f"Enhanced context with {len(memories)} memories") logger.debug(f"Enhanced context with {len(memories)} memories")
async def process_frame(self, frame: Frame, direction: FrameDirection): async def process_frame(self, frame: Frame, direction: FrameDirection):
@@ -241,9 +254,9 @@ class Mem0MemoryService(FrameProcessor):
if latest_user_message: if latest_user_message:
# Enhance context with memories before passing it downstream # Enhance context with memories before passing it downstream
self._enhance_context_with_memories(context, latest_user_message) await self._enhance_context_with_memories(context, latest_user_message)
# Store the conversation in Mem0. Only call this when user message is detected # Store the conversation in Mem0 as a background task
self._store_messages(context_messages) self.create_task(self._store_messages(context_messages), name="mem0_store")
# If we received an LLMMessagesFrame, create a new one with the enhanced messages # If we received an LLMMessagesFrame, create a new one with the enhanced messages
if messages is not None: if messages is not None: