ReMe/reme/tool/memory/memory_handler.py

210 lines
8.3 KiB
Python

from ...core.context import ServiceContext
from ...core.enumeration import MemoryType
from ...core.schema import MemoryNode
from ...core.vector_store import BaseVectorStore
class MemoryHandler:
"""Handler for managing memory nodes in the vector store."""
def __init__(self, memory_target: str, service_context: ServiceContext):
self.memory_target: str = memory_target
self.memory_type: MemoryType = service_context.memory_target_type_mapping[memory_target]
self.vector_store: BaseVectorStore = service_context.vector_stores["default"]
async def add_batch(self, memories: list[dict]) -> list[MemoryNode]:
"""Add multiple memory nodes and return their memory_ids."""
# First, delete existing memory nodes if memory_ids are provided
memory_ids_to_delete = [mem.get("memory_id") for mem in memories if mem.get("memory_id")]
if memory_ids_to_delete:
await self.vector_store.delete(memory_ids_to_delete)
# Create MemoryNode objects
memory_nodes = [
MemoryNode(
memory_type=self.memory_type,
memory_target=self.memory_target,
content=mem.get("content", ""),
when_to_use=mem.get("when_to_use", ""),
message_time=mem.get("message_time", ""),
ref_memory_id=mem.get("ref_memory_id", ""),
author=mem.get("author", ""),
score=mem.get("score", 0.0),
metadata=mem.get("metadata", {}),
)
for mem in memories
]
# Deduplicate memory_nodes by content (keep last occurrence)
memory_dict = {node.content: node for node in memory_nodes}
memory_nodes = list(memory_dict.values())
# Convert to VectorNodes and insert
vector_nodes = [node.to_vector_node() for node in memory_nodes]
await self.vector_store.insert(vector_nodes)
return memory_nodes
async def add(
self,
content: str,
when_to_use: str = "",
message_time: str = "",
ref_memory_id: str = "",
author: str = "",
score: float = 0.0,
**kwargs,
) -> MemoryNode:
"""Add a single memory node and return its memory_id."""
memory_dict = {
"content": content,
"when_to_use": when_to_use,
"message_time": message_time,
"ref_memory_id": ref_memory_id,
"author": author,
"score": score,
"metadata": kwargs,
}
memory_nodes = await self.add_batch([memory_dict])
return memory_nodes[0]
async def delete(self, memory_ids: str | list[str]):
"""Delete multiple memory nodes by their memory_ids."""
# Deduplicate if input is a list
if isinstance(memory_ids, list):
memory_ids = list(dict.fromkeys(memory_ids))
await self.vector_store.delete(memory_ids)
async def delete_all(self):
"""Delete all memory nodes."""
await self.vector_store.delete_all()
async def update_batch(self, updates: list[dict]) -> list[MemoryNode]:
"""Update multiple memory nodes with their memory_ids and new values using delete + add."""
# Deduplicate updates by memory_id (keep last occurrence)
updates_dict = {upd["memory_id"]: upd for upd in updates}
updates = list(updates_dict.values())
memory_ids = list(updates_dict.keys())
# Get existing nodes
vector_nodes = await self.vector_store.get(memory_ids)
if not isinstance(vector_nodes, list):
vector_nodes = [vector_nodes]
# Update and convert back
updated_nodes: list[MemoryNode] = []
for vector_node, update in zip(vector_nodes, updates):
memory_node = MemoryNode.from_vector_node(vector_node)
memory_node.memory_target = self.memory_target
memory_node.memory_type = self.memory_type
if "content" in update:
memory_node.content = update["content"]
if "when_to_use" in update:
memory_node.when_to_use = update["when_to_use"]
if "message_time" in update:
memory_node.message_time = update["message_time"]
if "ref_memory_id" in update:
memory_node.ref_memory_id = update["ref_memory_id"]
if "author" in update:
memory_node.author = update["author"]
if "score" in update:
memory_node.score = update["score"]
if "metadata" in update:
memory_node.metadata.update(update["metadata"])
updated_nodes.append(memory_node)
# Delete old nodes first
await self.vector_store.delete(memory_ids)
# Then add updated nodes
vector_nodes = [node.to_vector_node() for node in updated_nodes]
await self.vector_store.insert(vector_nodes)
return updated_nodes
async def update(
self,
memory_id: str,
content: str | None = None,
when_to_use: str | None = None,
message_time: str | None = None,
ref_memory_id: str | None = None,
author: str | None = None,
score: float | None = None,
**kwargs,
) -> MemoryNode:
"""Update a memory node's content, when_to_use, or other fields."""
update_dict: dict = {"memory_id": memory_id}
if content is not None:
update_dict["content"] = content
if when_to_use is not None:
update_dict["when_to_use"] = when_to_use
if message_time is not None:
update_dict["message_time"] = message_time
if ref_memory_id is not None:
update_dict["ref_memory_id"] = ref_memory_id
if author is not None:
update_dict["author"] = author
if score is not None:
update_dict["score"] = score
if kwargs is not None:
update_dict["metadata"] = kwargs
memory_nodes = await self.update_batch([update_dict])
return memory_nodes[0]
async def search(
self,
query: str | list[str],
limit: int = 5,
filters: dict | None = None,
**kwargs,
) -> list[MemoryNode]:
"""Search for similar memory nodes based on query text."""
filters = filters or {}
filters["memory_type"] = self.memory_type.value
filters["memory_target"] = self.memory_target
# Handle single query
if isinstance(query, str):
vector_nodes = await self.vector_store.search(query, limit=limit, filters=filters, **kwargs)
return [MemoryNode.from_vector_node(node) for node in vector_nodes]
# Handle multiple queries: search each query with the same limit
seen_ids: dict[str, MemoryNode] = {}
for q in query:
vector_nodes = await self.vector_store.search(q, limit=limit, filters=filters, **kwargs)
for vector_node in vector_nodes:
memory_node = MemoryNode.from_vector_node(vector_node)
if memory_node.memory_id not in seen_ids:
seen_ids[memory_node.memory_id] = memory_node
return list(seen_ids.values())
async def batch_search(self, searches: list[dict]) -> list[MemoryNode]:
"""Execute multiple search queries in batch and return deduplicated results."""
seen_ids: dict[str, MemoryNode] = {}
for search_params in searches:
search_result = await self.search(**search_params)
for memory_node in search_result:
if memory_node.memory_id not in seen_ids:
seen_ids[memory_node.memory_id] = memory_node
return list(seen_ids.values())
async def list(
self,
filters: dict | None = None,
limit: int | None = None,
sort_key: str | None = None,
reverse: bool = True,
) -> list[MemoryNode]:
"""List memory nodes with optional filtering and sorting."""
filters = filters or {}
filters["memory_type"] = self.memory_type.value
filters["memory_target"] = self.memory_target
vector_nodes = await self.vector_store.list(filters=filters, limit=limit, sort_key=sort_key, reverse=reverse)
return [MemoryNode.from_vector_node(node) for node in vector_nodes]