This commit is contained in:
Deshraj Yadav
2025-03-25 23:15:35 -07:00
parent e694971423
commit 1b8d326b49

View File

@@ -25,10 +25,10 @@ except ModuleNotFoundError as e:
class Mem0MemoryService(FrameProcessor): class Mem0MemoryService(FrameProcessor):
"""A standalone memory service that integrates with Mem0. """A standalone memory service that integrates with Mem0.
This service intercepts message frames in the pipeline, stores them in Mem0, This service intercepts message frames in the pipeline, stores them in Mem0,
and enhances context with relevant memories before passing them downstream. and enhances context with relevant memories before passing them downstream.
Args: Args:
api_key (str): The API key for accessing Mem0's API api_key (str): The API key for accessing Mem0's API
user_id (str): The user ID to associate with memories in Mem0 user_id (str): The user ID to associate with memories in Mem0
@@ -39,9 +39,7 @@ class Mem0MemoryService(FrameProcessor):
search_limit: int = Field(default=10, ge=1) search_limit: int = Field(default=10, ge=1)
search_threshold: float = Field(default=0.1, ge=0.0, le=1.0) search_threshold: float = Field(default=0.1, ge=0.0, le=1.0)
api_version: str = Field(default="v2") api_version: str = Field(default="v2")
system_prompt: str = Field( system_prompt: str = Field(default="Based on previous conversations, I recall: \n\n")
default="Based on previous conversations, I recall: \n\n"
)
add_as_system_message: bool = Field(default=True) add_as_system_message: bool = Field(default=True)
position: int = Field(default=1) position: int = Field(default=1)
@@ -61,7 +59,7 @@ class Mem0MemoryService(FrameProcessor):
# At least one of user_id, agent_id, or run_id must be provided # At least one of user_id, agent_id, or run_id must be provided
if not any([user_id, agent_id, run_id]): if not any([user_id, agent_id, run_id]):
raise ValueError("At least one of user_id, agent_id, or run_id must be provided") raise ValueError("At least one of user_id, agent_id, or run_id must be provided")
self.user_id = user_id self.user_id = user_id
self.agent_id = agent_id self.agent_id = agent_id
self.run_id = run_id self.run_id = run_id
@@ -82,7 +80,11 @@ class Mem0MemoryService(FrameProcessor):
""" """
try: try:
logger.debug(f"Storing {len(messages)} messages in Mem0") logger.debug(f"Storing {len(messages)} messages in Mem0")
params = {"messages": messages, "metadata": {"platform": "pipecat"}, "output_format": "v1.1"} params = {
"messages": messages,
"metadata": {"platform": "pipecat"},
"output_format": "v1.1",
}
for id in ["user_id", "agent_id", "run_id"]: for id in ["user_id", "agent_id", "run_id"]:
if getattr(self, id): if getattr(self, id):
params[id] = getattr(self, id) params[id] = getattr(self, id)
@@ -93,16 +95,20 @@ class Mem0MemoryService(FrameProcessor):
def _retrieve_memories(self, query: str) -> List[Dict[str, Any]]: def _retrieve_memories(self, query: str) -> List[Dict[str, Any]]:
"""Retrieve relevant memories from Mem0. """Retrieve relevant memories from Mem0.
Args: Args:
query: The query to search for relevant memories query: The query to search for relevant memories
Returns: Returns:
List of relevant memory dictionaries List of relevant memory dictionaries
""" """
try: try:
logger.debug(f"Retrieving memories for query: {query}") logger.debug(f"Retrieving memories for query: {query}")
id_pairs = [("user_id", self.user_id), ("agent_id", self.agent_id), ("run_id", self.run_id)] id_pairs = [
("user_id", self.user_id),
("agent_id", self.agent_id),
("run_id", self.run_id),
]
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 = {"AND": clauses} if clauses else {} filters = {"AND": clauses} if clauses else {}
results = self.memory_client.search( results = self.memory_client.search(
@@ -121,7 +127,7 @@ class Mem0MemoryService(FrameProcessor):
def _enhance_context_with_memories(self, context: OpenAILLMContext, query: str): def _enhance_context_with_memories(self, context: OpenAILLMContext, query: str):
"""Enhance the LLM context with relevant memories. """Enhance the LLM context with relevant memories.
Args: Args:
context: The OpenAILLMContext to enhance context: The OpenAILLMContext to enhance
query: The query to search for relevant memories query: The query to search for relevant memories
@@ -140,7 +146,7 @@ class Mem0MemoryService(FrameProcessor):
memory_text = self.system_prompt memory_text = self.system_prompt
for i, memory in enumerate(memories, 1): for i, memory in enumerate(memories, 1):
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: if self.add_as_system_message:
context.add_message({"role": "system", "content": memory_text}) context.add_message({"role": "system", "content": memory_text})
@@ -151,22 +157,22 @@ class Mem0MemoryService(FrameProcessor):
async def process_frame(self, frame: Frame, direction: FrameDirection): async def process_frame(self, frame: Frame, direction: FrameDirection):
"""Process incoming frames, intercept context frames for memory integration. """Process incoming frames, intercept context frames for memory integration.
Args: Args:
frame: The incoming frame to process frame: The incoming frame to process
direction: The direction of frame flow in the pipeline direction: The direction of frame flow in the pipeline
""" """
await super().process_frame(frame, direction) await super().process_frame(frame, direction)
context = None context = None
messages = None messages = None
if isinstance(frame, OpenAILLMContextFrame): if isinstance(frame, OpenAILLMContextFrame):
context = frame.context context = frame.context
elif isinstance(frame, LLMMessagesFrame): elif isinstance(frame, LLMMessagesFrame):
messages = frame.messages messages = frame.messages
context = OpenAILLMContext.from_messages(messages) context = OpenAILLMContext.from_messages(messages)
if context: if context:
try: try:
# Get the latest user message to use as a query for memory retrieval # Get the latest user message to use as a query for memory retrieval
@@ -180,7 +186,7 @@ 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) 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. Only call this when user message is detected
self._store_messages(context_messages) self._store_messages(context_messages)