mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-29 01:41:38 +00:00
feat(profiles): add profile management system with file and vector storage backends
- Add FileProfileBackend for filesystem-based profile persistence
- Add VectorProfileBackend for vector store-based profile management
- Create abstract BaseProfileBackend interface for profile operations
- Implement ProfileVectorHandler for vector-backed profile storage
- Add RetrieveProfile tool for semantic profile retrieval
- Update eval_reme.py to use user_message_s2 for retriever prompt
- Modify eval_reme.yaml to use {profiles} instead of {user_profile}
- Implement complete CRUD operations for profile management
- Add batch operations for efficient profile handling
- Include search functionality with semantic matching capabilities
- Add capacity limits and automatic cleanup for profile storage
This commit is contained in:
parent
6e431adaa0
commit
03259729c1
7 changed files with 656 additions and 2 deletions
|
|
@ -643,7 +643,7 @@ class LocomoEvaluator:
|
|||
},
|
||||
"personal_retriever": {
|
||||
"prompt_dict": {
|
||||
"user_message": self.retriever_prompt,
|
||||
"user_message_s2": self.retriever_prompt,
|
||||
},
|
||||
"params": {
|
||||
"return_memory_nodes": True,
|
||||
|
|
|
|||
|
|
@ -132,7 +132,7 @@ user_message_retrieve: |
|
|||
You are a Memory Retrieval Agent specialized in retrieving {memory_type} memories about {memory_target}.
|
||||
|
||||
## User Profile
|
||||
{user_profile}
|
||||
{profiles}
|
||||
|
||||
## User Question
|
||||
{context}
|
||||
|
|
|
|||
226
reme/memory/vector_tools/profiles/file_profile_backend.py
Normal file
226
reme/memory/vector_tools/profiles/file_profile_backend.py
Normal file
|
|
@ -0,0 +1,226 @@
|
|||
"""Filesystem-backed profile storage."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from .profile_backend import BaseProfileBackend
|
||||
from ....core.enumeration import MemoryType
|
||||
from ....core.schema import MemoryNode
|
||||
from ....core.utils import CacheHandler, deduplicate_memories
|
||||
|
||||
|
||||
class FileProfileBackend(BaseProfileBackend):
|
||||
"""Persist user profiles in local JSONL cache files."""
|
||||
|
||||
def __init__(self, profile_path: str | Path, memory_target: str, max_capacity: int = 50):
|
||||
super().__init__(memory_target=memory_target, max_capacity=max_capacity)
|
||||
self.cache_key: str = self.memory_target.replace(" ", "_").lower()
|
||||
self.cache_handler: CacheHandler = CacheHandler(profile_path)
|
||||
|
||||
def _load_nodes(self) -> list[MemoryNode]:
|
||||
cached_data = self.cache_handler.load(self.cache_key, auto_clean=False)
|
||||
if not cached_data:
|
||||
return []
|
||||
return [MemoryNode(**data) for data in cached_data]
|
||||
|
||||
def _save_nodes(self, nodes: list[MemoryNode], apply_limits: bool = True):
|
||||
if apply_limits:
|
||||
nodes = deduplicate_memories(nodes)
|
||||
|
||||
if len(nodes) > self.max_capacity:
|
||||
sorted_nodes = sorted(nodes, key=lambda n: n.message_time)
|
||||
removed_count = len(sorted_nodes) - self.max_capacity
|
||||
nodes = sorted_nodes[removed_count:]
|
||||
logger.info(
|
||||
f"Capacity limit reached: removed {removed_count} oldest profiles "
|
||||
f"(kept {len(nodes)}/{self.max_capacity})",
|
||||
)
|
||||
|
||||
nodes_data = [node.model_dump(exclude_none=True) for node in nodes]
|
||||
self.cache_handler.save(self.cache_key, nodes_data)
|
||||
logger.info(f"Saved {len(nodes)} profiles to {self.cache_key}")
|
||||
|
||||
def get_all_sync(self) -> list[MemoryNode]:
|
||||
nodes = self._load_nodes()
|
||||
nodes.sort(key=lambda n: n.message_time)
|
||||
return nodes
|
||||
|
||||
def get_by_sync(self, *, profile_id: str | None = None, profile_key: str | None = None) -> MemoryNode | None:
|
||||
if not profile_id and not profile_key:
|
||||
raise ValueError("Must provide either profile_id or profile_key")
|
||||
|
||||
for node in self._load_nodes():
|
||||
if profile_id and node.memory_id == profile_id:
|
||||
return node
|
||||
if profile_key and node.when_to_use == profile_key:
|
||||
return node
|
||||
return None
|
||||
|
||||
def delete_sync(self, profile_id: str | list[str]) -> bool | int:
|
||||
nodes = self._load_nodes()
|
||||
original_count = len(nodes)
|
||||
|
||||
if isinstance(profile_id, list):
|
||||
profile_ids_set = set(profile_id)
|
||||
nodes = [n for n in nodes if n.memory_id not in profile_ids_set]
|
||||
deleted_count = original_count - len(nodes)
|
||||
if deleted_count == 0:
|
||||
logger.warning(f"No profiles found to delete from {len(profile_id)} IDs")
|
||||
return 0
|
||||
|
||||
self._save_nodes(nodes, apply_limits=False)
|
||||
logger.info(f"Batch deleted {deleted_count} profiles")
|
||||
return deleted_count
|
||||
|
||||
nodes = [n for n in nodes if n.memory_id != profile_id]
|
||||
if len(nodes) == original_count:
|
||||
logger.warning(f"Profile {profile_id} not found")
|
||||
return False
|
||||
|
||||
self._save_nodes(nodes, apply_limits=False)
|
||||
logger.info(f"Deleted profile {profile_id}")
|
||||
return True
|
||||
|
||||
def delete_all_sync(self) -> int:
|
||||
nodes = self._load_nodes()
|
||||
count = len(nodes)
|
||||
self._save_nodes([], apply_limits=False)
|
||||
logger.info(f"Deleted all {count} profiles")
|
||||
return count
|
||||
|
||||
def add_sync(self, message_time: str, profile_key: str, profile_value: str, ref_memory_id: str = "") -> MemoryNode:
|
||||
nodes = self._load_nodes()
|
||||
|
||||
new_node = MemoryNode(
|
||||
memory_type=MemoryType.PERSONAL,
|
||||
memory_target=self.memory_target,
|
||||
when_to_use=profile_key,
|
||||
content=profile_value,
|
||||
message_time=message_time,
|
||||
ref_memory_id=ref_memory_id,
|
||||
)
|
||||
|
||||
original_count = len(nodes)
|
||||
nodes = [n for n in nodes if n.when_to_use != profile_key]
|
||||
if len(nodes) < original_count:
|
||||
logger.info(f"Removed {original_count - len(nodes)} duplicate profile(s) with key: {profile_key}")
|
||||
|
||||
nodes.append(new_node)
|
||||
self._save_nodes(nodes)
|
||||
logger.info(f"Added profile: {profile_key}={profile_value}")
|
||||
return new_node
|
||||
|
||||
def add_batch_sync(self, profiles: list[dict], ref_memory_id: str = "") -> list[MemoryNode]:
|
||||
if not profiles:
|
||||
return []
|
||||
|
||||
nodes = self._load_nodes()
|
||||
new_nodes = [
|
||||
MemoryNode(
|
||||
memory_type=MemoryType.PERSONAL,
|
||||
memory_target=self.memory_target,
|
||||
when_to_use=p.get("profile_key", ""),
|
||||
content=p.get("profile_value", ""),
|
||||
message_time=p.get("message_time", ""),
|
||||
ref_memory_id=ref_memory_id,
|
||||
)
|
||||
for p in profiles
|
||||
]
|
||||
|
||||
new_keys = {n.when_to_use for n in new_nodes}
|
||||
original_count = len(nodes)
|
||||
nodes = [n for n in nodes if n.when_to_use not in new_keys]
|
||||
if len(nodes) < original_count:
|
||||
logger.info(f"Removed {original_count - len(nodes)} duplicate profile(s) with matching keys")
|
||||
|
||||
nodes.extend(new_nodes)
|
||||
self._save_nodes(nodes)
|
||||
logger.info(f"Batch added {len(new_nodes)} profiles")
|
||||
return new_nodes
|
||||
|
||||
def update_sync(
|
||||
self,
|
||||
profile_id: str,
|
||||
message_time: str,
|
||||
profile_key: str,
|
||||
profile_value: str,
|
||||
) -> MemoryNode | None:
|
||||
nodes = self._load_nodes()
|
||||
target_node = None
|
||||
for node in nodes:
|
||||
if node.memory_id == profile_id:
|
||||
node.when_to_use = profile_key
|
||||
node.content = profile_value
|
||||
node.message_time = message_time
|
||||
target_node = node
|
||||
break
|
||||
|
||||
if target_node is None:
|
||||
logger.warning(f"Profile {profile_id} not found")
|
||||
return None
|
||||
|
||||
self._save_nodes(nodes, apply_limits=False)
|
||||
logger.info(f"Updated profile {profile_id}: {profile_key}={profile_value}")
|
||||
return target_node
|
||||
|
||||
def search_sync(self, query: str | list[str], limit: int = 5) -> list[MemoryNode]:
|
||||
queries = [query] if isinstance(query, str) else query
|
||||
query_terms = [q.strip().lower() for q in queries if q and q.strip()]
|
||||
if not query_terms:
|
||||
return []
|
||||
|
||||
scored_nodes = []
|
||||
for node in self.get_all_sync():
|
||||
profile_key = str(node.metadata.get("profile_key", node.when_to_use)).lower()
|
||||
haystack = f"{profile_key}: {node.content}".lower()
|
||||
score = 0
|
||||
for term in query_terms:
|
||||
if term in haystack:
|
||||
score += len(term) + 10
|
||||
else:
|
||||
token_hits = sum(1 for token in term.split() if token and token in haystack)
|
||||
score += token_hits
|
||||
|
||||
if score > 0:
|
||||
node.score = float(score)
|
||||
scored_nodes.append(node)
|
||||
|
||||
scored_nodes.sort(key=lambda n: (n.score, n.message_time), reverse=True)
|
||||
return scored_nodes[:limit]
|
||||
|
||||
async def get_all(self) -> list[MemoryNode]:
|
||||
return self.get_all_sync()
|
||||
|
||||
async def get_by(self, *, profile_id: str | None = None, profile_key: str | None = None) -> MemoryNode | None:
|
||||
return self.get_by_sync(profile_id=profile_id, profile_key=profile_key)
|
||||
|
||||
async def delete(self, profile_id: str | list[str]) -> bool | int:
|
||||
return self.delete_sync(profile_id)
|
||||
|
||||
async def delete_all(self) -> int:
|
||||
return self.delete_all_sync()
|
||||
|
||||
async def add(
|
||||
self,
|
||||
message_time: str,
|
||||
profile_key: str,
|
||||
profile_value: str,
|
||||
ref_memory_id: str = "",
|
||||
) -> MemoryNode:
|
||||
return self.add_sync(message_time, profile_key, profile_value, ref_memory_id)
|
||||
|
||||
async def add_batch(self, profiles: list[dict], ref_memory_id: str = "") -> list[MemoryNode]:
|
||||
return self.add_batch_sync(profiles, ref_memory_id)
|
||||
|
||||
async def update(
|
||||
self,
|
||||
profile_id: str,
|
||||
message_time: str,
|
||||
profile_key: str,
|
||||
profile_value: str,
|
||||
) -> MemoryNode | None:
|
||||
return self.update_sync(profile_id, message_time, profile_key, profile_value)
|
||||
|
||||
async def search(self, query: str | list[str], limit: int = 5) -> list[MemoryNode]:
|
||||
return self.search_sync(query, limit)
|
||||
51
reme/memory/vector_tools/profiles/profile_backend.py
Normal file
51
reme/memory/vector_tools/profiles/profile_backend.py
Normal file
|
|
@ -0,0 +1,51 @@
|
|||
"""Profile backend abstractions."""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
from ....core.schema import MemoryNode
|
||||
|
||||
|
||||
class BaseProfileBackend(ABC):
|
||||
"""Abstract interface for profile storage backends."""
|
||||
|
||||
def __init__(self, memory_target: str, max_capacity: int = 50):
|
||||
self.memory_target = memory_target
|
||||
self.max_capacity = max_capacity
|
||||
|
||||
@abstractmethod
|
||||
async def get_all(self) -> list[MemoryNode]:
|
||||
"""Return all profile rows for the current user."""
|
||||
|
||||
@abstractmethod
|
||||
async def get_by(self, *, profile_id: str | None = None, profile_key: str | None = None) -> MemoryNode | None:
|
||||
"""Return one profile row by id or key."""
|
||||
|
||||
@abstractmethod
|
||||
async def delete(self, profile_id: str | list[str]) -> bool | int:
|
||||
"""Delete one or more profile rows."""
|
||||
|
||||
@abstractmethod
|
||||
async def delete_all(self) -> int:
|
||||
"""Delete all profile rows for the current user."""
|
||||
|
||||
@abstractmethod
|
||||
async def add(self, message_time: str, profile_key: str, profile_value: str, ref_memory_id: str = "") -> MemoryNode:
|
||||
"""Add a single profile row."""
|
||||
|
||||
@abstractmethod
|
||||
async def add_batch(self, profiles: list[dict], ref_memory_id: str = "") -> list[MemoryNode]:
|
||||
"""Add multiple profile rows."""
|
||||
|
||||
@abstractmethod
|
||||
async def update(
|
||||
self,
|
||||
profile_id: str,
|
||||
message_time: str,
|
||||
profile_key: str,
|
||||
profile_value: str,
|
||||
) -> MemoryNode | None:
|
||||
"""Update one profile row."""
|
||||
|
||||
@abstractmethod
|
||||
async def search(self, query: str | list[str], limit: int = 5) -> list[MemoryNode]:
|
||||
"""Search profile rows relevant to the query."""
|
||||
220
reme/memory/vector_tools/profiles/profile_vector_handler.py
Normal file
220
reme/memory/vector_tools/profiles/profile_vector_handler.py
Normal file
|
|
@ -0,0 +1,220 @@
|
|||
"""Vector-backed handler for bounded user profiles."""
|
||||
|
||||
import hashlib
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from ....core import ServiceContext
|
||||
from ....core.enumeration import MemoryType
|
||||
from ....core.schema import MemoryNode
|
||||
from ....core.vector_store import BaseVectorStore
|
||||
|
||||
|
||||
class ProfileVectorHandler:
|
||||
"""Manage profile rows stored in a dedicated vector collection."""
|
||||
|
||||
PROFILE_KIND = "profile"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
memory_target: str,
|
||||
service_context: ServiceContext,
|
||||
vector_store_name: str = "profile",
|
||||
max_capacity: int = 50,
|
||||
):
|
||||
self.memory_target = memory_target
|
||||
self.service_context = service_context
|
||||
self.vector_store_name = vector_store_name
|
||||
self.max_capacity = max_capacity
|
||||
self.vector_store: BaseVectorStore = service_context.vector_stores[vector_store_name]
|
||||
|
||||
@staticmethod
|
||||
def build_retrieval_text(profile_key: str, profile_value: str) -> str:
|
||||
"""Build the text that will be embedded for semantic profile retrieval."""
|
||||
return f"{profile_key}: {profile_value}".strip(": ")
|
||||
|
||||
def build_profile_id(self, profile_key: str) -> str:
|
||||
"""Build a stable id from user and key."""
|
||||
hash_obj = hashlib.sha256(f"{self.memory_target}\n{profile_key}".encode("utf-8"))
|
||||
return hash_obj.hexdigest()[:16]
|
||||
|
||||
def _base_filters(self) -> dict:
|
||||
return {
|
||||
"memory_type": MemoryType.IDENTITY.value,
|
||||
"memory_target": self.memory_target,
|
||||
"profile_kind": self.PROFILE_KIND,
|
||||
}
|
||||
|
||||
def _build_profile_node(self, profile: dict, ref_memory_id: str = "") -> MemoryNode:
|
||||
profile_key = profile.get("profile_key", "").strip()
|
||||
profile_value = profile.get("profile_value", "").strip()
|
||||
message_time = profile.get("message_time", "")
|
||||
ref_id = profile.get("ref_memory_id", ref_memory_id)
|
||||
metadata = dict(profile.get("metadata", {}))
|
||||
metadata.update(
|
||||
{
|
||||
"profile_key": profile_key,
|
||||
"profile_kind": self.PROFILE_KIND,
|
||||
"profile_backend": "vector",
|
||||
},
|
||||
)
|
||||
return MemoryNode(
|
||||
memory_id=self.build_profile_id(profile_key),
|
||||
memory_type=MemoryType.IDENTITY,
|
||||
memory_target=self.memory_target,
|
||||
when_to_use=self.build_retrieval_text(profile_key, profile_value),
|
||||
content=profile_value,
|
||||
message_time=message_time,
|
||||
ref_memory_id=ref_id,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
async def get_all(self) -> list[MemoryNode]:
|
||||
vector_nodes = await self.vector_store.list(
|
||||
filters=self._base_filters(),
|
||||
sort_key="message_time",
|
||||
reverse=False,
|
||||
)
|
||||
return [MemoryNode.from_vector_node(node) for node in vector_nodes]
|
||||
|
||||
async def get_by(self, *, profile_id: str | None = None, profile_key: str | None = None) -> MemoryNode | None:
|
||||
if not profile_id and not profile_key:
|
||||
raise ValueError("Must provide either profile_id or profile_key")
|
||||
|
||||
if profile_id:
|
||||
try:
|
||||
vector_node = await self.vector_store.get(profile_id)
|
||||
except KeyError:
|
||||
logger.warning(f"Profile {profile_id} not found in vector store")
|
||||
return None
|
||||
if vector_node is None:
|
||||
logger.warning(f"Profile {profile_id} not found in vector store")
|
||||
return None
|
||||
memory_node = MemoryNode.from_vector_node(vector_node)
|
||||
if memory_node.memory_target != self.memory_target:
|
||||
return None
|
||||
if memory_node.memory_type is not MemoryType.IDENTITY:
|
||||
return None
|
||||
if memory_node.metadata.get("profile_kind") != self.PROFILE_KIND:
|
||||
return None
|
||||
return memory_node
|
||||
|
||||
profile_key = profile_key or ""
|
||||
vector_nodes = await self.vector_store.list(filters={**self._base_filters(), "profile_key": profile_key}, limit=1)
|
||||
if not vector_nodes:
|
||||
return None
|
||||
return MemoryNode.from_vector_node(vector_nodes[0])
|
||||
|
||||
async def delete(self, profile_id: str | list[str]) -> bool | int:
|
||||
if isinstance(profile_id, list):
|
||||
profile_ids = list(dict.fromkeys(pid for pid in profile_id if pid))
|
||||
if not profile_ids:
|
||||
return 0
|
||||
existing_nodes = []
|
||||
for pid in profile_ids:
|
||||
node = await self.get_by(profile_id=pid)
|
||||
if node is not None:
|
||||
existing_nodes.append(node)
|
||||
if not existing_nodes:
|
||||
return 0
|
||||
await self.vector_store.delete([node.memory_id for node in existing_nodes])
|
||||
return len(existing_nodes)
|
||||
|
||||
existing_node = await self.get_by(profile_id=profile_id)
|
||||
if existing_node is None:
|
||||
return False
|
||||
await self.vector_store.delete(existing_node.memory_id)
|
||||
return True
|
||||
|
||||
async def delete_all(self) -> int:
|
||||
nodes = await self.get_all()
|
||||
if not nodes:
|
||||
return 0
|
||||
await self.vector_store.delete([node.memory_id for node in nodes])
|
||||
return len(nodes)
|
||||
|
||||
async def add_batch(self, profiles: list[dict], ref_memory_id: str = "") -> list[MemoryNode]:
|
||||
if not profiles:
|
||||
return []
|
||||
|
||||
deduped_profiles: dict[str, dict] = {}
|
||||
for profile in profiles:
|
||||
profile_key = profile.get("profile_key", "").strip()
|
||||
if not profile_key:
|
||||
continue
|
||||
deduped_profiles[profile_key] = profile
|
||||
|
||||
new_nodes = [self._build_profile_node(profile, ref_memory_id=ref_memory_id) for profile in deduped_profiles.values()]
|
||||
if not new_nodes:
|
||||
return []
|
||||
|
||||
await self.vector_store.delete([node.memory_id for node in new_nodes])
|
||||
await self.vector_store.insert([node.to_vector_node() for node in new_nodes])
|
||||
await self.enforce_capacity()
|
||||
return new_nodes
|
||||
|
||||
async def add(self, message_time: str, profile_key: str, profile_value: str, ref_memory_id: str = "") -> MemoryNode:
|
||||
nodes = await self.add_batch(
|
||||
[
|
||||
{
|
||||
"message_time": message_time,
|
||||
"profile_key": profile_key,
|
||||
"profile_value": profile_value,
|
||||
},
|
||||
],
|
||||
ref_memory_id=ref_memory_id,
|
||||
)
|
||||
return nodes[0]
|
||||
|
||||
async def update(
|
||||
self,
|
||||
profile_id: str,
|
||||
message_time: str,
|
||||
profile_key: str,
|
||||
profile_value: str,
|
||||
) -> MemoryNode | None:
|
||||
existing_node = await self.get_by(profile_id=profile_id)
|
||||
if existing_node is None:
|
||||
return None
|
||||
|
||||
new_node = self._build_profile_node(
|
||||
{
|
||||
"message_time": message_time,
|
||||
"profile_key": profile_key,
|
||||
"profile_value": profile_value,
|
||||
"ref_memory_id": existing_node.ref_memory_id,
|
||||
"metadata": existing_node.metadata,
|
||||
},
|
||||
)
|
||||
|
||||
if existing_node.memory_id != new_node.memory_id:
|
||||
await self.vector_store.delete(existing_node.memory_id)
|
||||
else:
|
||||
await self.vector_store.delete(new_node.memory_id)
|
||||
|
||||
await self.vector_store.insert(new_node.to_vector_node())
|
||||
await self.enforce_capacity()
|
||||
return new_node
|
||||
|
||||
async def search(self, query: str | list[str], limit: int = 5) -> list[MemoryNode]:
|
||||
queries = [query] if isinstance(query, str) else query
|
||||
seen_nodes: dict[str, MemoryNode] = {}
|
||||
for item in queries:
|
||||
if not item or not item.strip():
|
||||
continue
|
||||
vector_nodes = await self.vector_store.search(item, limit=limit, filters=self._base_filters())
|
||||
for vector_node in vector_nodes:
|
||||
memory_node = MemoryNode.from_vector_node(vector_node)
|
||||
seen_nodes[memory_node.memory_id] = memory_node
|
||||
nodes = list(seen_nodes.values())
|
||||
nodes.sort(key=lambda node: (node.score, node.message_time), reverse=True)
|
||||
return nodes[:limit]
|
||||
|
||||
async def enforce_capacity(self):
|
||||
nodes = await self.get_all()
|
||||
overflow = len(nodes) - self.max_capacity
|
||||
if overflow <= 0:
|
||||
return
|
||||
|
||||
to_delete = [node.memory_id for node in nodes[:overflow]]
|
||||
await self.vector_store.delete(to_delete)
|
||||
102
reme/memory/vector_tools/profiles/retrieve_profile.py
Normal file
102
reme/memory/vector_tools/profiles/retrieve_profile.py
Normal file
|
|
@ -0,0 +1,102 @@
|
|||
"""Retrieve relevant profile rows."""
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from .profile_handler import ProfileHandler
|
||||
from ..base_memory_tool import BaseMemoryTool
|
||||
from ....core.schema import MemoryNode, ToolCall
|
||||
|
||||
|
||||
class RetrieveProfile(BaseMemoryTool):
|
||||
"""Tool to retrieve relevant profiles using the configured backend."""
|
||||
|
||||
def __init__(self, top_k: int = 5, enable_memory_target: bool = False, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.top_k = top_k
|
||||
self.enable_memory_target = enable_memory_target
|
||||
|
||||
def _build_query_parameters(self) -> dict:
|
||||
properties = {
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "query",
|
||||
},
|
||||
}
|
||||
required = ["query"]
|
||||
if self.enable_memory_target:
|
||||
properties["memory_target"] = {
|
||||
"type": "string",
|
||||
"description": "memory_target",
|
||||
}
|
||||
required.append("memory_target")
|
||||
return {
|
||||
"type": "object",
|
||||
"properties": properties,
|
||||
"required": required,
|
||||
}
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": "Retrieve relevant user profiles using semantic matching.",
|
||||
"parameters": self._build_query_parameters(),
|
||||
},
|
||||
)
|
||||
|
||||
def _build_multiple_tool_call(self) -> ToolCall:
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": "Retrieve relevant user profiles using semantic matching.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query_items": {
|
||||
"type": "array",
|
||||
"description": "List of query items.",
|
||||
"items": self._build_query_parameters(),
|
||||
},
|
||||
},
|
||||
"required": ["query_items"],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self):
|
||||
if self.enable_multiple:
|
||||
query_items = self.context.get("query_items", [])
|
||||
else:
|
||||
query_items = [self.context]
|
||||
|
||||
queries_by_target: dict[str, list[str]] = {}
|
||||
for item in query_items:
|
||||
target = item["memory_target"] if self.enable_memory_target else self.memory_target
|
||||
queries_by_target.setdefault(target, []).append(item["query"])
|
||||
|
||||
profile_nodes: list[MemoryNode] = []
|
||||
for target, queries in queries_by_target.items():
|
||||
profile_handler = self.get_profile_handler(target)
|
||||
nodes, _ = await profile_handler.aretrieve(
|
||||
query=queries,
|
||||
limit=self.top_k,
|
||||
add_profile_id=True,
|
||||
add_history_id=True,
|
||||
)
|
||||
profile_nodes.extend(nodes)
|
||||
|
||||
seen_ids = {node.memory_id: node for node in self.retrieved_nodes if node.memory_id}
|
||||
new_nodes = []
|
||||
for node in profile_nodes:
|
||||
if node.memory_id not in seen_ids:
|
||||
seen_ids[node.memory_id] = node
|
||||
new_nodes.append(node)
|
||||
self.retrieved_nodes.extend(new_nodes)
|
||||
|
||||
if not new_nodes:
|
||||
output = "No new profiles found."
|
||||
else:
|
||||
output = "\n".join(
|
||||
[ProfileHandler._format_node(node, add_profile_id=True, add_history_id=True) for node in new_nodes],
|
||||
)
|
||||
|
||||
logger.info(f"Retrieved {len(profile_nodes)} profiles, {len(new_nodes)} new after deduplication")
|
||||
return output
|
||||
55
reme/memory/vector_tools/profiles/vector_profile_backend.py
Normal file
55
reme/memory/vector_tools/profiles/vector_profile_backend.py
Normal file
|
|
@ -0,0 +1,55 @@
|
|||
"""Vector-backed profile storage."""
|
||||
|
||||
from .profile_backend import BaseProfileBackend
|
||||
from .profile_vector_handler import ProfileVectorHandler
|
||||
from ....core import ServiceContext
|
||||
from ....core.schema import MemoryNode
|
||||
|
||||
|
||||
class VectorProfileBackend(BaseProfileBackend):
|
||||
"""Persist user profiles in a dedicated vector store."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
memory_target: str,
|
||||
service_context: ServiceContext,
|
||||
vector_store_name: str = "profile",
|
||||
max_capacity: int = 50,
|
||||
):
|
||||
super().__init__(memory_target=memory_target, max_capacity=max_capacity)
|
||||
self.handler = ProfileVectorHandler(
|
||||
memory_target=memory_target,
|
||||
service_context=service_context,
|
||||
vector_store_name=vector_store_name,
|
||||
max_capacity=max_capacity,
|
||||
)
|
||||
|
||||
async def get_all(self) -> list[MemoryNode]:
|
||||
return await self.handler.get_all()
|
||||
|
||||
async def get_by(self, *, profile_id: str | None = None, profile_key: str | None = None) -> MemoryNode | None:
|
||||
return await self.handler.get_by(profile_id=profile_id, profile_key=profile_key)
|
||||
|
||||
async def delete(self, profile_id: str | list[str]) -> bool | int:
|
||||
return await self.handler.delete(profile_id)
|
||||
|
||||
async def delete_all(self) -> int:
|
||||
return await self.handler.delete_all()
|
||||
|
||||
async def add(self, message_time: str, profile_key: str, profile_value: str, ref_memory_id: str = "") -> MemoryNode:
|
||||
return await self.handler.add(message_time, profile_key, profile_value, ref_memory_id)
|
||||
|
||||
async def add_batch(self, profiles: list[dict], ref_memory_id: str = "") -> list[MemoryNode]:
|
||||
return await self.handler.add_batch(profiles, ref_memory_id)
|
||||
|
||||
async def update(
|
||||
self,
|
||||
profile_id: str,
|
||||
message_time: str,
|
||||
profile_key: str,
|
||||
profile_value: str,
|
||||
) -> MemoryNode | None:
|
||||
return await self.handler.update(profile_id, message_time, profile_key, profile_value)
|
||||
|
||||
async def search(self, query: str | list[str], limit: int = 5) -> list[MemoryNode]:
|
||||
return await self.handler.search(query, limit)
|
||||
Loading…
Add table
Reference in a new issue