fix(user profile): update locomo benchmark and update vector based profile code (#225)
Some checks failed
Pre-commit / run (ubuntu-latest) (push) Has been cancelled

* feat(reme): 添加配置选项以启用或禁用个人资料功能

- 在 ReMe 初始化方法中添加 enable_profile 参数,默认值为 True
- 根据 enable_profile 设置决定是否创建 profile 目录和设置 profile_dir
- 在 PersonalSummarizer 中根据 enable_profile 条件性地添加个人资料相关工具
- 在 PersonalRetriever 中根据 enable_profile 条件性地添加 ReadAllProfiles 工具
- 修改 profile_path 属性以在禁用个人资料时返回 None
- 修改 get_profile_handler 方法以在禁用个人资料时返回 None
- 为 enable_profile 参数添加文档说明其用于云向量存储场景

* refactor(benchmark): 重构LongMemEval基准测试中的ReMe实例管理

- 移除未使用的shutil导入
- 将固定的ReMe实例改为每个问题创建独立实例以实现隔离
- 更新LLM配置名称从qwen3-max-think到qwen-max-t
- 修改模型调用逻辑使用正确的model_name参数
- 添加qwen-flash和GPT-4o-mini等新模型配置
- 统一使用"User"作为用户名,通过集合名实现隔离
- 调整并发处理数从4降至1,批处理大小从10增至30
- 每个问题类型采样数从2增至4
- 添加异步上下文管理确保资源正确释放

* reformat 2 files

* refactor(benchmark): 重构长记忆评估中的模型配置

- 将原有的 eval_model_name 替换为专门的 retrieve_model_name 用于检索操作
- 添加对 qwen-max 模型配置的支持
- 更新参数解析器以支持新的检索模型参数
- 修改最大并发数默认值从 1 提升到 4
- 调整样本数量默认值从 4 减少到 1
- 统一模型参数命名规范,区分摘要、检索和评估模型
- 优化内存处理器初始化逻辑,支持独立的检索模型配置

* fix(benchmark): 移除数据路径默认值并设为必填参数

- 将LongMemEval评估脚本中的data_path参数改为必需参数
- 将HaluMem评估脚本中的data_path参数改为必需参数
- 删除了硬编码的默认文件路径配置
- 强制用户显式指定数据集文件路径以避免路径错误

* Update __init__.py

* Update __init__.py

* fix(benchmark): 修复ReMe评估中的模型配置和空值处理问题

- 移除了retrieve_memory调用中不需要的llm_config_name参数
- 修复了长字符串打印的换行格式问题
- 添加了eval_result为空时的初始化处理
- 在accuracy评估中加入了eval_model_name参数传递

* style(benchmark): 格式化模型名称打印输出

- 移除了多行字符串中的换行符和多余空格
- 将模型名称信息合并为单行连续显示
- 保持了原有的打印格式和信息完整性

* docs(readme): 更新文档添加实验结果表格

- 在英文版 README 中添加 🧪 Experiments 章节
- 添加 LoCoMo 和 HaluMem 两个基准测试的结果表格
- 在中文版 README_ZH 中添加 🧪 实验 章节
- 添加 LoCoMo 和 HaluMem 测试集的实验配置说明
- 添加完整的实验数据对比表格和评估协议说明

* docs(readme): 更新文档中的内存系统链接

- 为基于文件的记忆系统添加锚点链接
- 为基于向量库的记忆系统添加锚点链接
- 修复英文文档中的链接格式
- 修复中文文档中的链接格式和空行问题

* docs(readme): update experimental results section in documentation

- Remove outdated experimental data placeholder "Coming soon..."
- Add complete evaluation results for LoCoMo and HaluMem benchmarks
- Include detailed performance metrics tables for all memory methods
- Update experimental settings description with ReMe backbone details
- Align evaluation protocol information with LLM-as-a-Judge approach
- Maintain consistent formatting between English and Chinese documentation

* docs(benchmark): add quick start guides for halumem and longmemeval experiments

- Created HaluMem experiment quick start guide with ReMe integration setup
- Added detailed steps for installing ReMe environment using conda
- Included repository cloning instructions for HaluMem benchmark
- Provided complete command examples for running HaluMem experiments
- Created LongMeMEval quick start guide with data download procedures
- Added wget commands for downloading cleaned dataset files
- Included evaluation script instructions for computing experiment statistics
- Documented parameter configurations for different model types and batch sizes

* docs(longmemeval): update quickstart guide documentation

- Changed project name from Halumem to Longmemeval in title
- Updated description to reference Longmemeval experiments instead of Halumem
- Maintained existing ReMe integration instructions unchanged

* chore(logger): add test comment to logger configuration

- Added test comment in logger utility function
- Removed duplicate log handling by keeping the remove() call

* chore(logger): add test comment to logger configuration

- Added test comment in logger utility function
- Removed duplicate log handling by keeping the remove() call

* feat(core): add file logging capability to application

- Added log_to_file parameter to Application class constructor
- Integrated log_to_file option in logger initialization
- Updated ServiceContext to support file logging configuration
- Modified init_logger function to conditionally enable file logging
- Added log_to_file field to ServiceConfig schema
- Updated ReMe class to include file logging option
- Wrapped file logging setup in conditional check to prevent unnecessary operations

* docs(benchmark): update HaluMem quickstart guide with dataset download instructions

- Replace repository cloning with direct dataset download using curl
- Add commands to download HaluMem-Medium.jsonl and HaluMem-Long.jsonl files
- Include both official Hugging Face and mirror download sources
- Update data path reference from nested directory to local data folder
- Add dataset page link and mirror usage instructions for mainland China access

* feat(memory): add profile retrieval tool and refactor profile management

- Introduce RetrieveProfile tool for fetching specific user profiles
- Refactor ProfileHandler to support both filesystem and vector backends
- Add async methods to ProfileHandler with synchronous fallbacks
- Update PersonalRetriever to support two-stage profile and memory retrieval
- Enhance PersonalSummarizer with improved tool partitioning logic
- Add profile_backend, profile_store_name, and profile_max_capacity configuration options
- Replace direct ProfileHandler imports with get_profile_handler method
- Implement profile search functionality with dedicated prompts and workflows
- Add FileProfileBackend and VectorProfileBackend implementations
- Update base memory tool with new profile configuration parameters

* feat(profile): add custom profile collection name support

- Add profile_collection_name parameter to Application constructor
- Allow custom database collection name for vector profiles instead of default suffix
- Update profile vector store configuration logic to use custom collection name
- Modify _ensure_profile_vector_store_config to handle custom collection names
- Update docstring with detailed parameter descriptions for profile configuration options

* test(history): add single history id acceptance test for multiple mode

- Add test case to verify multiple-mode history lookup accepts a single history_id string
- Create FakeVectorStore stub with minimal implementation for ReadHistory tests
- Return requested history node from vector store mock
- Initialize ReadHistory tool with multiple mode enabled
- Add pylint disable comment for protected access to vector store property

* refactor(memory): update profile handler and vector tools with improved formatting and error handling

- Add module docstring to profiles/__init__.py
- Add pylint disable comments for no-name-in-module and missing-function-docstring
- Format long error message in ProfileHandler.sync_run method for better readability
- Reformat parameters in ProfileHandler.aadd method to separate lines
- Update model_copy call in reme.py to span multiple lines for better readability
- Format aadd_batch call in update_profile.py to span multiple lines

* 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

* docs(profiles): add comprehensive docstrings for profile backend and handler methods

- Added documentation for get_all_sync, get_by_sync, delete_sync, delete_all_sync methods
- Documented add_sync and add_batch_sync functionality with deduping behavior
- Added docstrings for update_sync and search_sync operations
- Updated ProfileHandler.format_node method with proper documentation
- Refactored private _format_node to public format_node method
- Added comprehensive documentation for profile vector handler operations
- Documented _vector_profile_matches, _get_by_profile_id, _get_by_profile_key helper methods
- Added docstrings for retrieve_profile functionality and formatting methods
This commit is contained in:
Zhouwk 2026-04-30 10:19:36 +08:00 committed by GitHub
parent e0d0e3e568
commit 72eabfa858
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 693 additions and 5 deletions

View file

@ -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,

View file

@ -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}

View file

@ -0,0 +1,234 @@
"""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]:
"""Load all profile nodes from cache, ordered by ``message_time``."""
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:
"""Return the first node matching ``profile_id`` or ``profile_key``."""
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:
"""Remove one id, many ids, or none; returns bool, count, or 0/false if nothing removed."""
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:
"""Clear every cached profile for this target; returns how many were stored."""
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:
"""Append a profile row, replacing any existing row with the same key."""
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]:
"""Insert many profiles in one write, deduping by key against existing rows."""
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:
"""Update fields for ``profile_id``; return ``None`` if that id is missing."""
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]:
"""Simple substring/token match over key and content, best matches first."""
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)

View 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."""

View file

@ -115,7 +115,8 @@ class ProfileHandler:
return await self.backend.search(query=query, limit=limit)
@staticmethod
def _format_node(node: MemoryNode, add_profile_id: bool = False, add_history_id: bool = False) -> str:
def format_node(node: MemoryNode, add_profile_id: bool = False, add_history_id: bool = False) -> str:
"""Render a profile ``MemoryNode`` as a single-line string for tools/logs."""
parts = []
profile_key = str(node.metadata.get("profile_key", node.when_to_use))
@ -134,7 +135,7 @@ class ProfileHandler:
async def aread_all(self, add_profile_id: bool = False, add_history_id: bool = False) -> str:
nodes = await self.aget_all()
formatted_profiles = [self._format_node(node, add_profile_id, add_history_id) for node in nodes]
formatted_profiles = [self.format_node(node, add_profile_id, add_history_id) for node in nodes]
logger.info(f"Read {len(formatted_profiles)} profiles from {self.cache_key}")
return "\n".join(formatted_profiles).strip()
@ -146,7 +147,7 @@ class ProfileHandler:
add_history_id: bool = False,
) -> tuple[list[MemoryNode], str]:
nodes = await self.asearch(query=query, limit=limit)
formatted_profiles = [self._format_node(node, add_profile_id, add_history_id) for node in nodes]
formatted_profiles = [self.format_node(node, add_profile_id, add_history_id) for node in nodes]
return nodes, "\n".join(formatted_profiles).strip()
def delete(self, profile_id: str | list[str]) -> bool | int:

View file

@ -0,0 +1,245 @@
"""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:
"""Filters shared by all profile rows in the vector collection."""
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:
"""Turn a profile dict into a ``MemoryNode`` for upsert into the vector store."""
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,
)
def _vector_profile_matches(self, memory_node: MemoryNode) -> bool:
"""True if ``memory_node`` belongs to this handler's target and profile kind."""
if memory_node.memory_target != self.memory_target:
return False
if memory_node.memory_type is not MemoryType.IDENTITY:
return False
if memory_node.metadata.get("profile_kind") != self.PROFILE_KIND:
return False
return True
async def _get_by_profile_id(self, profile_id: str) -> MemoryNode | None:
"""Load by vector id and validate filters."""
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 not self._vector_profile_matches(memory_node):
return None
return memory_node
async def _get_by_profile_key(self, profile_key: str) -> MemoryNode | None:
"""Load the single row matching ``profile_key`` under base filters."""
filters = {**self._base_filters(), "profile_key": profile_key}
vector_nodes = await self.vector_store.list(filters=filters, limit=1)
if not vector_nodes:
return None
return MemoryNode.from_vector_node(vector_nodes[0])
async def get_all(self) -> list[MemoryNode]:
"""List every profile row for this memory target, sorted by store."""
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:
"""Return one profile by stable id or by logical profile key."""
if not profile_id and not profile_key:
raise ValueError("Must provide either profile_id or profile_key")
if profile_id:
return await self._get_by_profile_id(profile_id)
return await self._get_by_profile_key(profile_key or "")
async def delete(self, profile_id: str | list[str]) -> bool | int:
"""Delete one id, many ids, or report zero/false when nothing matched."""
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:
"""Remove all profile vectors for this target; returns how many were deleted."""
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]:
"""Upsert many profiles at once (last dict wins per key), then enforce capacity."""
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:
"""Insert or replace a single profile row."""
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:
"""Replace content and key for ``profile_id``; return ``None`` if missing."""
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]:
"""Semantic search with de-duplication across multiple query strings."""
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):
"""Drop oldest rows when count exceeds ``max_capacity``."""
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)

View 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

View 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)