From 72eabfa858bd331d021310435f77bb718931ad47 Mon Sep 17 00:00:00 2001 From: Zhouwk <57825291+nitwtog@users.noreply.github.com> Date: Thu, 30 Apr 2026 10:19:36 +0800 Subject: [PATCH 1/3] fix(user profile): update locomo benchmark and update vector based profile code (#225) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * 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 --- benchmark/locomo/eval_reme.py | 2 +- benchmark/locomo/eval_reme.yaml | 2 +- .../profiles/file_profile_backend.py | 234 +++++++++++++++++ .../vector_tools/profiles/profile_backend.py | 51 ++++ .../vector_tools/profiles/profile_handler.py | 7 +- .../profiles/profile_vector_handler.py | 245 ++++++++++++++++++ .../vector_tools/profiles/retrieve_profile.py | 102 ++++++++ .../profiles/vector_profile_backend.py | 55 ++++ 8 files changed, 693 insertions(+), 5 deletions(-) create mode 100644 reme/memory/vector_tools/profiles/file_profile_backend.py create mode 100644 reme/memory/vector_tools/profiles/profile_backend.py create mode 100644 reme/memory/vector_tools/profiles/profile_vector_handler.py create mode 100644 reme/memory/vector_tools/profiles/retrieve_profile.py create mode 100644 reme/memory/vector_tools/profiles/vector_profile_backend.py diff --git a/benchmark/locomo/eval_reme.py b/benchmark/locomo/eval_reme.py index 82d91031..bb424aee 100644 --- a/benchmark/locomo/eval_reme.py +++ b/benchmark/locomo/eval_reme.py @@ -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, diff --git a/benchmark/locomo/eval_reme.yaml b/benchmark/locomo/eval_reme.yaml index 113431b3..5228884a 100644 --- a/benchmark/locomo/eval_reme.yaml +++ b/benchmark/locomo/eval_reme.yaml @@ -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} diff --git a/reme/memory/vector_tools/profiles/file_profile_backend.py b/reme/memory/vector_tools/profiles/file_profile_backend.py new file mode 100644 index 00000000..e6226459 --- /dev/null +++ b/reme/memory/vector_tools/profiles/file_profile_backend.py @@ -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) diff --git a/reme/memory/vector_tools/profiles/profile_backend.py b/reme/memory/vector_tools/profiles/profile_backend.py new file mode 100644 index 00000000..8ffcb9bc --- /dev/null +++ b/reme/memory/vector_tools/profiles/profile_backend.py @@ -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.""" diff --git a/reme/memory/vector_tools/profiles/profile_handler.py b/reme/memory/vector_tools/profiles/profile_handler.py index 0df3707a..145caa6e 100644 --- a/reme/memory/vector_tools/profiles/profile_handler.py +++ b/reme/memory/vector_tools/profiles/profile_handler.py @@ -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: diff --git a/reme/memory/vector_tools/profiles/profile_vector_handler.py b/reme/memory/vector_tools/profiles/profile_vector_handler.py new file mode 100644 index 00000000..f03e6a96 --- /dev/null +++ b/reme/memory/vector_tools/profiles/profile_vector_handler.py @@ -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) diff --git a/reme/memory/vector_tools/profiles/retrieve_profile.py b/reme/memory/vector_tools/profiles/retrieve_profile.py new file mode 100644 index 00000000..38337bd8 --- /dev/null +++ b/reme/memory/vector_tools/profiles/retrieve_profile.py @@ -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 diff --git a/reme/memory/vector_tools/profiles/vector_profile_backend.py b/reme/memory/vector_tools/profiles/vector_profile_backend.py new file mode 100644 index 00000000..3147cc92 --- /dev/null +++ b/reme/memory/vector_tools/profiles/vector_profile_backend.py @@ -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) From f42cf60706611afdfa57dd3c09bca106c93c368f Mon Sep 17 00:00:00 2001 From: lichen2015 Date: Fri, 8 May 2026 17:12:11 +0800 Subject: [PATCH 2/3] add zvec vector/file store (#218) --- README.md | 2 +- README_ZH.md | 2 +- docs/vector_store_api_guide.md | 33 +- reme/core/file_store/__init__.py | 3 + reme/core/file_store/zvec_file_store.py | 573 +++++++++++++ reme/core/vector_store/__init__.py | 3 + reme/core/vector_store/zvec_vector_store.py | 809 +++++++++++++++++ tests/test_file_store.py | 45 +- tests/test_vector_store.py | 39 +- tests/test_zvec_vector_store.py | 906 ++++++++++++++++++++ tests/vector/test_reme_vector.py | 2 +- 11 files changed, 2409 insertions(+), 8 deletions(-) create mode 100644 reme/core/file_store/zvec_file_store.py create mode 100644 reme/core/vector_store/zvec_vector_store.py create mode 100644 tests/test_zvec_vector_store.py diff --git a/README.md b/README.md index f2f413c0..b4bb48c4 100644 --- a/README.md +++ b/README.md @@ -506,7 +506,7 @@ async def main(): "dimensions": 1024, }, default_vector_store_config={ - "backend": "local", # Supports local/chroma/qdrant/elasticsearch/obvec + "backend": "local", # Supports local/chroma/qdrant/elasticsearch/obvec/zvec }, ) await reme.start() diff --git a/README_ZH.md b/README_ZH.md index 7e6ccd58..210a11ce 100644 --- a/README_ZH.md +++ b/README_ZH.md @@ -486,7 +486,7 @@ async def main(): "dimensions": 1024, }, default_vector_store_config={ - "backend": "local", # 支持 local/chroma/qdrant/elasticsearch/obvec + "backend": "local", # 支持 local/chroma/qdrant/elasticsearch/obvec/zvec }, ) await reme.start() diff --git a/docs/vector_store_api_guide.md b/docs/vector_store_api_guide.md index b0a06eab..89881ef2 100644 --- a/docs/vector_store_api_guide.md +++ b/docs/vector_store_api_guide.md @@ -34,6 +34,7 @@ FlowLLM provides multiple Vector Store implementations tailored to different use - **ChromaVectorStore** ([source code](https://github.com/flowllm-ai/flowllm/blob/main/flowllm/core/vector_store/chroma_vector_store.py)): Based on ChromaDB, providing persistent storage and metadata filtering capabilities. - **EsVectorStore** ([source code](https://github.com/flowllm-ai/flowllm/blob/main/flowllm/core/vector_store/es_vector_store.py)): Built on Elasticsearch, enabling powerful combined full-text and vector search functionalities. - **ObVecVectorStore** ([source code](https://github.com/agentscope-ai/ReMe/blob/main/reme/core/vector_store/obvec_vector_store.py)): Uses [pyobvector](https://pypi.org/project/pyobvector/) against **OceanBase** or **seekdb** (MySQL-compatible wire protocol). Suitable when you already run OceanBase/seekdb or need a SQL-native vector table with HNSW-style ANN search and JSON metadata filters. +- **ZvecVectorStore** ([source code](https://github.com/agentscope-ai/ReMe/blob/main/reme/core/vector_store/zvec_vector_store.py)): Built on zvec, a high-performance local vector database with strong-schema support and HNSW indexing. Suitable for single-machine deployments requiring fast vector search. All Vector Store implementations inherit from **BaseVectorStore** ([source code](https://github.com/agentscope-ai/ReMe/blob/main/reme/core/vector_store/base_vector_store.py)) in ReMe, ensuring a consistent interface specification. @@ -130,6 +131,11 @@ docker run -d --name reme_seekdb -p 2881:2881 -e ROOT_PASSWORD= python tests/test_vector_store.py --obvec ``` +### ZvecVectorStore Configuration + +- **db_path**: Local storage path for persistent mode (required). +- **dimension**: Dimensionality of the embedding vectors (default: `1024`). +- **distance**: Distance metric — supports `cosine`, `l2`, `ip` (default: `cosine`). ## Configuration File Examples @@ -151,7 +157,7 @@ vector_store.default.params.= ### Configuration Field Descriptions -- **`backend`** (required): Vector store backend type. Options: `local`, `memory`, `chroma`, `qdrant`, `elasticsearch`, `obvec`. +- **`backend`** (required): Vector store backend type. Options: `local`, `memory`, `chroma`, `qdrant`, `elasticsearch`, `obvec`, `zvec`. - **`embedding_model`** (required): Name of the embedding model configuration, referencing the `embedding_model` section. - **`params`** (optional): Dictionary of backend-specific parameters passed to the vector store constructor. @@ -347,6 +353,30 @@ vector_stores.default.password=your-root-password ReMe service YAML uses the key `vector_stores` (plural); CLI overrides use the same nested paths. +#### 7. ZvecVectorStore Configuration + +Persistent local storage based on zvec with HNSW indexing and strong-schema support. + +**Implementation**: [`reme/core/vector_store/zvec_vector_store.py`](https://github.com/agentscope-ai/ReMe/blob/main/reme/core/vector_store/zvec_vector_store.py) + +```yaml +vector_store: + default: + backend: zvec + embedding_model: default + params: + db_path: "./zvec_vector_store" # Local storage path (required) + dimension: 1024 # Vector dimension (optional; default: 1024) + distance: "cosine" # Distance metric (optional; default: cosine; options: cosine, l2, ip) +``` + +```shell +vector_store.default.backend=zvec +vector_store.default.params.db_path=./zvec_vector_store +vector_store.default.params.dimension=1024 +vector_store.default.params.distance=cosine +``` + ### Complete Configuration Example Below is a complete `default.yaml` example including both embedding model and vector store configurations: @@ -405,6 +435,7 @@ Two types of metadata filtering are supported: - **Development & Testing**: Use MemoryVectorStore or LocalVectorStore—no additional services required. - **Small-Scale Applications**: Use LocalVectorStore or ChromaVectorStore for simplicity and ease of use. - **Production Environments**: Use QdrantVectorStore, EsVectorStore, or ObVecVectorStore (OceanBase/seekdb) for high performance and scalability, depending on your existing infrastructure. +- **High-Performance Local Search**: Use ZvecVectorStore for single-machine deployments requiring fast HNSW-based vector search with local persistence. - **Hybrid Search**: Use EsVectorStore to combine vector search with full-text search capabilities. - **OceanBase / seekdb**: Use ObVecVectorStore when you standardize on pyobvector and SQL-accessible vector tables. diff --git a/reme/core/file_store/__init__.py b/reme/core/file_store/__init__.py index 1358df52..a8457406 100644 --- a/reme/core/file_store/__init__.py +++ b/reme/core/file_store/__init__.py @@ -9,6 +9,7 @@ from .base_file_store import BaseFileStore from .chroma_file_store import ChromaFileStore from .local_file_store import LocalFileStore from .sqlite_file_store import SqliteFileStore +from .zvec_file_store import ZvecFileStore from ..registry_factory import R __all__ = [ @@ -16,8 +17,10 @@ __all__ = [ "ChromaFileStore", "LocalFileStore", "SqliteFileStore", + "ZvecFileStore", ] R.file_stores.register("sqlite")(SqliteFileStore) R.file_stores.register("chroma")(ChromaFileStore) R.file_stores.register("local")(LocalFileStore) +R.file_stores.register("zvec")(ZvecFileStore) diff --git a/reme/core/file_store/zvec_file_store.py b/reme/core/file_store/zvec_file_store.py new file mode 100644 index 00000000..3d162819 --- /dev/null +++ b/reme/core/file_store/zvec_file_store.py @@ -0,0 +1,573 @@ +"""Zvec storage backend for file store.""" + +from __future__ import annotations + +import json +import time +from pathlib import Path +from typing import Any + +from .base_file_store import BaseFileStore +from ..enumeration import MemorySource +from ..schema import FileMetadata, MemoryChunk, MemorySearchResult +from ..utils import get_logger + +logger = get_logger() + +_ZVEC_IMPORT_ERROR: Exception | None = None + +try: + import zvec # type: ignore[import-untyped] + from zvec import ( + CollectionOption, + CollectionSchema, + DataType, + Doc, + FieldSchema, + HnswIndexParam, + InvertIndexParam, + VectorQuery, + VectorSchema, + ) + from zvec.typing import MetricType +except Exception as e: + _ZVEC_IMPORT_ERROR = e + zvec = None # type: ignore[assignment] + + +# zvec max topk (will be lifted to 100,000 in zvec v0.3.2+) +_ZVEC_MAX_TOPK = 1024 + +# Default vector field name +_DEFAULT_VECTOR_FIELD = "embedding" + + +def _escape(value: str) -> str: + """Escape a string value for zvec filter expressions.""" + return value.replace("'", "\\'") + + +def _build_file_store_schema(name: str, dimension: int) -> CollectionSchema: + """Build a zvec CollectionSchema for file store chunks.""" + return CollectionSchema( + name=name, + fields=[ + FieldSchema("content", DataType.STRING, nullable=True, index_param=InvertIndexParam()), + FieldSchema("path", DataType.STRING, nullable=True, index_param=InvertIndexParam()), + FieldSchema("source", DataType.STRING, nullable=True, index_param=InvertIndexParam()), + FieldSchema("start_line", DataType.INT64, nullable=True), + FieldSchema("end_line", DataType.INT64, nullable=True), + FieldSchema("hash", DataType.STRING, nullable=True), + FieldSchema("updated_at", DataType.INT64, nullable=True), + FieldSchema("file_metadata", DataType.STRING, nullable=True), + ], + vectors=[ + VectorSchema( + name=_DEFAULT_VECTOR_FIELD, + data_type=DataType.VECTOR_FP32, + dimension=dimension, + index_param=HnswIndexParam(metric_type=MetricType.COSINE), + ), + ], + ) + + +def _chunk_to_doc(chunk: MemoryChunk, file_meta_json: str = "{}") -> Doc: + """Convert a MemoryChunk to a zvec Doc.""" + fields: dict[str, Any] = { + "content": chunk.text, + "path": chunk.path, + "source": chunk.source.value if chunk.source else "", + "start_line": chunk.start_line, + "end_line": chunk.end_line, + "hash": chunk.hash, + "updated_at": int(time.time() * 1000), + "file_metadata": file_meta_json, + } + vectors: dict[str, Any] = {} + if chunk.embedding is not None: + vectors[_DEFAULT_VECTOR_FIELD] = chunk.embedding + return Doc(id=chunk.id, fields=fields, vectors=vectors) + + +def _doc_to_chunk(doc: Doc) -> MemoryChunk: + """Convert a zvec Doc to a MemoryChunk.""" + raw_vector = doc.vector(_DEFAULT_VECTOR_FIELD) + vector = raw_vector if isinstance(raw_vector, list) and len(raw_vector) > 0 else None + return MemoryChunk( + id=str(doc.id), + path=str(doc.field("path") or ""), + source=MemorySource(str(doc.field("source") or "")), + start_line=int(doc.field("start_line") or 0), + end_line=int(doc.field("end_line") or 0), + text=str(doc.field("content") or ""), + hash=str(doc.field("hash") or ""), + embedding=vector, + ) + + +def _build_source_filter(sources: list[MemorySource] | None) -> str | None: + """Build a zvec filter expression for source filtering.""" + if not sources: + return None + if len(sources) == 1: + return f"source='{_escape(sources[0].value)}'" + vals = ", ".join(f"'{_escape(s.value)}'" for s in sources) + return f"source IN ({vals})" + + +class ZvecFileStore(BaseFileStore): + """Zvec file storage with vector and keyword search. + + Provides zvec-backed persistent storage with: + - Vector similarity search (native zvec HNSW) + - Keyword search (Python substring matching on fetched results) + - Hybrid search (weighted fusion of vector and keyword results) + + Note: + Keyword search operates on chunks fetched from zvec, which is subject + to the topk limit (1024 in zvec < v0.3.2, 100,000 in v0.3.2+). + For collections with more chunks than the topk limit, keyword search + may not scan all documents. + """ + + def __init__( + self, + store_name: str, + db_path: str | Path, + embedding_model: Any | None = None, + vector_enabled: bool = False, + fts_enabled: bool = True, + dimension: int = 1024, + **kwargs: Any, + ): + if _ZVEC_IMPORT_ERROR is not None: + raise ImportError( + "Zvec requires extra dependencies. Install with `pip install zvec`", + ) from _ZVEC_IMPORT_ERROR + + super().__init__( + store_name=store_name, + db_path=db_path, + embedding_model=embedding_model, + vector_enabled=vector_enabled, + fts_enabled=fts_enabled, + **kwargs, + ) + + self.dimension = dimension + self._collection = None + self._initialized = False + self._metadata_file: Path = self.db_path / f"{store_name}_file_metadata.json" + self._metadata_cache: dict[str, dict[str, FileMetadata]] = {} + + @property + def collection_name(self) -> str: + """Get the name of the zvec collection for this store.""" + return f"chunks_{self.store_name}" + + # ------------------------------------------------------------------ + # Lifecycle + # ------------------------------------------------------------------ + + async def start(self) -> None: + """Initialize zvec engine and open the collection.""" + if not self._initialized: + try: + zvec.init() + except RuntimeError: + pass + self._initialized = True + + self.db_path.mkdir(parents=True, exist_ok=True) + collection_path = str(self.db_path / self.collection_name) + option = CollectionOption(read_only=False, enable_mmap=True) + + try: + self._collection = zvec.open(collection_path, option) + logger.info(f"Opened existing zvec file store collection: {collection_path}") + except Exception: + schema = _build_file_store_schema(self.collection_name, self.dimension) + self._collection = zvec.create_and_open( + path=collection_path, + schema=schema, + option=option, + ) + logger.info(f"Created new zvec file store collection: {collection_path}") + + self._metadata_cache = await self._load_metadata() + + async def close(self) -> None: + """Close zvec collection and persist metadata.""" + if self._metadata_cache: + await self._save_metadata(self._metadata_cache) + + if self._collection is not None: + try: + self._collection.flush() + except Exception as e: + logger.warning(f"Failed to flush collection on close: {e}") + self._collection = None + + # ------------------------------------------------------------------ + # Metadata management + # ------------------------------------------------------------------ + + async def _load_metadata(self) -> dict[str, dict[str, FileMetadata]]: + """Load file metadata from JSON file.""" + if not self._metadata_file.exists(): + return {} + try: + data = json.loads(self._metadata_file.read_text(encoding="utf-8")) + result: dict[str, dict[str, FileMetadata]] = {} + for source, files in data.items(): + result[source] = {} + for path, meta in files.items(): + result[source][path] = FileMetadata(**meta) + return result + except Exception as e: + logger.warning(f"Failed to load metadata from {self._metadata_file}: {e}") + return {} + + async def _save_metadata(self, metadata: dict[str, dict[str, FileMetadata]]) -> None: + """Save file metadata to JSON file.""" + try: + out: dict[str, dict[str, dict]] = {} + for source, files in metadata.items(): + out[source] = {} + for path, meta in files.items(): + out[source][path] = { + "path": meta.path, + "hash": meta.hash, + "mtime_ms": meta.mtime_ms, + "size": meta.size, + "chunk_count": meta.chunk_count, + } + self._metadata_file.write_text( + json.dumps(out, indent=2, ensure_ascii=False), + encoding="utf-8", + ) + except Exception as e: + logger.error(f"Failed to save metadata to {self._metadata_file}: {e}") + + # ------------------------------------------------------------------ + # CRUD operations + # ------------------------------------------------------------------ + + async def upsert_file( + self, + file_meta: FileMetadata, + source: MemorySource, + chunks: list[MemoryChunk], + ) -> None: + """Insert or update a file and its chunks.""" + if not chunks: + return + + # Delete existing chunks for this file first + await self.delete_file(file_meta.path, source) + + # Generate embeddings + chunks = await self.get_chunk_embeddings(chunks) + + file_meta_json = json.dumps( + { + "path": file_meta.path, + "hash": file_meta.hash, + "mtime_ms": file_meta.mtime_ms, + "size": file_meta.size, + "chunk_count": len(chunks), + }, + ensure_ascii=False, + ) + + docs = [_chunk_to_doc(c, file_meta_json) for c in chunks] + self._collection.insert(docs) + + # Update metadata cache + if source.value not in self._metadata_cache: + self._metadata_cache[source.value] = {} + self._metadata_cache[source.value][file_meta.path] = FileMetadata( + hash=file_meta.hash, + mtime_ms=file_meta.mtime_ms, + size=file_meta.size, + path=file_meta.path, + chunk_count=len(chunks), + ) + + async def delete_file(self, path: str, source: MemorySource) -> None: + """Delete a file and all its chunks.""" + filter_expr = f"path='{_escape(path)}' AND source='{_escape(source.value)}'" + results = self._collection.query(topk=_ZVEC_MAX_TOPK, filter=filter_expr, include_vector=False) + + ids_to_delete = [doc.id for doc in results] + if ids_to_delete: + self._collection.delete(ids_to_delete) + + if source.value in self._metadata_cache: + self._metadata_cache[source.value].pop(path, None) + + async def delete_file_chunks(self, path: str, chunk_ids: list[str]) -> None: + """Delete specific chunks for a file.""" + if not chunk_ids: + return + self._collection.delete(chunk_ids) + + async def upsert_chunks( + self, + chunks: list[MemoryChunk], + source: MemorySource, + ) -> None: + """Insert or update specific chunks.""" + if not chunks: + return + + chunks = await self.get_chunk_embeddings(chunks) + docs = [_chunk_to_doc(c) for c in chunks] + self._collection.upsert(docs) + + # ------------------------------------------------------------------ + # Listing and metadata + # ------------------------------------------------------------------ + + async def list_files(self, source: MemorySource) -> list[str]: + """List all indexed files for a source.""" + if source.value not in self._metadata_cache: + return [] + return list(self._metadata_cache[source.value].keys()) + + async def get_file_metadata( + self, + path: str, + source: MemorySource, + ) -> FileMetadata | None: + """Get file metadata.""" + if source.value not in self._metadata_cache: + return None + return self._metadata_cache[source.value].get(path) + + async def update_file_metadata(self, file_meta: FileMetadata, source: MemorySource) -> None: + """Update file metadata without affecting chunks.""" + if source.value not in self._metadata_cache: + self._metadata_cache[source.value] = {} + self._metadata_cache[source.value][file_meta.path] = FileMetadata( + hash=file_meta.hash, + mtime_ms=file_meta.mtime_ms, + size=file_meta.size, + path=file_meta.path, + chunk_count=file_meta.chunk_count, + ) + + async def get_file_chunks( + self, + path: str, + source: MemorySource, + ) -> list[MemoryChunk]: + """Get all chunks for a file.""" + filter_expr = f"path='{_escape(path)}' AND source='{_escape(source.value)}'" + results = self._collection.query( + topk=_ZVEC_MAX_TOPK, + filter=filter_expr, + include_vector=True, + ) + chunks = [_doc_to_chunk(doc) for doc in results] + chunks.sort(key=lambda c: c.start_line) + return chunks + + # ------------------------------------------------------------------ + # Search + # ------------------------------------------------------------------ + + async def vector_search( + self, + query: str, + limit: int, + sources: list[MemorySource] | None = None, + ) -> list[MemorySearchResult]: + """Perform vector similarity search.""" + if not self.vector_enabled or not query: + return [] + + query_embedding = await self.get_embedding(query) + if not query_embedding: + return [] + + filter_expr = _build_source_filter(sources) + vq = VectorQuery(field_name=_DEFAULT_VECTOR_FIELD, vector=query_embedding) + + try: + results = self._collection.query( + vectors=vq, + topk=min(limit, _ZVEC_MAX_TOPK), + filter=filter_expr, + include_vector=False, + ) + except Exception as e: + logger.error(f"Vector search failed: {e}") + return [] + + search_results = [] + for doc in results: + score = doc.score if doc.score is not None else 0.0 + # zvec cosine score might need normalization depending on version + search_results.append( + MemorySearchResult( + path=str(doc.field("path") or ""), + start_line=int(doc.field("start_line") or 0), + end_line=int(doc.field("end_line") or 0), + score=score, + snippet=str(doc.field("content") or ""), + source=MemorySource(str(doc.field("source") or "")), + raw_metric=score, + ), + ) + + search_results.sort(key=lambda r: r.score, reverse=True) + return search_results[:limit] + + async def keyword_search( + self, + query: str, + limit: int, + sources: list[MemorySource] | None = None, + ) -> list[MemorySearchResult]: + """Perform keyword search via Python substring matching. + + Fetches chunks from zvec (subject to topk limit) then matches + keywords in Python. For collections larger than the topk limit, + not all documents are scanned. + """ + if not self.fts_enabled or not query: + return [] + + words = query.split() + if not words: + return [] + + # Fetch candidate chunks from zvec + filter_expr = _build_source_filter(sources) + results = self._collection.query( + topk=_ZVEC_MAX_TOPK, + filter=filter_expr, + include_vector=False, + ) + + query_lower = query.lower() + words_lower = [w.lower() for w in words] + n_words = len(words) + + search_results = [] + for doc in results: + text = str(doc.field("content") or "") + text_lower = text.lower() + match_count = sum(1 for w in words_lower if w in text_lower) + if match_count == 0: + continue + + base_score = match_count / n_words + phrase_bonus = 0.2 if n_words > 1 and query_lower in text_lower else 0.0 + score = min(1.0, base_score + phrase_bonus) + + search_results.append( + MemorySearchResult( + path=str(doc.field("path") or ""), + start_line=int(doc.field("start_line") or 0), + end_line=int(doc.field("end_line") or 0), + score=score, + snippet=text, + source=MemorySource(str(doc.field("source") or "")), + ), + ) + + search_results.sort(key=lambda r: r.score, reverse=True) + return search_results[:limit] + + async def hybrid_search( + self, + query: str, + limit: int, + sources: list[MemorySource] | None = None, + vector_weight: float = 0.7, + candidate_multiplier: float = 3.0, + ) -> list[MemorySearchResult]: + """Perform hybrid search combining vector and keyword search.""" + assert 0.0 <= vector_weight <= 1.0, f"vector_weight must be between 0 and 1, got {vector_weight}" + + candidates = min(200, max(1, int(limit * candidate_multiplier))) + text_weight = 1.0 - vector_weight + + if self.vector_enabled and self.fts_enabled: + keyword_results = await self.keyword_search(query, candidates, sources) + vector_results = await self.vector_search(query, candidates, sources) + + if not keyword_results: + return vector_results[:limit] + elif not vector_results: + return keyword_results[:limit] + else: + return self._merge_hybrid_results( + vector=vector_results, + keyword=keyword_results, + vector_weight=vector_weight, + text_weight=text_weight, + )[:limit] + elif self.vector_enabled: + return await self.vector_search(query, limit, sources) + elif self.fts_enabled: + return await self.keyword_search(query, limit, sources) + else: + return [] + + @staticmethod + def _merge_hybrid_results( + vector: list[MemorySearchResult], + keyword: list[MemorySearchResult], + vector_weight: float, + text_weight: float, + ) -> list[MemorySearchResult]: + """Merge vector and keyword search results with weighted scoring.""" + merged: dict[str, MemorySearchResult] = {} + + for result in vector: + result.score = result.score * vector_weight + merged[result.merge_key] = result + + for result in keyword: + key = result.merge_key + if key in merged: + merged[key].score += result.score * text_weight + else: + result.score = result.score * text_weight + merged[key] = result + + results = list(merged.values()) + results.sort(key=lambda r: r.score, reverse=True) + return results + + # ------------------------------------------------------------------ + # Maintenance + # ------------------------------------------------------------------ + + async def clear_all(self) -> None: + """Clear all indexed data.""" + # Delete all documents + stats = self._collection.stats + count = stats.doc_count if stats else 0 + if count > 0: + try: + self._collection.delete_by_filter("content!=''") + except Exception: + remaining = count + while remaining > 0: + batch = self._collection.query( + topk=min(remaining, _ZVEC_MAX_TOPK), + include_vector=False, + ) + if not batch: + break + self._collection.delete([doc.id for doc in batch]) + remaining -= len(batch) + + self._metadata_cache = {} + await self._save_metadata({}) + logger.info(f"Cleared all data from zvec file store: {self.collection_name}") diff --git a/reme/core/vector_store/__init__.py b/reme/core/vector_store/__init__.py index 8429b911..0426fd64 100644 --- a/reme/core/vector_store/__init__.py +++ b/reme/core/vector_store/__init__.py @@ -7,6 +7,7 @@ from .local_vector_store import LocalVectorStore from .obvec_vector_store import ObVecVectorStore from .pgvector_store import PGVectorStore from .qdrant_vector_store import QdrantVectorStore +from .zvec_vector_store import ZvecVectorStore from ..registry_factory import R __all__ = [ @@ -17,6 +18,7 @@ __all__ = [ "ObVecVectorStore", "PGVectorStore", "QdrantVectorStore", + "ZvecVectorStore", ] R.vector_stores.register("chroma")(ChromaVectorStore) @@ -25,3 +27,4 @@ R.vector_stores.register("local")(LocalVectorStore) R.vector_stores.register("obvec")(ObVecVectorStore) R.vector_stores.register("pgvector")(PGVectorStore) R.vector_stores.register("qdrant")(QdrantVectorStore) +R.vector_stores.register("zvec")(ZvecVectorStore) diff --git a/reme/core/vector_store/zvec_vector_store.py b/reme/core/vector_store/zvec_vector_store.py new file mode 100644 index 00000000..dde12a0b --- /dev/null +++ b/reme/core/vector_store/zvec_vector_store.py @@ -0,0 +1,809 @@ +"""Zvec vector store implementation for the ReMe framework.""" + +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any + +from loguru import logger + +from .base_vector_store import BaseVectorStore +from ..embedding import BaseEmbeddingModel +from ..schema import VectorNode + +_ZVEC_IMPORT_ERROR: Exception | None = None + +try: + import zvec # type: ignore[import-untyped] + from zvec import ( + CollectionOption, + CollectionSchema, + DataType, + Doc, + FieldSchema, + HnswIndexParam, + InvertIndexParam, + VectorQuery, + VectorSchema, + ) + from zvec.typing import MetricType +except Exception as e: + _ZVEC_IMPORT_ERROR = e + zvec = None # type: ignore[assignment] + + +# Default vector field name used inside zvec collections +_DEFAULT_VECTOR_FIELD = "embedding" + +# Default scalar content field name for storing text +_CONTENT_FIELD = "content" + +# Field name for JSON-serialized metadata +_METADATA_FIELD = "metadata" + +# Metadata fields promoted to top-level zvec schema columns for native filtering. +# These are the most commonly filtered keys in ReMe's memory system. +# Defining them as independent schema columns allows zvec to perform +# filtering at the database level instead of Python post-filtering. +# Format: {metadata_key: (zvec_data_type_str, has_inverted_index)} +_PROMOTED_FIELD_SPECS: dict[str, tuple[str, bool]] = { + "memory_type": ("STRING", True), # Inverted index for exact match filtering + "memory_target": ("STRING", True), # Inverted index for exact match filtering + "author": ("STRING", False), + "time_int": ("INT64", False), # Numeric for range queries +} + +# zvec data-type string → DataType enum mapping (populated after import) +_DATATYPE_MAP: dict[str, Any] = {} # filled in _build_collection_schema + + +def _escape_zvec_string(value: str) -> str: + """Escape a string value for use in zvec filter expressions.""" + return value.replace("'", "\\'") + + +def _build_zvec_filter( + filters: dict | None, + promoted_fields: set[str], +) -> tuple[str | None, dict | None]: + """Split ReMe filter dict into a zvec native filter expression and remaining post-filters. + + For filter keys that correspond to promoted schema fields, native + zvec filter expressions are generated. Non-promoted keys are + kept for Python post-filtering. + + Args: + filters: ReMe-style filter dictionary. + promoted_fields: Set of metadata keys that exist as top-level schema columns. + + Returns: + (native_filter_expr, post_filter_dict) — either may be None. + """ + if not filters: + return None, None + + native_conditions: list[str] = [] + post_filters: dict = {} + + for key, value in filters.items(): + if key.startswith("$"): + # Compound operators ($or, $and, $not) — keep for post-filtering + post_filters[key] = value + continue + + if key not in promoted_fields: + # Not a promoted field — use post-filtering + post_filters[key] = value + continue + + # Build native filter condition for promoted fields + field_type = _PROMOTED_FIELD_SPECS.get(key, ("STRING", False))[0] + + if isinstance(value, list) and len(value) == 2: + # Range query: [start, end] + if field_type == "INT64": + native_conditions.append(f"{key} >= {value[0]} AND {key} <= {value[1]}") + else: + # STRING range — use >= and <= with string escaping + native_conditions.append( + f"{key} >= '{_escape_zvec_string(str(value[0]))}' " + f"AND {key} <= '{_escape_zvec_string(str(value[1]))}'", + ) + elif isinstance(value, bool): + native_conditions.append(f"{key} = {str(value).upper()}") + elif isinstance(value, (int, float)): + native_conditions.append(f"{key} = {value}") + elif isinstance(value, str): + native_conditions.append(f"{key} = '{_escape_zvec_string(value)}'") + else: + # Unsupported type — fall back to post-filtering + post_filters[key] = value + + native_filter = " AND ".join(native_conditions) if native_conditions else None + return native_filter, post_filters if post_filters else None + + +def _metric_type_from_str(metric: str) -> Any: + """Convert a string metric name to zvec MetricType enum value.""" + if zvec is None: + return None + mapping = { + "cosine": MetricType.COSINE, + "l2": MetricType.L2, + "ip": MetricType.IP, + } + return mapping.get(metric.lower(), MetricType.COSINE) + + +def _build_collection_schema( + name: str, + dimension: int, + metric: str = "cosine", +) -> CollectionSchema: + """Build a zvec CollectionSchema for ReMe usage. + + The schema contains: + - "content" (STRING, inverted index) — text content + - "metadata" (STRING) — JSON-serialized metadata dictionary + - Promoted metadata fields (STRING / INT64) — for native zvec filtering + - "embedding" (VECTOR_FP32, dimension, HNSW index) — the vector field + + Promoted fields are commonly filtered metadata keys defined as top-level + schema columns so that zvec can perform filtering natively instead of + Python post-filtering. The full metadata is still stored as JSON in the + "metadata" field for complete round-trip serialization. + + zvec automatically manages the document ID (string type); we do NOT + define an "id" field in the schema. + """ + # Populate the DataType map on first call + if not _DATATYPE_MAP: + _DATATYPE_MAP.update( + { + "STRING": DataType.STRING, + "INT64": DataType.INT64, + }, + ) + + distance = _metric_type_from_str(metric) + + # Base fields + fields = [ + FieldSchema("content", DataType.STRING, nullable=True, index_param=InvertIndexParam()), + FieldSchema("metadata", DataType.STRING, nullable=True), + ] + + # Add promoted metadata fields as top-level schema columns + for field_name, (type_str, has_inv_index) in _PROMOTED_FIELD_SPECS.items(): + dt = _DATATYPE_MAP[type_str] + idx_param = InvertIndexParam() if has_inv_index else None + fields.append(FieldSchema(field_name, dt, nullable=True, index_param=idx_param)) + + return CollectionSchema( + name=name, + fields=fields, + vectors=[ + VectorSchema( + name=_DEFAULT_VECTOR_FIELD, + data_type=DataType.VECTOR_FP32, + dimension=dimension, + index_param=HnswIndexParam(metric_type=distance), + ), + ], + ) + + +def _vector_node_to_doc(node: VectorNode) -> Doc: + """Convert a ReMe VectorNode to a zvec Doc. + + Metadata is serialized as a JSON string into the "metadata" field. + The "score" key is excluded since it is a computed value, not stored data. + Promoted metadata fields are also extracted as top-level Doc fields + for native zvec filtering. + The vector is placed under the default vector field name. + The zvec Doc id must be a string. + """ + # Filter out computed score before serialization + meta_to_store = {k: v for k, v in node.metadata.items() if k != "score"} + + fields: dict[str, Any] = { + "content": node.content, + "metadata": json.dumps(meta_to_store) if meta_to_store else "{}", + } + + # Extract promoted metadata fields as top-level schema columns + for field_name, (type_str, _) in _PROMOTED_FIELD_SPECS.items(): + value = meta_to_store.get(field_name) + if value is not None: + # Ensure correct type: INT64 fields must be int + if type_str == "INT64" and not isinstance(value, int): + try: + value = int(value) + except (ValueError, TypeError): + continue + fields[field_name] = value + + vectors: dict[str, Any] = {} + if node.vector is not None: + vectors[_DEFAULT_VECTOR_FIELD] = node.vector + + return Doc(id=str(node.vector_id), fields=fields, vectors=vectors) + + +def _doc_to_vector_node(doc: Doc, include_score: bool = False) -> VectorNode: + """Convert a zvec Doc back to a ReMe VectorNode. + + The "metadata" field is parsed from JSON. The "content" field becomes + the node content. If ``include_score`` is True, the search score is + added to the metadata dictionary. + """ + metadata: dict[str, str | bool | int | float] = {} + + # Parse JSON metadata + raw_metadata = doc.field("metadata") + if raw_metadata: + try: + parsed = json.loads(raw_metadata) + if isinstance(parsed, dict): + metadata.update(parsed) + except (json.JSONDecodeError, TypeError): + logger.warning(f"Failed to parse metadata JSON: {raw_metadata}") + + if include_score and doc.score is not None: + metadata["score"] = doc.score + + # Extract vector — doc.vector() returns list or empty dict + raw_vector = doc.vector(_DEFAULT_VECTOR_FIELD) + vector = raw_vector if isinstance(raw_vector, list) and len(raw_vector) > 0 else None + + content = doc.field("content") or "" + + return VectorNode( + vector_id=str(doc.id), + content=str(content), + vector=vector, + metadata=metadata, + ) + + +def _apply_filters_post(nodes: list[VectorNode], filters: dict | None) -> list[VectorNode]: + """Apply ReMe-style filter dict as post-filtering on metadata. + + Used as a fallback for metadata keys that are NOT promoted to top-level + schema columns (and thus cannot be filtered natively by zvec). Promoted + fields are handled by zvec's native ``filter`` parameter instead. + + Supports: + - Exact match: {"field": value} + - Range query: {"field": [start, end]} + """ + if not filters: + return nodes + + filtered = [] + for node in nodes: + match = True + for key, value in filters.items(): + if key.startswith("$"): + # Skip compound operators for post-filtering + continue + node_value = node.metadata.get(key) + + # Range query: [start, end] + if isinstance(value, list) and len(value) == 2: + if node_value is None: + match = False + break + try: + if not value[0] <= node_value <= value[1]: + match = False + break + except TypeError: + match = False + break + else: + # Exact match + if node_value != value: + match = False + break + + if match: + filtered.append(node) + + return filtered + + +class ZvecVectorStore(BaseVectorStore): + """Zvec-based vector store implementation. + + Zvec is a high-performance vector database. This adapter bridges the + ReMe ``BaseVectorStore`` interface with zvec's Python API. + + Supports local persistent storage via ``db_path``. + + Args: + collection_name: Name of the vector collection. + db_path: Local storage path for persistent mode. + embedding_model: Model used for generating vector embeddings. + dimension: Dimensionality of the embedding vectors (default: 1024). + distance: Distance metric — cosine / l2 / ip (default: cosine). + **kwargs: Additional zvec-specific configuration. + """ + + def __init__( + self, + collection_name: str, + db_path: str | Path, + embedding_model: BaseEmbeddingModel, + dimension: int = 1024, + distance: str = "cosine", + **kwargs: Any, + ): + """Initialize the Zvec vector store.""" + if _ZVEC_IMPORT_ERROR is not None: + raise ImportError( + "Zvec requires extra dependencies. Install with `pip install zvec`", + ) from _ZVEC_IMPORT_ERROR + + super().__init__( + collection_name=collection_name, + db_path=db_path, + embedding_model=embedding_model, + **kwargs, + ) + + self.dimension = dimension + self.distance = distance + self._collection = None + self._initialized = False + # Set of promoted field names that exist in the current collection's schema. + # Populated during start() by inspecting the schema. Only fields present + # in the schema can use native zvec filtering; the rest fall back to + # Python post-filtering. + self._promoted_fields_in_schema: set[str] = set() + + # ------------------------------------------------------------------ + # Lifecycle + # ------------------------------------------------------------------ + + async def start(self) -> None: + """Initialize the Zvec engine and open the collection. + + Calls ``zvec.init()`` once, then tries to ``zvec.open()`` an existing + collection or ``zvec.create_and_open()`` a new one. + After opening, detects which promoted fields exist in the schema + and attempts to add missing numeric fields via ``add_column``. + """ + if not self._initialized: + try: + zvec.init() + except RuntimeError: + # Already initialized — safe to ignore + pass + self._initialized = True + + self.db_path.mkdir(parents=True, exist_ok=True) + collection_path = str(self.db_path / self.collection_name) + + option = CollectionOption(read_only=False, enable_mmap=True) + + try: + # Try opening an existing collection first + self._collection = zvec.open(collection_path, option) + logger.info(f"Opened existing Zvec collection at {collection_path}") + except Exception: + # Collection doesn't exist — create it + schema = _build_collection_schema( + name=self.collection_name, + dimension=self.dimension, + metric=self.distance, + ) + self._collection = zvec.create_and_open( + path=collection_path, + schema=schema, + option=option, + ) + logger.info(f"Created new Zvec collection at {collection_path}") + + # Detect which promoted fields exist in the current schema + self._detect_promoted_fields() + + # Try to add missing numeric promoted fields to existing collections + # (zvec's add_column only supports numeric types: INT64, FLOAT, etc.) + self._ensure_numeric_promoted_columns() + + async def close(self) -> None: + """Flush pending writes and release the collection handle.""" + if self._collection is not None: + try: + self._collection.flush() + except Exception as e: + logger.warning(f"Failed to flush collection on close: {e}") + self._collection = None + logger.info(f"Zvec vector store for collection {self.collection_name} closed") + + # ------------------------------------------------------------------ + # Collection management + # ------------------------------------------------------------------ + + async def list_collections(self) -> list[str]: + """Retrieve a list of collection names in the db_path directory. + + Zvec doesn't have a global ``list_collections`` API; we scan the + db_path directory for zvec collection folders. + """ + if not self.db_path.exists(): + return [] + collections = [] + for child in self.db_path.iterdir(): + if child.is_dir(): + collections.append(child.name) + return collections + + async def create_collection(self, collection_name: str, **kwargs) -> None: + """Create a new collection with the specified name and distance metric.""" + if not self._initialized: + try: + zvec.init() + except RuntimeError: + pass + self._initialized = True + + self.db_path.mkdir(parents=True, exist_ok=True) + collection_path = str(self.db_path / collection_name) + + dimension = kwargs.get("dimension", self.dimension) + metric = kwargs.get("distance_metric", self.distance) + + schema = _build_collection_schema( + name=collection_name, + dimension=dimension, + metric=metric, + ) + option = CollectionOption(read_only=False, enable_mmap=True) + + collection = zvec.create_and_open(path=collection_path, schema=schema, option=option) + if collection_name == self.collection_name: + self._collection = collection + logger.info(f"Created collection `{collection_name}`") + + async def delete_collection(self, collection_name: str, **kwargs) -> None: + """Permanently remove a collection from disk.""" + # If it's the active collection, destroy it via zvec API + if self._collection is not None and collection_name == self.collection_name: + try: + self._collection.destroy() + self._collection = None + deleted = True + except Exception as _e: + logger.warning(f"Failed to destroy collection {collection_name}: {_e}") + deleted = False + else: + # For non-active collections, remove the directory + collection_path = self.db_path / collection_name + if collection_path.exists(): + import shutil + + shutil.rmtree(collection_path, ignore_errors=True) + deleted = True + else: + deleted = False + + logger.info(f"Deleted collection {collection_name}: {deleted}") + + async def copy_collection(self, collection_name: str, **kwargs) -> None: + """Duplicate the current collection to a new one with the given name. + + Uses ``shutil.copytree`` to directly copy the collection directory on + disk, which is both faster and complete — it avoids the topk limit of + ``list()`` (max 1024 docs) that would cause data loss for large + collections. + + The source collection is flushed before copying to ensure all + pending writes are persisted to disk. + """ + import shutil + + # Flush source collection so all data is on disk + if self._collection is not None: + self._collection.flush() + + src_path = self.db_path / self.collection_name + dst_path = self.db_path / collection_name + + if not src_path.exists(): + logger.warning(f"Source collection directory not found: {src_path}") + return + + if dst_path.exists(): + logger.warning(f"Target collection already exists: {dst_path}, removing it first") + shutil.rmtree(dst_path, ignore_errors=True) + + shutil.copytree(src_path, dst_path) + logger.info( + f"Copied collection {self.collection_name} to {collection_name} " + f"(directory copy: {src_path} -> {dst_path})", + ) + + # ------------------------------------------------------------------ + # CRUD operations + # ------------------------------------------------------------------ + + async def insert(self, nodes: VectorNode | list[VectorNode], **kwargs) -> None: + """Add one or more vector nodes into the current collection. + + Automatically generates embeddings for nodes that lack vectors. + """ + if isinstance(nodes, VectorNode): + nodes = [nodes] + if not nodes: + return + + # Batch generate embeddings for nodes that need them + nodes_without_vectors = [n for n in nodes if n.vector is None] + if nodes_without_vectors: + nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) + vector_map = {n.vector_id: n for n in nodes_with_vectors} + nodes_to_insert = [vector_map.get(n.vector_id, n) if n.vector is None else n for n in nodes] + else: + nodes_to_insert = nodes + + batch_size = kwargs.get("batch_size", 100) + + for i in range(0, len(nodes_to_insert), batch_size): + batch = nodes_to_insert[i : i + batch_size] + docs = [_vector_node_to_doc(n) for n in batch] + self._collection.insert(docs) + + logger.info(f"Inserted {len(nodes_to_insert)} nodes into {self.collection_name}") + + async def search( + self, + query: str, + limit: int = 5, + filters: dict | None = None, + **kwargs, + ) -> list[VectorNode]: + """Find the most similar vector nodes based on a text query. + + Uses zvec's ``query()`` method with a ``VectorQuery`` built from the + embedding of the query text. Promoted metadata fields are filtered + natively via zvec's ``filter`` parameter; remaining filters are + applied as post-filtering in Python. + """ + query_vector = await self.get_embedding(query) + + vq = VectorQuery( + field_name=_DEFAULT_VECTOR_FIELD, + vector=query_vector, + ) + include_vector = kwargs.get("include_embeddings", False) + + # Split filters: native zvec filter vs Python post-filter + native_filter, post_filters = _build_zvec_filter(filters, self._promoted_fields_in_schema) + + # Over-fetch to compensate for post-filtering + _ZVEC_MAX_TOPK = 1024 + # When post-filters remain, we need to fetch more results because + # many may be filtered out. Use the maximum allowed to minimize misses. + fetch_limit = _ZVEC_MAX_TOPK if post_filters else min(limit, _ZVEC_MAX_TOPK) + + results = self._collection.query( + vectors=vq, + topk=fetch_limit, + filter=native_filter, + include_vector=include_vector, + ) + + nodes = [_doc_to_vector_node(doc, include_score=True) for doc in results] + + # Post-filter on non-promoted metadata fields + nodes = _apply_filters_post(nodes, post_filters) + + score_threshold = kwargs.get("score_threshold") + if score_threshold is not None: + nodes = [n for n in nodes if n.metadata.get("score", 0) >= score_threshold] + + return nodes[:limit] + + async def delete(self, vector_ids: str | list[str], **kwargs) -> None: + """Remove specific vectors from the collection using their identifiers.""" + if isinstance(vector_ids, str): + vector_ids = [vector_ids] + if not vector_ids: + return + + self._collection.delete(vector_ids) + logger.info(f"Deleted {len(vector_ids)} nodes from {self.collection_name}") + + async def delete_all(self, **kwargs) -> None: + """Remove all vectors from the collection. + + Uses zvec's ``delete_by_filter`` with a condition that matches all + documents (content is not empty), or falls back to query + delete + in batches (zvec topk max is 1024). + """ + stats = self._collection.stats + count = stats.doc_count if stats else 0 + if count > 0: + try: + # Use delete_by_filter for efficiency + self._collection.delete_by_filter("content!=''") + except Exception: + # Fallback: fetch all IDs in batches then delete + _ZVEC_MAX_TOPK = 1024 + remaining = count + while remaining > 0: + all_docs = self._collection.query(topk=min(remaining, _ZVEC_MAX_TOPK), include_vector=False) + if not all_docs: + break + ids = [doc.id for doc in all_docs] + self._collection.delete(ids) + remaining -= len(ids) + logger.info(f"Deleted all {count} nodes from {self.collection_name}") + + async def update(self, nodes: VectorNode | list[VectorNode], **kwargs) -> None: + """Update existing vectors using zvec's ``upsert``. + + Automatically regenerates embeddings for nodes whose content changed + but lack an updated vector. + """ + if isinstance(nodes, VectorNode): + nodes = [nodes] + if not nodes: + return + + # Batch generate embeddings for nodes that need them + nodes_without_vectors = [n for n in nodes if n.vector is None and n.content] + if nodes_without_vectors: + nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) + vector_map = {n.vector_id: n for n in nodes_with_vectors} + nodes_to_update = [vector_map.get(n.vector_id, n) if n.vector is None and n.content else n for n in nodes] + else: + nodes_to_update = nodes + + docs = [_vector_node_to_doc(n) for n in nodes_to_update] + self._collection.upsert(docs) + logger.info(f"Updated {len(nodes_to_update)} nodes in {self.collection_name}") + + async def get(self, vector_ids: str | list[str]) -> VectorNode | list[VectorNode]: + """Fetch specific vector nodes from the collection by their IDs.""" + is_single = isinstance(vector_ids, str) + ids = [vector_ids] if is_single else vector_ids + + result_dict = self._collection.fetch(ids) + nodes = [_doc_to_vector_node(doc) for doc in result_dict.values()] + return nodes[0] if is_single and nodes else (nodes if not is_single else None) + + async def list( + self, + filters: dict | None = None, + limit: int | None = None, + sort_key: str | None = None, + reverse: bool = True, + ) -> list[VectorNode]: + """Retrieve vectors matching optional metadata filters. + + Uses zvec's ``query()`` without a vector query to list all documents. + Promoted metadata fields are filtered natively via zvec's ``filter`` + parameter; remaining filters are applied as post-filtering in Python. + + Args: + filters: Dictionary of filter conditions to match vectors. + limit: Maximum number of vectors to return. + sort_key: Key to sort the results by (in metadata). + reverse: If True, sort in descending order; otherwise ascending. + """ + # Split filters: native zvec filter vs Python post-filter + native_filter, post_filters = _build_zvec_filter(filters, self._promoted_fields_in_schema) + + # Determine fetch limit — zvec max topk is 1024 (will be lifted to 100,000 in zvec v0.3.2+) + _ZVEC_MAX_TOPK = 1024 + fetch_limit = min(limit or _ZVEC_MAX_TOPK, _ZVEC_MAX_TOPK) + if sort_key or post_filters: + fetch_limit = _ZVEC_MAX_TOPK # fetch max and sort/filter in Python + + results = self._collection.query( + topk=fetch_limit, + filter=native_filter, + include_vector=True, + ) + + nodes = [_doc_to_vector_node(doc) for doc in results] + + # Post-filter on non-promoted metadata fields + nodes = _apply_filters_post(nodes, post_filters) + + # Apply sorting if sort_key is provided + if sort_key: + + def _sort_key_func(node: VectorNode): + value = node.metadata.get(sort_key) + if value is None: + return float("-inf") if not reverse else float("inf") + return value + + nodes.sort(key=_sort_key_func, reverse=reverse) + + if limit is not None: + nodes = nodes[:limit] + + return nodes + + # ------------------------------------------------------------------ + # Helpers + # ------------------------------------------------------------------ + + def _detect_promoted_fields(self) -> None: + """Detect which promoted fields exist in the current collection's schema. + + Compares the set of promoted field names against the actual schema + and populates ``_promoted_fields_in_schema`` accordingly. Only fields + present in the schema can use native zvec filtering. + """ + if self._collection is None: + return + + try: + schema = self._collection.schema + existing_fields = {f.name for f in schema.fields} if schema.fields else set() + except Exception as e: + logger.warning(f"Failed to read collection schema: {e}") + existing_fields = set() + + self._promoted_fields_in_schema = set(_PROMOTED_FIELD_SPECS.keys()) & existing_fields + + missing = set(_PROMOTED_FIELD_SPECS.keys()) - existing_fields + if missing: + logger.info( + f"Promoted fields not in schema (will use post-filtering): {missing}", + ) + + def _ensure_numeric_promoted_columns(self) -> None: + """Add missing numeric promoted fields to existing collections. + + zvec's ``add_column`` only supports numeric types (INT64, FLOAT, etc.). + STRING fields cannot be added via ``add_column`` and must be defined + at collection creation time. For those, we fall back to post-filtering. + """ + if self._collection is None: + return + + missing = set(_PROMOTED_FIELD_SPECS.keys()) - self._promoted_fields_in_schema + if not missing: + return + + # Populate the DataType map if needed + if not _DATATYPE_MAP: + _DATATYPE_MAP.update( + { + "STRING": DataType.STRING, + "INT64": DataType.INT64, + }, + ) + + for field_name in missing: + type_str, _ = _PROMOTED_FIELD_SPECS[field_name] + # Only numeric types can be added via add_column + if type_str not in ("INT64", "INT32", "FLOAT", "DOUBLE"): + continue + try: + dt = _DATATYPE_MAP[type_str] + self._collection.add_column(FieldSchema(field_name, dt, nullable=True)) + self._promoted_fields_in_schema.add(field_name) + logger.info(f"Added promoted column '{field_name}' to existing collection") + except Exception as e: + logger.warning(f"Failed to add column '{field_name}': {e}") + + async def count(self) -> int: + """Return the total number of documents in the current collection.""" + stats = self._collection.stats + return stats.doc_count if stats else 0 + + async def reset(self): + """Reset the current collection by destroying and recreating it.""" + logger.warning(f"Resetting collection {self.collection_name}...") + await self.delete_collection(self.collection_name) + await self.create_collection(self.collection_name) + logger.info(f"Collection {self.collection_name} has been reset") diff --git a/tests/test_file_store.py b/tests/test_file_store.py index fd44332d..0e34764f 100644 --- a/tests/test_file_store.py +++ b/tests/test_file_store.py @@ -28,6 +28,7 @@ from reme.core.file_store.base_file_store import BaseFileStore from reme.core.file_store.chroma_file_store import ChromaFileStore from reme.core.file_store.local_file_store import LocalFileStore from reme.core.file_store.sqlite_file_store import SqliteFileStore +from reme.core.file_store.zvec_file_store import ZvecFileStore from reme.core.schema.file_metadata import FileMetadata from reme.core.schema.memory_chunk import MemoryChunk from reme.core.utils import load_env @@ -53,6 +54,10 @@ class TestConfig: CHROMA_DB_PATH = "./test_file_store_chroma" CHROMA_FTS_ENABLED = True + # ZvecFileStore settings + ZVEC_DB_PATH = "./test_file_store_zvec" + ZVEC_FTS_ENABLED = True + # LocalFileStore settings LOCAL_DB_PATH = "./test_file_store_local" LOCAL_FTS_ENABLED = True @@ -199,6 +204,8 @@ def get_store_type(store: BaseFileStore) -> str: return "chroma" elif isinstance(store, LocalFileStore): return "local" + elif isinstance(store, ZvecFileStore): + return "zvec" else: raise ValueError(f"Unknown file store type: {type(store)}") @@ -242,6 +249,14 @@ def create_file_store(store_type: str) -> BaseFileStore: embedding_model=embedding_model, fts_enabled=config.LOCAL_FTS_ENABLED, ) + elif store_type == "zvec": + return ZvecFileStore( + store_name=config.NAME, + db_path=config.ZVEC_DB_PATH, + embedding_model=embedding_model, + fts_enabled=config.ZVEC_FTS_ENABLED, + dimension=config.EMBEDDING_DIMENSIONS, + ) else: raise ValueError(f"Unknown store type: {store_type}") @@ -284,6 +299,13 @@ async def test_start_store(store: BaseFileStore, _store_name: str): assert isinstance(store._files, dict), "Files index should be a dict" logger.info(f"✓ LocalFileStore ready (chunks file: {store._chunks_file})") + # Verify ZvecFileStore initialized + if isinstance(store, ZvecFileStore): + # pylint: disable=protected-access + assert store._collection is not None, "Zvec collection should be initialized" + assert store._initialized, "Zvec engine should be initialized" + logger.info(f"✓ ZvecFileStore ready (collection: {store.collection_name})") + async def test_upsert_file(store: BaseFileStore, _store_name: str) -> tuple[FileMetadata, List[MemoryChunk]]: """Test file and chunks insertion.""" @@ -1013,6 +1035,18 @@ async def cleanup_store(store: BaseFileStore, store_type: str): json_file.unlink() logger.info(f"✓ Cleaned up file: {json_file}") + # Clean up zvec directory and metadata file + if store_type == "zvec": + config = TestConfig() + db_dir = Path(config.ZVEC_DB_PATH) + if db_dir.exists(): + shutil.rmtree(db_dir) + logger.info(f"✓ Cleaned up directory: {db_dir}") + metadata_file = db_dir.parent / f"{config.NAME}_file_metadata.json" + if metadata_file.exists(): + metadata_file.unlink() + logger.info(f"✓ Cleaned up metadata file: {metadata_file}") + logger.info("✓ Cleanup completed") except Exception as e: logger.error(f"Cleanup error: {e}") @@ -1049,6 +1083,11 @@ Examples: action="store_true", help="Test LocalFileStore", ) + parser.add_argument( + "--zvec", + action="store_true", + help="Test ZvecFileStore", + ) parser.add_argument( "--all", action="store_true", @@ -1065,6 +1104,7 @@ Examples: ("sqlite", "SqliteFileStore"), ("chroma", "ChromaFileStore"), ("local", "LocalFileStore"), + ("zvec", "ZvecFileStore"), ] else: # Build list based on individual flags @@ -1074,6 +1114,8 @@ Examples: stores_to_test.append(("chroma", "ChromaFileStore")) if args.local: stores_to_test.append(("local", "LocalFileStore")) + if args.zvec: + stores_to_test.append(("zvec", "ZvecFileStore")) if not stores_to_test: # Default to all file stores if no argument provided @@ -1081,9 +1123,10 @@ Examples: ("sqlite", "SqliteFileStore"), ("chroma", "ChromaFileStore"), ("local", "LocalFileStore"), + ("zvec", "ZvecFileStore"), ] print("No file store specified, defaulting to test all file stores") - print("Use --sqlite, --chroma, or --local to test specific ones\n") + print("Use --sqlite, --chroma, --local, or --zvec to test specific ones\n") # Run tests for each file store for store_type, store_name in stores_to_test: diff --git a/tests/test_vector_store.py b/tests/test_vector_store.py index 7f7264be..9b39ab9b 100644 --- a/tests/test_vector_store.py +++ b/tests/test_vector_store.py @@ -2,7 +2,7 @@ """Unified test suite for vector store implementations. This module provides comprehensive test coverage for LocalVectorStore, ESVectorStore, -PGVectorStore, QdrantVectorStore, ChromaVectorStore, and ObVecVectorStore implementations. +PGVectorStore, QdrantVectorStore, ChromaVectorStore, ObVecVectorStore and ZvecVectorStore implementations. Tests can be run for specific vector stores or all implementations. Usage: @@ -12,6 +12,7 @@ Usage: python test_vector_store.py --qdrant # Test QdrantVectorStore only python test_vector_store.py --chroma # Test ChromaVectorStore only python test_vector_store.py --obvec # Test ObVecVectorStore only (needs seekdb / OceanBase) + python test_vector_store.py --zvec # Test ZvecVectorStore only python test_vector_store.py --all # Test all vector stores """ @@ -36,6 +37,7 @@ from reme.core.vector_store import ( ObVecVectorStore, PGVectorStore, QdrantVectorStore, + ZvecVectorStore, ) load_env() @@ -90,6 +92,9 @@ class TestConfig: OBVEC_PASSWORD = os.environ.get("OBVEC_PASSWORD", "root") OBVEC_DATABASE = os.environ.get("OBVEC_DATABASE", "test") + # ZvecVectorStore settings + ZVEC_PATH = "./test_vector_store_zvec" # For local persistent mode + # Embedding model settings EMBEDDING_MODEL_NAME = "text-embedding-v4" EMBEDDING_DIMENSIONS = 64 @@ -192,6 +197,7 @@ class SampleDataGenerator: # ==================== Vector Store Factory ==================== +# pylint: disable=too-many-return-statements def get_store_type(store: BaseVectorStore) -> str: """Get the type identifier of a vector store instance. @@ -199,7 +205,7 @@ def get_store_type(store: BaseVectorStore) -> str: store: Vector store instance Returns: - str: Type identifier ("local", "es", "pgvector", "qdrant", "chroma", or "obvec") + str: Type identifier ("local", "es", "pgvector", "qdrant", "chroma", "obvec", or "zvec") """ if isinstance(store, LocalVectorStore): return "local" @@ -213,10 +219,13 @@ def get_store_type(store: BaseVectorStore) -> str: return "chroma" elif isinstance(store, ObVecVectorStore): return "obvec" + elif isinstance(store, ZvecVectorStore): + return "zvec" else: raise ValueError(f"Unknown vector store type: {type(store)}") +# pylint: disable=too-many-return-statements def create_vector_store(store_type: str, collection_name: str) -> BaseVectorStore: """Create a vector store instance based on type. @@ -295,6 +304,14 @@ def create_vector_store(store_type: str, collection_name: str) -> BaseVectorStor index_metric="cosine", index_ef_search=100, ) + elif store_type == "zvec": + return ZvecVectorStore( + collection_name=collection_name, + embedding_model=embedding_model, + db_path=config.ZVEC_PATH or tempfile.mkdtemp(prefix="test_zvec_"), + dimension=config.EMBEDDING_DIMENSIONS, + distance="cosine", + ) else: raise ValueError(f"Unknown store type: {store_type}") @@ -1790,6 +1807,13 @@ async def cleanup_store(store: BaseVectorStore, store_type: str): shutil.rmtree(obvec_dir, ignore_errors=True) logger.info(f"Cleaned up obvec temp directory: {obvec_dir}") + # Clean up local directory if ZvecVectorStore + if store_type == "zvec" and config.ZVEC_PATH: + test_dir = Path(config.ZVEC_PATH) + if test_dir.exists(): + shutil.rmtree(test_dir) + logger.info(f"Cleaned up zvec directory: {config.ZVEC_PATH}") + logger.info("✓ Cleanup completed") except Exception as e: logger.error(f"Cleanup error: {e}") @@ -1844,6 +1868,11 @@ Examples: action="store_true", help="Test ObVecVectorStore", ) + parser.add_argument( + "--zvec", + action="store_true", + help="Test ZvecVectorStore", + ) parser.add_argument( "--all", action="store_true", @@ -1863,6 +1892,7 @@ Examples: ("qdrant", "QdrantVectorStore"), ("chroma", "ChromaVectorStore"), ("obvec", "ObVecVectorStore"), + ("zvec", "ZvecVectorStore"), ] else: # Build list based on individual flags @@ -1878,6 +1908,8 @@ Examples: stores_to_test.append(("chroma", "ChromaVectorStore")) if args.obvec: stores_to_test.append(("obvec", "ObVecVectorStore")) + if args.zvec: + stores_to_test.append(("zvec", "ZvecVectorStore")) if not stores_to_test: # Default to all vector stores if no argument provided @@ -1888,10 +1920,11 @@ Examples: ("qdrant", "QdrantVectorStore"), ("chroma", "ChromaVectorStore"), ("obvec", "ObVecVectorStore"), + ("zvec", "ZvecVectorStore"), ] print("No vector store specified, defaulting to test all vector stores") print( - "Use --local/--es/--pgvector/--qdrant/--chroma/--obvec to test specific ones\n", + "Use --local/--es/--pgvector/--qdrant/--chroma/--obvec/--zvec to test specific ones\n", ) # Run tests for each vector store diff --git a/tests/test_zvec_vector_store.py b/tests/test_zvec_vector_store.py new file mode 100644 index 00000000..59836f24 --- /dev/null +++ b/tests/test_zvec_vector_store.py @@ -0,0 +1,906 @@ +"""Test suite for ZvecVectorStore implementation. + +Comprehensive tests covering CRUD operations, search, filtering, +collection management, and edge cases for the zvec vector store adapter. + +Usage: + python -m pytest tests/test_zvec_vector_store.py -v + python tests/test_zvec_vector_store.py +""" + +# pylint: disable=redefined-outer-name,unused-argument + +from __future__ import annotations + +import asyncio +import shutil +import tempfile +from pathlib import Path +from typing import List +from uuid import uuid4 + +import pytest + +from loguru import logger + +from reme.core.schema import VectorNode +from reme.core.vector_store import ZvecVectorStore + +# --------------------------------------------------------------------------- +# Skip entire module if zvec native library is not installed +# --------------------------------------------------------------------------- +try: + import zvec as _zvec # noqa: F401 — just checking availability +except ImportError: + pytest.skip("zvec native library not installed", allow_module_location=True) + + +# ==================== Configuration ==================== + + +class TestConfig: + """Configuration for zvec test execution.""" + + ZVEC_ROOT_PATH = tempfile.mkdtemp(prefix="test_zvec_") + EMBEDDING_DIMENSION = 64 # Small dimension for faster tests + TEST_COLLECTION_PREFIX = "test_zvec_vs" + + +# ==================== Sample Data ==================== + + +def create_sample_nodes(prefix: str = "") -> List[VectorNode]: + """Create sample VectorNode instances for testing.""" + id_prefix = f"{prefix}_" if prefix else "" + return [ + VectorNode( + vector_id=f"{id_prefix}node1", + content="Artificial intelligence is a technology that simulates human intelligence.", + metadata={ + "node_type": "tech", + "category": "AI", + "source": "research", + "priority": "high", + "year": "2023", + }, + ), + VectorNode( + vector_id=f"{id_prefix}node2", + content="Machine learning is a subset of artificial intelligence.", + metadata={ + "node_type": "tech", + "category": "ML", + "source": "research", + "priority": "high", + "year": "2022", + }, + ), + VectorNode( + vector_id=f"{id_prefix}node3", + content="Deep learning uses neural networks with multiple layers.", + metadata={ + "node_type": "tech_new", + "category": "DL", + "source": "blog", + "priority": "medium", + "year": "2024", + }, + ), + VectorNode( + vector_id=f"{id_prefix}node4", + content="I love eating delicious seafood, especially fresh fish.", + metadata={ + "node_type": "food", + "category": "preference", + "source": "personal", + "priority": "low", + "year": "2023", + }, + ), + VectorNode( + vector_id=f"{id_prefix}node5", + content="Natural language processing enables computers to understand human language.", + metadata={ + "node_type": "tech", + "category": "NLP", + "source": "research", + "priority": "high", + "year": "2024", + }, + ), + ] + + +# ==================== Fixtures ==================== + + +class MockEmbeddingModel: + """A mock embedding model that generates deterministic random vectors. + + Avoids external API calls during testing. Produces unit-normalized + vectors so that cosine similarity works correctly. + """ + + def __init__(self, dimension: int = 64): + self.dimension = dimension + + async def get_embedding(self, query: str) -> list[float]: + """Generate a deterministic embedding from a query string.""" + import hashlib + import struct + + h = hashlib.sha256(query.encode()).digest() + # Repeat hash to fill dimension + full_hash = b"" + while len(full_hash) < self.dimension * 4: + full_hash += hashlib.sha256(h + full_hash).digest() + + vec = list(struct.unpack(f"<{self.dimension}f", full_hash[: self.dimension * 4])) + # Normalize to unit vector + norm = sum(x * x for x in vec) ** 0.5 + if norm > 0: + vec = [x / norm for x in vec] + return vec + + async def get_embeddings(self, queries: list[str]) -> list[list[float]]: + """Generate embeddings for multiple queries.""" + return [await self.get_embedding(q) for q in queries] + + async def get_node_embedding(self, node: VectorNode) -> VectorNode: + """Assign embedding to a single node.""" + if node.content: + node.vector = await self.get_embedding(node.content) + return node + + async def get_node_embeddings(self, nodes: list[VectorNode]) -> list[VectorNode]: + """Assign embeddings to multiple nodes.""" + return [await self.get_node_embedding(n) for n in nodes] + + +@pytest.fixture +def embedding_model(): + """Provide a MockEmbeddingModel for tests.""" + return MockEmbeddingModel(dimension=TestConfig.EMBEDDING_DIMENSION) + + +@pytest.fixture +def zvec_store(embedding_model, tmp_path): + """Create and start a ZvecVectorStore for testing. + + Yields the store and cleans up afterwards. + """ + collection_name = f"{TestConfig.TEST_COLLECTION_PREFIX}_{uuid4().hex[:8]}" + store = ZvecVectorStore( + collection_name=collection_name, + db_path=str(tmp_path / "zvec_db"), + embedding_model=embedding_model, + dimension=TestConfig.EMBEDDING_DIMENSION, + distance="cosine", + ) + + async def _setup(): + await store.start() + return store + + store = asyncio.get_event_loop().run_until_complete(_setup()) + yield store + + async def _teardown(): + try: + await store.close() + except Exception: + pass + # Clean up temp directory + db_path = Path(str(tmp_path / "zvec_db")) + if db_path.exists(): + shutil.rmtree(db_path, ignore_errors=True) + + asyncio.get_event_loop().run_until_complete(_teardown()) + + +# ==================== Helper ==================== + + +def run(coro): + """Run an async coroutine in the current event loop.""" + return asyncio.get_event_loop().run_until_complete(coro) + + +# ==================== Test: Collection Lifecycle ==================== + + +class TestCollectionLifecycle: + """Tests for collection creation, listing, deletion, and copy.""" + + def test_create_collection(self, zvec_store): + """Test that a collection is created during start().""" + collections = run(zvec_store.list_collections()) + assert zvec_store.collection_name in collections + + def test_list_collections(self, zvec_store): + """Test listing collections.""" + collections = run(zvec_store.list_collections()) + assert isinstance(collections, list) + assert len(collections) >= 1 + + def test_delete_collection(self, zvec_store, embedding_model, tmp_path): + """Test deleting a collection.""" + # Create a secondary collection + coll_name = f"del_test_{uuid4().hex[:8]}" + store2 = ZvecVectorStore( + collection_name=coll_name, + db_path=str(tmp_path / "zvec_db"), + embedding_model=embedding_model, + dimension=TestConfig.EMBEDDING_DIMENSION, + ) + run(store2.start()) + + collections = run(zvec_store.list_collections()) + assert coll_name in collections + + run(zvec_store.delete_collection(coll_name)) + + collections = run(zvec_store.list_collections()) + assert coll_name not in collections + + def test_copy_collection(self, zvec_store, embedding_model, tmp_path): + """Test copying a collection.""" + # Insert some data first + nodes = create_sample_nodes("copy") + run(zvec_store.insert(nodes)) + + copy_name = f"copy_test_{uuid4().hex[:8]}" + run(zvec_store.copy_collection(copy_name)) + + # Verify copy exists + collections = run(zvec_store.list_collections()) + assert copy_name in collections + + # Clean up + run(zvec_store.delete_collection(copy_name)) + + +# ==================== Test: Insert ==================== + + +class TestInsert: + """Tests for node insertion (single and batch).""" + + def test_insert_single_node(self, zvec_store): + """Test inserting a single node.""" + node = VectorNode( + vector_id="single_1", + content="This is a single node insertion test", + metadata={"test_type": "single_insert"}, + ) + run(zvec_store.insert(node)) + + result = run(zvec_store.get("single_1")) + assert result is not None + assert result.vector_id == "single_1" + assert "single node" in result.content + + def test_insert_batch_nodes(self, zvec_store): + """Test inserting multiple nodes in batch.""" + nodes = create_sample_nodes("batch") + run(zvec_store.insert(nodes)) + + all_nodes = run(zvec_store.list(limit=10)) + assert len(all_nodes) >= len(nodes) + + def test_insert_node_with_vector(self, zvec_store): + """Test inserting a node that already has a vector.""" + node = VectorNode( + vector_id="prevec_1", + content="Node with pre-computed vector", + vector=[0.1] * TestConfig.EMBEDDING_DIMENSION, + metadata={"test_type": "pre_vector"}, + ) + run(zvec_store.insert(node)) + + result = run(zvec_store.get("prevec_1")) + assert result is not None + assert result.vector is not None + + +# ==================== Test: Search ==================== + + +class TestSearch: + """Tests for vector similarity search.""" + + @pytest.fixture(autouse=True) + def _insert_sample_data(self, zvec_store): + """Insert sample data before each search test.""" + nodes = create_sample_nodes("search") + run(zvec_store.insert(nodes)) + + def test_basic_search(self, zvec_store): + """Test basic vector search.""" + results = run(zvec_store.search(query="What is artificial intelligence?", limit=3)) + assert len(results) > 0 + for r in results: + assert isinstance(r, VectorNode) + assert r.content + + def test_search_with_limit(self, zvec_store): + """Test search with various limits.""" + results = run(zvec_store.search(query="technology", limit=2)) + assert len(results) <= 2 + + def test_search_with_filter(self, zvec_store): + """Test vector search with metadata filter.""" + results = run( + zvec_store.search( + query="What is artificial intelligence?", + limit=5, + filters={"node_type": "tech"}, + ), + ) + # All results should have node_type == "tech" + for r in results: + assert r.metadata.get("node_type") == "tech" + + def test_search_with_multiple_filters(self, zvec_store): + """Test search with multiple metadata filters (AND).""" + results = run( + zvec_store.search( + query="What is artificial intelligence?", + limit=5, + filters={"node_type": "tech", "source": "research"}, + ), + ) + for r in results: + assert r.metadata.get("node_type") == "tech" + assert r.metadata.get("source") == "research" + + def test_search_relevance_ranking(self, zvec_store): + """Test that search results have scores and are relevant.""" + results = run(zvec_store.search(query="artificial intelligence", limit=5)) + assert len(results) > 0 + # All results should have a score + for r in results: + assert "score" in r.metadata + assert r.metadata["score"] > 0 + # The top result should be highly relevant (AI content matches AI query) + top_content = results[0].content.lower() + assert "artificial intelligence" in top_content or "intelligence" in top_content or "ai" in top_content + + +# ==================== Test: Get ==================== + + +class TestGet: + """Tests for retrieving nodes by ID.""" + + @pytest.fixture(autouse=True) + def _insert_sample_data(self, zvec_store): + """Insert sample data before each get test.""" + nodes = create_sample_nodes("get") + run(zvec_store.insert(nodes)) + + def test_get_single_id(self, zvec_store): + """Test retrieving a single node by ID.""" + result = run(zvec_store.get("get_node1")) + assert result is not None + assert result.vector_id == "get_node1" + + def test_get_multiple_ids(self, zvec_store): + """Test retrieving multiple nodes by IDs.""" + results = run(zvec_store.get(["get_node1", "get_node2"])) + assert isinstance(results, list) + assert len(results) >= 2 + result_ids = {r.vector_id for r in results} + assert "get_node1" in result_ids + assert "get_node2" in result_ids + + def test_get_nonexistent_id(self, zvec_store): + """Test retrieving a non-existent ID.""" + result = run(zvec_store.get("nonexistent_id_xyz")) + assert result is None or result == [] + + +# ==================== Test: List ==================== + + +class TestList: + """Tests for listing nodes with optional filters and sorting.""" + + @pytest.fixture(autouse=True) + def _insert_sample_data(self, zvec_store): + """Insert sample data before each list test.""" + nodes = create_sample_nodes("list") + run(zvec_store.insert(nodes)) + + def test_list_all(self, zvec_store): + """Test listing all nodes.""" + results = run(zvec_store.list(limit=20)) + assert len(results) > 0 + + def test_list_with_filter(self, zvec_store): + """Test listing nodes with metadata filter.""" + results = run(zvec_store.list(filters={"category": "AI"}, limit=10)) + for r in results: + assert r.metadata.get("category") == "AI" + + def test_list_with_sorting(self, zvec_store): + """Test listing with sorting by metadata key.""" + # Insert nodes with numeric metadata for sorting + sort_nodes = [ + VectorNode( + vector_id=f"sort_{i}", + content=f"Sort test node {i}", + metadata={"rating": str(50 + i * 5), "test_type": "sort_test"}, + ) + for i in range(10) + ] + run(zvec_store.insert(sort_nodes)) + + results = run( + zvec_store.list( + filters={"test_type": "sort_test"}, + sort_key="rating", + reverse=True, + limit=5, + ), + ) + assert len(results) <= 5 + # Verify descending order + ratings = [r.metadata.get("rating") for r in results] + for i in range(len(ratings) - 1): + assert ratings[i] >= ratings[i + 1] + + +# ==================== Test: Update ==================== + + +class TestUpdate: + """Tests for updating existing nodes.""" + + @pytest.fixture(autouse=True) + def _insert_sample_data(self, zvec_store): + """Insert sample data before each update test.""" + nodes = create_sample_nodes("upd") + run(zvec_store.insert(nodes)) + + def test_update_single_node(self, zvec_store): + """Test updating a single node's content and metadata.""" + updated = VectorNode( + vector_id="upd_node2", + content="Machine learning is a powerful subset of AI that learns from data.", + metadata={ + "node_type": "tech", + "category": "ML", + "updated": "true", + }, + ) + run(zvec_store.update(updated)) + + result = run(zvec_store.get("upd_node2")) + assert result is not None + assert result.metadata.get("updated") == "true" + + def test_update_batch(self, zvec_store): + """Test batch updating multiple nodes.""" + updates = [ + VectorNode( + vector_id="upd_node1", + content="Updated content for node 1", + metadata={"node_type": "tech", "batch_updated": "true"}, + ), + VectorNode( + vector_id="upd_node3", + content="Updated content for node 3", + metadata={"node_type": "tech_new", "batch_updated": "true"}, + ), + ] + run(zvec_store.update(updates)) + + results = run(zvec_store.get(["upd_node1", "upd_node3"])) + if isinstance(results, list): + for r in results: + assert r.metadata.get("batch_updated") == "true" + + +# ==================== Test: Delete ==================== + + +class TestDelete: + """Tests for deleting nodes.""" + + @pytest.fixture(autouse=True) + def _insert_sample_data(self, zvec_store): + """Insert sample data before each delete test.""" + nodes = create_sample_nodes("del") + run(zvec_store.insert(nodes)) + + def test_delete_single(self, zvec_store): + """Test deleting a single node by ID.""" + run(zvec_store.delete("del_node4")) + + # Verify deletion + result = run(zvec_store.get("del_node4")) + assert result is None or result == [] + + def test_delete_batch(self, zvec_store): + """Test batch deleting multiple nodes by IDs.""" + # First insert some extra nodes to delete + extra_nodes = [ + VectorNode( + vector_id=f"del_extra_{i}", + content=f"Extra node {i} for batch delete test", + metadata={"test_type": "batch_delete"}, + ) + for i in range(5) + ] + run(zvec_store.insert(extra_nodes)) + + ids = [f"del_extra_{i}" for i in range(5)] + run(zvec_store.delete(ids)) + + # Verify all deleted + for nid in ids: + result = run(zvec_store.get(nid)) + assert result is None or result == [] + + def test_delete_all(self, zvec_store): + """Test deleting all nodes from the collection.""" + run(zvec_store.delete_all()) + # Collection should be empty now + remaining = run(zvec_store.list(limit=100)) + assert len(remaining) == 0 + + +# ==================== Test: Edge Cases ==================== + + +class TestEdgeCases: + """Tests for edge cases and boundary conditions.""" + + def test_empty_content(self, zvec_store): + """Test inserting a node with empty content.""" + node = VectorNode( + vector_id="edge_empty", + content="", + metadata={"type": "empty"}, + ) + # Empty content may fail embedding — that's OK, we just want to see it handled + try: + run(zvec_store.insert([node])) + except Exception: + pass # Expected if embedding fails on empty string + + def test_long_content(self, zvec_store): + """Test inserting a node with very long content.""" + node = VectorNode( + vector_id="edge_long", + content="A" * 5000, + metadata={"type": "long_content"}, + ) + run(zvec_store.insert([node])) + result = run(zvec_store.get("edge_long")) + assert result is not None + assert len(result.content) == 5000 + + def test_special_characters(self, zvec_store): + """Test content with special characters.""" + node = VectorNode( + vector_id="edge_special", + content="Special chars: @#$%^&*()[]{}|;:',.<>?/~`", + metadata={"type": "special_chars"}, + ) + run(zvec_store.insert([node])) + result = run(zvec_store.get("edge_special")) + assert result is not None + assert "@#$%" in result.content + + def test_unicode_content(self, zvec_store): + """Test content with Unicode characters.""" + node = VectorNode( + vector_id="edge_unicode", + content="Unicode test: 你好世界 مرحبا Привет", + metadata={"type": "unicode"}, + ) + run(zvec_store.insert([node])) + result = run(zvec_store.get("edge_unicode")) + assert result is not None + assert "你好世界" in result.content + + def test_nonexistent_id(self, zvec_store): + """Test getting a non-existent ID.""" + result = run(zvec_store.get("nonexistent_xyz_999")) + assert result is None or result == [] + + def test_metadata_with_empty_string_value(self, zvec_store): + """Test metadata containing empty string values.""" + node = VectorNode( + vector_id="edge_meta_empty", + content="Testing empty metadata values", + metadata={"field1": "value1", "field2": "", "field3": "value3"}, + ) + run(zvec_store.insert([node])) + result = run(zvec_store.get("edge_meta_empty")) + assert result is not None + + def test_search_nonexistent_filter(self, zvec_store): + """Test search with a filter value that doesn't match anything.""" + nodes = create_sample_nodes("edge_filter") + run(zvec_store.insert(nodes)) + + results = run( + zvec_store.search( + query="test", + limit=10, + filters={"category": "NONEXISTENT_CATEGORY"}, + ), + ) + assert len(results) == 0 + + +# ==================== Test: Batch Operations ==================== + + +class TestBatchOperations: + """Tests for large-scale batch insert, update, and delete.""" + + def test_batch_insert_100_nodes(self, zvec_store): + """Test inserting 100 nodes in batch.""" + batch_nodes = [ + VectorNode( + vector_id=f"batch_{i}", + content=f"This is batch test content number {i} about technology and science.", + metadata={ + "batch_id": str(i // 10), + "index": str(i), + "category": ["tech", "science", "business"][i % 3], + }, + ) + for i in range(100) + ] + run(zvec_store.insert(batch_nodes)) + + all_nodes = run(zvec_store.list(limit=150)) + assert len(all_nodes) >= 100 + + def test_batch_update_20_nodes(self, zvec_store): + """Test batch updating 20 nodes.""" + # Insert first + nodes = [ + VectorNode( + vector_id=f"bupd_{i}", + content=f"Batch update test {i}", + metadata={"index": str(i)}, + ) + for i in range(30) + ] + run(zvec_store.insert(nodes)) + + # Update first 20 + updates = [ + VectorNode( + vector_id=f"bupd_{i}", + content=f"UPDATED content {i}", + metadata={"index": str(i), "updated": "true"}, + ) + for i in range(20) + ] + run(zvec_store.update(updates)) + + # Verify + results = run(zvec_store.list(filters={"updated": "true"}, limit=30)) + assert len(results) >= 20 + + def test_batch_delete_50_nodes(self, zvec_store): + """Test batch deleting 50 nodes.""" + # Insert + nodes = [ + VectorNode( + vector_id=f"bdel_{i}", + content=f"Batch delete test {i}", + metadata={"index": str(i)}, + ) + for i in range(50) + ] + run(zvec_store.insert(nodes)) + + # Delete + ids = [f"bdel_{i}" for i in range(50)] + run(zvec_store.delete(ids)) + + # Verify + remaining = run(zvec_store.list(limit=200)) + batch_remaining = [n for n in remaining if n.vector_id.startswith("bdel_")] + assert len(batch_remaining) == 0 + + +# ==================== Test: Concurrent Operations ==================== + + +class TestConcurrentOperations: + """Tests for concurrent read/write operations.""" + + def test_concurrent_inserts_and_searches(self, zvec_store): + """Test that concurrent inserts and searches work without errors.""" + + async def _run(): + # Concurrent inserts + insert_tasks = [] + for i in range(5): + batch = [ + VectorNode( + vector_id=f"conc_{i}_{j}", + content=f"Concurrent test content {i}-{j}", + metadata={"thread_id": str(i)}, + ) + for j in range(10) + ] + insert_tasks.append(zvec_store.insert(batch)) + + await asyncio.gather(*insert_tasks) + + # Concurrent searches + search_tasks = [zvec_store.search(query="concurrent test", limit=5) for _ in range(5)] + search_results = await asyncio.gather(*search_tasks) + + # All searches should return results + for results in search_results: + assert len(results) > 0 + + run(_run()) + + +# ==================== Test: Data Model Conversion ==================== + + +class TestDataModelConversion: + """Tests for VectorNode <-> zvec Doc conversion helpers.""" + + def test_vector_node_to_doc_roundtrip(self, zvec_store): + """Test that VectorNode -> Doc -> VectorNode roundtrip preserves data.""" + from reme.core.vector_store.zvec_vector_store import ( + _vector_node_to_doc, + _doc_to_vector_node, + ) + + original = VectorNode( + vector_id="roundtrip_1", + content="Roundtrip test content", + vector=[0.1] * TestConfig.EMBEDDING_DIMENSION, + metadata={"key1": "value1", "key2": "42", "key3": "true"}, + ) + + doc = _vector_node_to_doc(original) + assert doc.id == "roundtrip_1" + assert doc.field("content") == "Roundtrip test content" + + restored = _doc_to_vector_node(doc, include_score=False) + assert restored.vector_id == "roundtrip_1" + assert restored.content == "Roundtrip test content" + assert restored.metadata.get("key1") == "value1" + + def test_post_filter_exact_match(self): + """Test post-filtering with exact match.""" + from reme.core.vector_store.zvec_vector_store import _apply_filters_post + + nodes = [ + VectorNode(vector_id="1", content="a", metadata={"category": "AI"}), + VectorNode(vector_id="2", content="b", metadata={"category": "ML"}), + VectorNode(vector_id="3", content="c", metadata={"category": "AI"}), + ] + + filtered = _apply_filters_post(nodes, {"category": "AI"}) + assert len(filtered) == 2 + assert all(n.metadata["category"] == "AI" for n in filtered) + + def test_post_filter_range_query(self): + """Test post-filtering with range query.""" + from reme.core.vector_store.zvec_vector_store import _apply_filters_post + + nodes = [ + VectorNode(vector_id="1", content="a", metadata={"year": 2022}), + VectorNode(vector_id="2", content="b", metadata={"year": 2023}), + VectorNode(vector_id="3", content="c", metadata={"year": 2024}), + ] + + filtered = _apply_filters_post(nodes, {"year": [2023, 2024]}) + assert len(filtered) == 2 + + def test_post_filter_none_and_empty(self): + """Test post-filtering with None and empty filters.""" + from reme.core.vector_store.zvec_vector_store import _apply_filters_post + + nodes = [VectorNode(vector_id="1", content="a", metadata={})] + + # None filter returns all + assert _apply_filters_post(nodes, None) == nodes + # Empty filter returns all + assert _apply_filters_post(nodes, {}) == nodes + + def test_score_excluded_from_stored_metadata(self): + """Test that score is excluded when converting VectorNode to Doc.""" + from reme.core.vector_store.zvec_vector_store import _vector_node_to_doc + + node = VectorNode( + vector_id="score_test", + content="test", + vector=[0.1] * TestConfig.EMBEDDING_DIMENSION, + metadata={"key1": "val1", "score": 0.95}, + ) + + doc = _vector_node_to_doc(node) + # The metadata JSON should NOT contain the score key + import json + + stored_meta = json.loads(doc.field("metadata")) + assert "score" not in stored_meta + assert "key1" in stored_meta + + +# ==================== Main Entry Point ==================== + + +async def run_standalone_tests(): + """Run tests standalone (without pytest) for quick validation.""" + tmp_dir = tempfile.mkdtemp(prefix="test_zvec_standalone_") + embedding_model = MockEmbeddingModel(dimension=TestConfig.EMBEDDING_DIMENSION) + + store = ZvecVectorStore( + collection_name="standalone_test", + db_path=tmp_dir, + embedding_model=embedding_model, + dimension=TestConfig.EMBEDDING_DIMENSION, + distance="cosine", + ) + + try: + await store.start() + logger.info("✓ Store started") + + # Insert + nodes = create_sample_nodes("std") + await store.insert(nodes) + logger.info(f"✓ Inserted {len(nodes)} nodes") + + # Search + results = await store.search(query="artificial intelligence", limit=3) + logger.info(f"✓ Search returned {len(results)} results") + for r in results: + logger.info(f" - {r.vector_id}: {r.content[:50]}... (score={r.metadata.get('score')})") + + # Get + result = await store.get("std_node1") + logger.info(f"✓ Get: {result.vector_id if result else 'None'}") + + # List + all_nodes = await store.list(limit=10) + logger.info(f"✓ List: {len(all_nodes)} nodes") + + # Update + await store.update( + VectorNode( + vector_id="std_node1", + content="Updated content", + metadata={"updated": "true"}, + ), + ) + result = await store.get("std_node1") + logger.info(f"✓ Update: metadata.updated={result.metadata.get('updated') if result else 'N/A'}") + + # Delete + await store.delete("std_node4") + result = await store.get("std_node4") + logger.info(f"✓ Delete: {'gone' if result is None or result == [] else 'still exists'}") + + # Count + count = await store.count() + logger.info(f"✓ Count: {count} nodes") + + logger.info("✓ All standalone tests passed!") + + finally: + await store.close() + shutil.rmtree(tmp_dir, ignore_errors=True) + + +if __name__ == "__main__": + asyncio.run(run_standalone_tests()) diff --git a/tests/vector/test_reme_vector.py b/tests/vector/test_reme_vector.py index 3829c841..8a37cb97 100644 --- a/tests/vector/test_reme_vector.py +++ b/tests/vector/test_reme_vector.py @@ -20,7 +20,7 @@ async def main(): "dimensions": 1024, }, default_vector_store_config={ - "backend": "local", # 支持 local/chroma/qdrant/elasticsearch + "backend": "local", # 支持 local/chroma/qdrant/elasticsearch/zvec }, ) await reme.start() From d72f5fc581b1e770f134dedf9e998d6f2aac907f Mon Sep 17 00:00:00 2001 From: yangtiancheng-ali Date: Sat, 9 May 2026 10:30:37 +0800 Subject: [PATCH 3/3] feat(vector_store): add Hologres vector store implementation (#226) --- README.md | 2 +- docs/index.md | 2 +- docs/vector_store_api_guide.md | 52 +- reme/core/vector_store/__init__.py | 3 + reme/core/vector_store/hologres_store.py | 633 +++++++++++++++++++++++ tests/test_vector_store.py | 56 +- 6 files changed, 737 insertions(+), 11 deletions(-) create mode 100644 reme/core/vector_store/hologres_store.py diff --git a/README.md b/README.md index b4bb48c4..b408332b 100644 --- a/README.md +++ b/README.md @@ -506,7 +506,7 @@ async def main(): "dimensions": 1024, }, default_vector_store_config={ - "backend": "local", # Supports local/chroma/qdrant/elasticsearch/obvec/zvec + "backend": "local", # Supports local/chroma/qdrant/elasticsearch/obvec/zvec/hologres }, ) await reme.start() diff --git a/docs/index.md b/docs/index.md index d7e2e716..d0e9d9d2 100644 --- a/docs/index.md +++ b/docs/index.md @@ -139,7 +139,7 @@ response = requests.post("http://localhost:8002/retrieve_task_memory", json={ ## 📚 Resources - **[Installation Guide](installation.md)**, **[Quick Start](quick_start.md)**: Get started quickly with practical examples -- **[Vector Storage Setup](vector_store_api_guide.md)**: Configure local, Elasticsearch, Qdrant, ChromaDB, or ObVec (OceanBase / seekdb via pyobvector) storage and usage +- **[Vector Storage Setup](vector_store_api_guide.md)**: Configure local, Elasticsearch, Qdrant, ChromaDB, ObVec (OceanBase / seekdb via pyobvector) or Hologres storage and usage - **[MCP Guide](mcp_quick_start.md)**: Create MCP services - **[Personal Memory](personal_memory/personal_memory.md)**, **[Task Memory](task_memory/task_memory.md)** & **[Tool Memory](tool_memory/tool_memory.md)**: Operators used in personal memory, task memory and tool memory. You can modify the config to customize the pipelines. - **[Example Collection](./cookbook/appworld/quickstart.md)**: Real use cases and best practices diff --git a/docs/vector_store_api_guide.md b/docs/vector_store_api_guide.md index 89881ef2..53ced358 100644 --- a/docs/vector_store_api_guide.md +++ b/docs/vector_store_api_guide.md @@ -34,6 +34,7 @@ FlowLLM provides multiple Vector Store implementations tailored to different use - **ChromaVectorStore** ([source code](https://github.com/flowllm-ai/flowllm/blob/main/flowllm/core/vector_store/chroma_vector_store.py)): Based on ChromaDB, providing persistent storage and metadata filtering capabilities. - **EsVectorStore** ([source code](https://github.com/flowllm-ai/flowllm/blob/main/flowllm/core/vector_store/es_vector_store.py)): Built on Elasticsearch, enabling powerful combined full-text and vector search functionalities. - **ObVecVectorStore** ([source code](https://github.com/agentscope-ai/ReMe/blob/main/reme/core/vector_store/obvec_vector_store.py)): Uses [pyobvector](https://pypi.org/project/pyobvector/) against **OceanBase** or **seekdb** (MySQL-compatible wire protocol). Suitable when you already run OceanBase/seekdb or need a SQL-native vector table with HNSW-style ANN search and JSON metadata filters. +- **HologresVectorStore** ([source code](https://github.com/agentscope-ai/ReMe/blob/main/reme/core/vector_store/hologres_store.py)): Uses [asyncpg](https://pypi.org/project/asyncpg/) against **Hologres** (PostgreSQL-compatible). Leverages native `float4[]` vector storage with built-in HGraph index for approximate nearest neighbor search. Suitable when you already run Hologres or need high-performance vector search with JSONB metadata filtering in a PostgreSQL-compatible environment. - **ZvecVectorStore** ([source code](https://github.com/agentscope-ai/ReMe/blob/main/reme/core/vector_store/zvec_vector_store.py)): Built on zvec, a high-performance local vector database with strong-schema support and HNSW indexing. Suitable for single-machine deployments requiring fast vector search. All Vector Store implementations inherit from **BaseVectorStore** ([source code](https://github.com/agentscope-ai/ReMe/blob/main/reme/core/vector_store/base_vector_store.py)) in ReMe, ensuring a consistent interface specification. @@ -137,6 +138,20 @@ OBVEC_PASSWORD= python tests/test_vector_store.py --obvec - **dimension**: Dimensionality of the embedding vectors (default: `1024`). - **distance**: Distance metric — supports `cosine`, `l2`, `ip` (default: `cosine`). +### HologresVectorStore Configuration + +- **host**: Hologres host address (default: `localhost`). +- **port**: Hologres port (default: `80`). +- **database**: Database name (default: `postgres`). +- **user**: Database user (default: `postgres`). +- **password**: Database password. +- **schema**: PostgreSQL schema name (default: `public`). +- **min_size**: Minimum connections in pool (default: `1`). +- **max_size**: Maximum connections in pool (default: `10`). +- **dsn**: Full DSN connection string. When provided, overrides `host`, `port`, `database`, `user`, and `password`. +- **distance_method**: Distance method for the HGraph index: `Cosine`, `InnerProduct`, or `Euclidean` (default: `Cosine`). +- **collection_name**: Table name for the collection (from `VectorStoreConfig`, default `reme`). + ## Configuration File Examples Configure Vector Store in `flowllm/config/default.yaml` under the `vector_store` section. The basic structure is as follows: @@ -157,7 +172,7 @@ vector_store.default.params.= ### Configuration Field Descriptions -- **`backend`** (required): Vector store backend type. Options: `local`, `memory`, `chroma`, `qdrant`, `elasticsearch`, `obvec`, `zvec`. +- **`backend`** (required): Vector store backend type. Options: `local`, `memory`, `chroma`, `qdrant`, `elasticsearch`, `obvec`, `zvec`, `hologres`. - **`embedding_model`** (required): Name of the embedding model configuration, referencing the `embedding_model` section. - **`params`** (optional): Dictionary of backend-specific parameters passed to the vector store constructor. @@ -353,7 +368,37 @@ vector_stores.default.password=your-root-password ReMe service YAML uses the key `vector_stores` (plural); CLI overrides use the same nested paths. -#### 7. ZvecVectorStore Configuration +#### 7. HologresVectorStore Configuration + +**Implementation**: [`reme/core/vector_store/hologres_store.py`](https://github.com/agentscope-ai/ReMe/blob/main/reme/core/vector_store/hologres_store.py) + +**Example (Hologres instance)**: + +```yaml +vector_stores: + default: + backend: hologres + embedding_model: default + collection_name: reme + host: "your-hologres-host" + port: 80 + database: "postgres" + user: "postgres" + password: "your-password" + schema: "public" + distance_method: "Cosine" +``` + +```shell +vector_stores.default.backend=hologres +vector_stores.default.host=your-hologres-host +vector_stores.default.port=80 +vector_stores.default.user=postgres +vector_stores.default.password=your-password +vector_stores.default.database=postgres +``` + +#### 8. ZvecVectorStore Configuration Persistent local storage based on zvec with HNSW indexing and strong-schema support. @@ -434,10 +479,11 @@ Two types of metadata filtering are supported: - **Development & Testing**: Use MemoryVectorStore or LocalVectorStore—no additional services required. - **Small-Scale Applications**: Use LocalVectorStore or ChromaVectorStore for simplicity and ease of use. -- **Production Environments**: Use QdrantVectorStore, EsVectorStore, or ObVecVectorStore (OceanBase/seekdb) for high performance and scalability, depending on your existing infrastructure. +- **Production Environments**: Use QdrantVectorStore, EsVectorStore, ObVecVectorStore (OceanBase/seekdb), or HologresVectorStore for high performance and scalability, depending on your existing infrastructure. - **High-Performance Local Search**: Use ZvecVectorStore for single-machine deployments requiring fast HNSW-based vector search with local persistence. - **Hybrid Search**: Use EsVectorStore to combine vector search with full-text search capabilities. - **OceanBase / seekdb**: Use ObVecVectorStore when you standardize on pyobvector and SQL-accessible vector tables. +- **Hologres**: Use HologresVectorStore when you run Hologres and need native HGraph-indexed vector search with PostgreSQL-compatible SQL and JSONB metadata filtering. ## Important Notes diff --git a/reme/core/vector_store/__init__.py b/reme/core/vector_store/__init__.py index 0426fd64..7b5a60bb 100644 --- a/reme/core/vector_store/__init__.py +++ b/reme/core/vector_store/__init__.py @@ -3,6 +3,7 @@ from .base_vector_store import BaseVectorStore from .chroma_vector_store import ChromaVectorStore from .es_vector_store import ESVectorStore +from .hologres_store import HologresVectorStore from .local_vector_store import LocalVectorStore from .obvec_vector_store import ObVecVectorStore from .pgvector_store import PGVectorStore @@ -14,6 +15,7 @@ __all__ = [ "BaseVectorStore", "ChromaVectorStore", "ESVectorStore", + "HologresVectorStore", "LocalVectorStore", "ObVecVectorStore", "PGVectorStore", @@ -23,6 +25,7 @@ __all__ = [ R.vector_stores.register("chroma")(ChromaVectorStore) R.vector_stores.register("es")(ESVectorStore) +R.vector_stores.register("hologres")(HologresVectorStore) R.vector_stores.register("local")(LocalVectorStore) R.vector_stores.register("obvec")(ObVecVectorStore) R.vector_stores.register("pgvector")(PGVectorStore) diff --git a/reme/core/vector_store/hologres_store.py b/reme/core/vector_store/hologres_store.py new file mode 100644 index 00000000..270dc1bf --- /dev/null +++ b/reme/core/vector_store/hologres_store.py @@ -0,0 +1,633 @@ +"""Hologres implementation for vector storage and retrieval.""" + +import json +import re +from pathlib import Path +from typing import Any + +from loguru import logger + +from .base_vector_store import BaseVectorStore +from ..embedding import BaseEmbeddingModel +from ..schema import VectorNode + +_ASYNCPG_IMPORT_ERROR: Exception | None = None + +try: + import asyncpg + from asyncpg import Pool +except Exception as e: + _ASYNCPG_IMPORT_ERROR = e + asyncpg = None + Pool = None + + +class HologresVectorStore(BaseVectorStore): + """Vector store implementation using Hologres for efficient similarity search. + + Hologres uses native float4[] arrays for vector storage with built-in + HGraph index for approximate nearest neighbor search, unlike pgvector + which requires an extension. + """ + + @staticmethod + def _validate_table_name(name: str) -> None: + """Validate table name to prevent SQL injection.""" + if not name: + raise ValueError("Table name cannot be empty") + if len(name) > 63: + raise ValueError(f"Table name too long: {len(name)} characters (max 63)") + if not re.match(r"^[a-zA-Z_][a-zA-Z0-9_]*$", name): + raise ValueError( + f"Invalid table name: {name}. Must start with letter or underscore, " + "and contain only alphanumeric characters and underscores.", + ) + + def __init__( + self, + collection_name: str, + db_path: str | Path, + embedding_model: BaseEmbeddingModel, + host: str = "localhost", + port: int = 80, + database: str = "postgres", + user: str = "postgres", + password: str = "", + schema: str = "public", + min_size: int = 1, + max_size: int = 10, + dsn: str | None = None, + distance_method: str = "Cosine", + **kwargs, + ): + """Initialize the Hologres vector store with connection parameters. + + Args: + collection_name: Name of the collection (table). + db_path: Database path (used by base class). + embedding_model: Embedding model for generating vectors. + host: Hologres host address. + port: Hologres port (default 80 for Hologres). + database: Database name. + user: Database user. + password: Database password. + schema: PostgreSQL schema name (default "public"). + min_size: Minimum connections in pool. + max_size: Maximum connections in pool. + dsn: Full DSN connection string (overrides individual params). + distance_method: Distance method for HGraph index (Cosine, InnerProduct, Euclidean). + """ + if _ASYNCPG_IMPORT_ERROR is not None: + raise ImportError( + "Hologres vector store requires asyncpg. Install with `pip install asyncpg`", + ) from _ASYNCPG_IMPORT_ERROR + + self._validate_table_name(collection_name) + self._validate_table_name(schema) + + super().__init__( + collection_name=collection_name, + db_path=db_path, + embedding_model=embedding_model, + **kwargs, + ) + + self.dsn = dsn + self.host = host + self.port = port + self.database = database + self.user = user + self.password = password + self.schema = schema + self.min_size = min_size + self.max_size = max_size + self.distance_method = distance_method + self._pool: Pool | None = None + self.embedding_model_dims = embedding_model.dimensions + + @property + def _qualified_name(self) -> str: + """Return the schema-qualified table name (e.g. 'my_schema.my_table').""" + return f"{self.schema}.{self.collection_name}" + + def _qualify(self, table_name: str) -> str: + """Return a schema-qualified name for an arbitrary table.""" + return f"{self.schema}.{table_name}" + + @staticmethod + async def _hologres_reset(conn): + """Custom reset for Hologres connections.""" + await conn.execute( + """ + SELECT pg_advisory_unlock_all(); + CLOSE ALL; + RESET ALL; + """, + ) + + async def _get_pool(self) -> Pool: + """Create or return the existing asyncpg connection pool.""" + if self._pool is None: + if self.dsn: + self._pool = await asyncpg.create_pool( + dsn=self.dsn, + min_size=self.min_size, + max_size=self.max_size, + reset=self._hologres_reset, + ) + else: + self._pool = await asyncpg.create_pool( + host=self.host, + port=self.port, + database=self.database, + user=self.user, + password=self.password, + min_size=self.min_size, + max_size=self.max_size, + reset=self._hologres_reset, + ) + + # Ensure schema exists + async with self._pool.acquire() as conn: + await conn.execute(f"CREATE SCHEMA IF NOT EXISTS {self.schema}") + + logger.info(f"Hologres connection pool created for database {self.database}") + + return self._pool + + @staticmethod + def _vector_to_pg_array(vector: list[float]) -> str: + """Convert a Python list of floats to PostgreSQL array literal format.""" + return "{" + ",".join(map(str, vector)) + "}" + + @staticmethod + def _pg_array_to_vector(pg_array) -> list[float] | None: + """Convert a PostgreSQL array result to a Python list of floats.""" + if pg_array is None: + return None + if isinstance(pg_array, list): + return [float(x) for x in pg_array] + # Handle string format like {1.0,2.0,3.0} + raw = str(pg_array) + if raw.startswith("{") and raw.endswith("}"): + return [float(x) for x in raw[1:-1].split(",")] + return None + + async def list_collections(self) -> list[str]: + """List all available table names in the current schema.""" + pool = await self._get_pool() + async with pool.acquire() as conn: + rows = await conn.fetch( + "SELECT table_name FROM information_schema.tables WHERE table_schema = $1", + self.schema, + ) + return [row["table_name"] for row in rows] + + async def create_collection(self, collection_name: str, **kwargs): + """Create a new Hologres table with vector support and HGraph index.""" + self._validate_table_name(collection_name) + pool = await self._get_pool() + dimensions = kwargs.get("dimensions", self.embedding_model_dims) + qualified = self._qualify(collection_name) + + async with pool.acquire() as conn: + create_sql = f""" + CREATE TABLE IF NOT EXISTS {qualified} ( + id TEXT PRIMARY KEY, + content TEXT, + vector float4[] CHECK (array_ndims(vector) = 1 AND array_length(vector, 1) = {dimensions}), + metadata JSONB + ) + WITH ( + vectors = '{{ + "vector": {{ + "algorithm": "HGraph", + "distance_method": "{self.distance_method}", + "builder_params": {{ + "base_quantization_type": "rabitq", + "rabitq_use_fht":true, + "graph_storage_type": "compressed", + "max_total_size_to_merge_mb": 4096, + "max_degree": 64, + "ef_construction": 400, + "precise_quantization_type": "fp32", + "use_reorder": true + }} + }} + }}' + ) + """ + await conn.execute(create_sql) + + logger.info(f"Created Hologres collection {qualified} with dimensions={dimensions}") + + async def delete_collection(self, collection_name: str, **kwargs): + """Remove the specified collection table from the database.""" + self._validate_table_name(collection_name) + pool = await self._get_pool() + qualified = self._qualify(collection_name) + async with pool.acquire() as conn: + await conn.execute(f"DROP TABLE IF EXISTS {qualified}") + logger.info(f"Deleted collection {qualified}") + + async def copy_collection(self, collection_name: str, **kwargs): + """Duplicate the structure and content of the current collection to a new table.""" + self._validate_table_name(collection_name) + pool = await self._get_pool() + qualified_src = self._qualified_name + qualified_dst = self._qualify(collection_name) + + async with pool.acquire() as conn: + columns = await conn.fetch( + """ + SELECT column_name, data_type, udt_name + FROM information_schema.columns + WHERE table_name = $1 AND table_schema = $2 + """, + self.collection_name, + self.schema, + ) + + if not columns: + raise ValueError(f"Source collection {qualified_src} does not exist") + + # Create new table with primary key, then add data + await conn.execute( + f""" + SET hg_experimental_enable_create_table_like_properties = true; + CALL hg_create_table_like('{qualified_dst}', 'select * from {qualified_src}') + """, + ) + await conn.execute(f"INSERT INTO {qualified_dst} SELECT * FROM {qualified_src} ;") + + logger.info(f"Copied collection {qualified_src} to {qualified_dst}") + + async def insert(self, nodes: VectorNode | list[VectorNode], **kwargs): + """Insert or upsert vector nodes into the Hologres collection.""" + if isinstance(nodes, VectorNode): + nodes = [nodes] + + if not nodes: + return + + nodes_without_vectors = [node for node in nodes if node.vector is None] + if nodes_without_vectors: + nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) + vector_map = {n.vector_id: n for n in nodes_with_vectors} + nodes_to_insert = [vector_map.get(n.vector_id, n) if n.vector is None else n for n in nodes] + else: + nodes_to_insert = nodes + + pool = await self._get_pool() + data = [ + ( + node.vector_id, + node.content, + node.vector, + json.dumps(node.metadata), + ) + for node in nodes_to_insert + ] + + async with pool.acquire() as conn: + on_conflict = kwargs.get("on_conflict", "update") + + if on_conflict == "update": + await conn.executemany( + f""" + INSERT INTO {self._qualified_name} (id, content, vector, metadata) + VALUES ($1, $2, $3::float4[], $4::jsonb) + ON CONFLICT (id) DO UPDATE SET + content = EXCLUDED.content, + vector = EXCLUDED.vector, + metadata = EXCLUDED.metadata + """, + data, + ) + elif on_conflict == "ignore": + await conn.executemany( + f""" + INSERT INTO {self._qualified_name} (id, content, vector, metadata) + VALUES ($1, $2, $3::float4[], $4::jsonb) + ON CONFLICT (id) DO NOTHING + """, + data, + ) + else: + await conn.executemany( + f""" + INSERT INTO {self._qualified_name} (id, content, vector, metadata) + VALUES ($1, $2, $3::float4[], $4::jsonb) + """, + data, + ) + + logger.info(f"Inserted {len(nodes_to_insert)} documents into {self._qualified_name}") + + @staticmethod + def _build_filter_clause(filters: dict | None) -> tuple[str, list]: + """Generate an SQL WHERE clause and parameter list from a filter dictionary. + + Supports two filter formats: + 1. Range query: {"field": [start_value, end_value]} + 2. Exact match: {"field": value} + """ + if not filters: + return "", [] + + conditions = [] + params = [] + param_idx = 1 + + for key, value in filters.items(): + if not key.replace("_", "").replace(".", "").isalnum(): + raise ValueError( + f"Invalid metadata key: {key}. Only alphanumeric characters, underscore and dot are allowed.", + ) + + if isinstance(value, list) and len(value) == 2: + if isinstance(value[0], (int, float)) and isinstance(value[1], (int, float)): + conditions.append( + f"(metadata->>'{key}')::numeric >= ${param_idx} AND " + f"(metadata->>'{key}')::numeric <= ${param_idx + 1}", + ) + else: + conditions.append(f"metadata->>'{key}' >= ${param_idx} AND metadata->>'{key}' <= ${param_idx + 1}") + params.extend([value[0], value[1]]) + param_idx += 2 + else: + conditions.append(f"metadata->>'{key}' = ${param_idx}") + params.append(str(value)) + param_idx += 1 + + filter_clause = "WHERE " + " AND ".join(conditions) if conditions else "" + return filter_clause, params + + async def search( + self, + query: str, + limit: int = 5, + filters: dict | None = None, + **kwargs, + ) -> list[VectorNode]: + """Perform vector similarity search using Hologres approx_cosine_distance.""" + query_vector = await self.get_embedding(query) + vector_str = self._vector_to_pg_array(query_vector) + pool = await self._get_pool() + + filter_clause, filter_params = self._build_filter_clause(filters) + + # filter_params use $1..$N, limit uses $(N+1) + limit_placeholder = f"${len(filter_params) + 1}" + + async with pool.acquire() as conn: + sql = f""" + SELECT id, content, vector, metadata, + approx_cosine_distance(vector, '{vector_str}') AS distance + FROM {self._qualified_name} + {filter_clause} + ORDER BY distance DESC + LIMIT {limit_placeholder} + """ + rows = await conn.fetch(sql, *filter_params, limit) + + results = [] + score_threshold = kwargs.get("score_threshold") + + for row in rows: + distance = float(row["distance"]) + # approx_cosine_distance returns cosine similarity (higher = more similar) + score = distance + if score_threshold is not None and score < score_threshold: + continue + + vector_data = self._pg_array_to_vector(row["vector"]) + + metadata = row["metadata"] if row["metadata"] else {} + if isinstance(metadata, str): + metadata = json.loads(metadata) + + metadata["score"] = score + metadata["_distance"] = 1 - score + + node = VectorNode( + vector_id=row["id"], + content=row["content"] or "", + vector=vector_data, + metadata=metadata, + ) + results.append(node) + + return results + + async def delete(self, vector_ids: str | list[str], **kwargs): + """Remove specific vector records from the collection by their IDs.""" + if isinstance(vector_ids, str): + vector_ids = [vector_ids] + + if not vector_ids: + return + + pool = await self._get_pool() + async with pool.acquire() as conn: + placeholders = ", ".join([f"${i + 1}" for i in range(len(vector_ids))]) + await conn.execute( + f"DELETE FROM {self._qualified_name} WHERE id IN ({placeholders})", + *vector_ids, + ) + + logger.info(f"Deleted {len(vector_ids)} documents from {self._qualified_name}") + + async def delete_all(self, **kwargs): + """Remove all vectors from the collection.""" + pool = await self._get_pool() + async with pool.acquire() as conn: + result = await conn.execute(f"DELETE FROM {self._qualified_name}") + + logger.info(f"Deleted all documents from {self._qualified_name} result={result}") + + async def update(self, nodes: VectorNode | list[VectorNode], **kwargs): + """Update existing vector nodes with new content, embeddings, or metadata.""" + if isinstance(nodes, VectorNode): + nodes = [nodes] + + if not nodes: + return + + nodes_without_vectors = [node for node in nodes if node.vector is None and node.content] + if nodes_without_vectors: + nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) + vector_map = {n.vector_id: n for n in nodes_with_vectors} + nodes_to_update = [vector_map.get(n.vector_id, n) if n.vector is None and n.content else n for n in nodes] + else: + nodes_to_update = nodes + + pool = await self._get_pool() + async with pool.acquire() as conn: + for node in nodes_to_update: + update_fields = [] + params = [] + idx = 1 + + if node.content: + update_fields.append(f"content = ${idx}") + params.append(node.content) + idx += 1 + + if node.vector: + update_fields.append(f"vector = ${idx}::float4[]") + params.append(node.vector) + idx += 1 + + if node.metadata: + update_fields.append(f"metadata = ${idx}::jsonb") + params.append(json.dumps(node.metadata)) + idx += 1 + + if update_fields: + params.append(node.vector_id) + await conn.execute( + f"UPDATE {self._qualified_name} SET {', '.join(update_fields)} WHERE id = ${idx}", + *params, + ) + + logger.info(f"Updated {len(nodes_to_update)} documents in {self._qualified_name}") + + async def get(self, vector_ids: str | list[str]) -> VectorNode | list[VectorNode] | None: + """Retrieve vector nodes by their unique identifiers.""" + single_result = isinstance(vector_ids, str) + if single_result: + vector_ids = [vector_ids] + + if not vector_ids: + return [] if not single_result else None + + pool = await self._get_pool() + async with pool.acquire() as conn: + placeholders = ", ".join([f"${i + 1}" for i in range(len(vector_ids))]) + rows = await conn.fetch( + f"SELECT id, content, vector, metadata FROM {self._qualified_name} WHERE id IN ({placeholders})", + *vector_ids, + ) + + results = [] + for row in rows: + vector_data = self._pg_array_to_vector(row["vector"]) + + metadata = row["metadata"] if row["metadata"] else {} + if isinstance(metadata, str): + metadata = json.loads(metadata) + + results.append( + VectorNode( + vector_id=row["id"], + content=row["content"] or "", + vector=vector_data, + metadata=metadata, + ), + ) + + if single_result: + return results[0] if results else None + return results + + async def list( + self, + filters: dict | None = None, + limit: int | None = None, + sort_key: str | None = None, + reverse: bool = False, + ) -> list[VectorNode]: + """Return a list of vector nodes matching the provided filters and limit. + + Args: + filters: Dictionary of filter conditions to match vectors + limit: Maximum number of vectors to return + sort_key: Key to sort the results by (e.g., field name in metadata). None for no sorting + reverse: If True, sort in descending order; if False, sort in ascending order + """ + pool = await self._get_pool() + filter_clause, filter_params = self._build_filter_clause(filters) + + order_clause = "" + if sort_key: + order_direction = "DESC" if reverse else "ASC" + order_clause = f"ORDER BY metadata->>'{sort_key}' {order_direction}" + + limit_clause = "" + if limit: + limit_clause = f"LIMIT ${len(filter_params) + 1}" + filter_params.append(limit) + + async with pool.acquire() as conn: + sql = f""" + SELECT id, content, vector, metadata + FROM {self._qualified_name} + {filter_clause} + {order_clause} + {limit_clause} + """ + rows = await conn.fetch(sql, *filter_params) + + results = [] + for row in rows: + vector_data = self._pg_array_to_vector(row["vector"]) + + metadata = row["metadata"] if row["metadata"] else {} + if isinstance(metadata, str): + metadata = json.loads(metadata) + + results.append( + VectorNode( + vector_id=row["id"], + content=row["content"] or "", + vector=vector_data, + metadata=metadata, + ), + ) + + return results + + async def collection_info(self) -> dict[str, Any]: + """Fetch metadata including record count and disk usage for the collection.""" + pool = await self._get_pool() + qualified = self._qualified_name + + async with pool.acquire() as conn: + count = await conn.fetchval(f"SELECT COUNT(*) FROM {qualified}") + size = await conn.fetchval(f"SELECT pg_size_pretty(pg_total_relation_size('{qualified}'))") + + return { + "name": qualified, + "count": count, + "size": size, + } + + async def reset(self): + """Purge all data by dropping and recreating the collection table.""" + logger.warning(f"Resetting collection {self._qualified_name}...") + await self.delete_collection(self.collection_name) + await self.create_collection(self.collection_name) + + async def reset_collection(self, collection_name: str): + """Reset collection with table name validation.""" + self._validate_table_name(collection_name) + self.collection_name = collection_name + await self.create_collection(collection_name) + logger.info(f"Collection reset to {self._qualified_name}") + + async def start(self) -> None: + """Initialize the PGVector store. + + Creates the connection pool and ensures the collection table exists. + """ + await self._get_pool() + await super().start() + logger.info(f"Hologres collection {self._qualified_name} initialized") + + async def close(self): + """Terminate the database connection pool.""" + if self._pool is not None: + await self._pool.close() + self._pool = None + logger.info("Hologres connection pool closed") diff --git a/tests/test_vector_store.py b/tests/test_vector_store.py index 9b39ab9b..346be967 100644 --- a/tests/test_vector_store.py +++ b/tests/test_vector_store.py @@ -2,8 +2,8 @@ """Unified test suite for vector store implementations. This module provides comprehensive test coverage for LocalVectorStore, ESVectorStore, -PGVectorStore, QdrantVectorStore, ChromaVectorStore, ObVecVectorStore and ZvecVectorStore implementations. -Tests can be run for specific vector stores or all implementations. +PGVectorStore, QdrantVectorStore, ChromaVectorStore, ObVecVectorStore, HologresVectorStore, and +ZvecVectorStore implementations. Tests can be run for specific vector stores or all implementations. Usage: python test_vector_store.py --local # Test LocalVectorStore only @@ -12,6 +12,7 @@ Usage: python test_vector_store.py --qdrant # Test QdrantVectorStore only python test_vector_store.py --chroma # Test ChromaVectorStore only python test_vector_store.py --obvec # Test ObVecVectorStore only (needs seekdb / OceanBase) + python test_vector_store.py --hologres # Test HologresVectorStore only python test_vector_store.py --zvec # Test ZvecVectorStore only python test_vector_store.py --all # Test all vector stores """ @@ -32,6 +33,7 @@ from reme.core.utils import load_env, cosine_similarity from reme.core.vector_store import ( BaseVectorStore, ChromaVectorStore, + HologresVectorStore, LocalVectorStore, ESVectorStore, ObVecVectorStore, @@ -92,6 +94,19 @@ class TestConfig: OBVEC_PASSWORD = os.environ.get("OBVEC_PASSWORD", "root") OBVEC_DATABASE = os.environ.get("OBVEC_DATABASE", "test") + # HologresVectorStore settings + HOLOGRES_DSN = os.environ.get( + "HOLOGRES_DSN", + "", + ) # Full DSN connection string (overrides host/port/database/user/password) + HOLOGRES_HOST = os.environ.get("HOLOGRES_HOST", "localhost") + HOLOGRES_PORT = int(os.environ.get("HOLOGRES_PORT", "80")) + HOLOGRES_DATABASE = os.environ.get("HOLOGRES_DATABASE", "postgres") + HOLOGRES_USER = os.environ.get("HOLOGRES_USER", "postgres") + HOLOGRES_PASSWORD = os.environ.get("HOLOGRES_PASSWORD", "") + HOLOGRES_SCHEMA = os.environ.get("HOLOGRES_SCHEMA", "public") + HOLOGRES_MIN_SIZE = 1 + HOLOGRES_MAX_SIZE = 5 # ZvecVectorStore settings ZVEC_PATH = "./test_vector_store_zvec" # For local persistent mode @@ -205,7 +220,7 @@ def get_store_type(store: BaseVectorStore) -> str: store: Vector store instance Returns: - str: Type identifier ("local", "es", "pgvector", "qdrant", "chroma", "obvec", or "zvec") + str: Type identifier ("local", "es", "pgvector", "qdrant", "chroma", "obvec", "zvec", or "hologres") """ if isinstance(store, LocalVectorStore): return "local" @@ -221,6 +236,8 @@ def get_store_type(store: BaseVectorStore) -> str: return "obvec" elif isinstance(store, ZvecVectorStore): return "zvec" + elif isinstance(store, HologresVectorStore): + return "hologres" else: raise ValueError(f"Unknown vector store type: {type(store)}") @@ -230,7 +247,7 @@ def create_vector_store(store_type: str, collection_name: str) -> BaseVectorStor """Create a vector store instance based on type. Args: - store_type: Type of vector store ("local", "es", "pgvector", "qdrant", "chroma", or "obvec") + store_type: Type of vector store ("local", "es", "pgvector", "qdrant", "chroma", "obvec", or "hologres") collection_name: Name of the collection Returns: @@ -312,6 +329,23 @@ def create_vector_store(store_type: str, collection_name: str) -> BaseVectorStor dimension=config.EMBEDDING_DIMENSIONS, distance="cosine", ) + elif store_type == "hologres": + kwargs = { + "collection_name": collection_name, + "embedding_model": embedding_model, + "db_path": tempfile.mkdtemp(prefix="test_hologres_"), + "host": config.HOLOGRES_HOST, + "port": config.HOLOGRES_PORT, + "database": config.HOLOGRES_DATABASE, + "user": config.HOLOGRES_USER, + "password": config.HOLOGRES_PASSWORD, + "schema": config.HOLOGRES_SCHEMA, + "min_size": config.HOLOGRES_MIN_SIZE, + "max_size": config.HOLOGRES_MAX_SIZE, + } + if config.HOLOGRES_DSN: + kwargs["dsn"] = config.HOLOGRES_DSN + return HologresVectorStore(**kwargs) else: raise ValueError(f"Unknown store type: {store_type}") @@ -631,7 +665,7 @@ async def test_copy_collection(store: BaseVectorStore, store_name: str): # Elasticsearch, PostgreSQL and OceanBase require lowercase table/index names store_type = get_store_type(store) - if store_type in ("es", "pgvector", "obvec"): + if store_type in ("es", "pgvector", "obvec", "hologres"): copy_collection_name = copy_collection_name.lower() # Clean up if exists @@ -1835,6 +1869,7 @@ Examples: python test_vector_store.py --qdrant # Test QdrantVectorStore only python test_vector_store.py --chroma # Test ChromaVectorStore only python test_vector_store.py --obvec # Test ObVecVectorStore (seekdb / OceanBase) + python test_vector_store.py --hologres # Test HologresVectorStore python test_vector_store.py --all # Test all vector stores """, ) @@ -1868,6 +1903,11 @@ Examples: action="store_true", help="Test ObVecVectorStore", ) + parser.add_argument( + "--hologres", + action="store_true", + help="Test HologresVectorStore", + ) parser.add_argument( "--zvec", action="store_true", @@ -1892,6 +1932,7 @@ Examples: ("qdrant", "QdrantVectorStore"), ("chroma", "ChromaVectorStore"), ("obvec", "ObVecVectorStore"), + ("hologres", "HologresVectorStore"), ("zvec", "ZvecVectorStore"), ] else: @@ -1908,6 +1949,8 @@ Examples: stores_to_test.append(("chroma", "ChromaVectorStore")) if args.obvec: stores_to_test.append(("obvec", "ObVecVectorStore")) + if args.hologres: + stores_to_test.append(("hologres", "HologresVectorStore")) if args.zvec: stores_to_test.append(("zvec", "ZvecVectorStore")) @@ -1920,11 +1963,12 @@ Examples: ("qdrant", "QdrantVectorStore"), ("chroma", "ChromaVectorStore"), ("obvec", "ObVecVectorStore"), + ("hologres", "HologresVectorStore"), ("zvec", "ZvecVectorStore"), ] print("No vector store specified, defaulting to test all vector stores") print( - "Use --local/--es/--pgvector/--qdrant/--chroma/--obvec/--zvec to test specific ones\n", + "Use --local/--es/--pgvector/--qdrant/--chroma/--obvec/--zvec/--hologres to test specific ones\n", ) # Run tests for each vector store