feat(core): add ReMe V3 implementation with optional MCP client and enhanced filtering

This commit is contained in:
jinli.yl 2026-01-17 01:15:49 +08:00
parent 4d312ea682
commit e6ad682ede
32 changed files with 2369 additions and 106 deletions

View file

@ -29,6 +29,7 @@ class UserStats:
dialogues_per_session: list[int] # 每个 session 的对话数量
dialogue_lengths_per_session: list[int] # 每个 session 的对话总长度(字符数)
num_chunks_after_split: int # 按 5000 字符分割后的 chunk 数量
session_time_ranges: list[tuple[Any, Any]] # 每个 session 的 (开始时间, 结束时间)
@dataclass
@ -169,6 +170,7 @@ class DatasetAnalyzer:
dialogues_per_session = []
dialogue_lengths_per_session = []
session_time_ranges = []
total_chunks = 0
for session in sessions:
@ -179,6 +181,11 @@ class DatasetAnalyzer:
dialogues_per_session.append(num_dialogues)
dialogue_lengths_per_session.append(dialogue_length)
# 收集 session 的时间范围
start_time = session.get("start_time", None)
end_time = session.get("end_time", None)
session_time_ranges.append((start_time, end_time))
# 计算这个 session 分割后的 chunk 数量
num_chunks = self.split_session_into_chunks(dialogue, max_length=5000)
total_chunks += num_chunks
@ -202,7 +209,8 @@ class DatasetAnalyzer:
num_sessions=len(sessions),
dialogues_per_session=dialogues_per_session,
dialogue_lengths_per_session=dialogue_lengths_per_session,
num_chunks_after_split=total_chunks
num_chunks_after_split=total_chunks,
session_time_ranges=session_time_ranges
)
self.user_stats_list.append(user_stats)
@ -392,6 +400,32 @@ class DatasetAnalyzer:
print(f" 平均每 Session 对话长度: {avg_length:.2f} 字符")
print()
def print_first_user_session_times(self):
"""打印第一个用户的每个 session 的时间范围"""
if not self.user_stats_list:
print("\n没有用户数据")
return
first_user = self.user_stats_list[0]
print("\n" + "=" * 80)
print(f"第一个用户的 Session 时间统计")
print("=" * 80 + "\n")
print(f"用户名: {first_user.user_name}")
print(f"UUID: {first_user.uuid}")
print(f"总 Session 数: {first_user.num_sessions}\n")
print("-" * 80)
print(f"{'Session #':<12} {'开始时间':<30} {'结束时间':<30}")
print("-" * 80)
for idx, (start_time, end_time) in enumerate(first_user.session_time_ranges, 1):
start_str = str(start_time) if start_time is not None else "无"
end_str = str(end_time) if end_time is not None else "无"
print(f"{idx:<12} {start_str:<30} {end_str:<30}")
print("=" * 80)
def print_user_split_summary(self):
"""打印每个用户的分割统计摘要(表格形式)"""
print("\n" + "=" * 80)
@ -469,7 +503,11 @@ class DatasetAnalyzer:
if u.dialogue_lengths_per_session else 0
),
"dialogues_per_session": u.dialogues_per_session,
"dialogue_lengths_per_session": u.dialogue_lengths_per_session
"dialogue_lengths_per_session": u.dialogue_lengths_per_session,
"session_time_ranges": [
{"start_time": start, "end_time": end}
for start, end in u.session_time_ranges
]
}
for u in self.user_stats_list
]
@ -498,6 +536,9 @@ def main(data_path: str, output_path: str = None, show_per_user: bool = False):
# 打印摘要
analyzer.print_summary(stats)
# 打印第一个用户的 session 时间统计
analyzer.print_first_user_session_times()
# 打印每个用户的分割统计摘要(始终显示)
analyzer.print_user_split_summary()

View file

@ -0,0 +1,668 @@
"""
HaluMem Benchmark Evaluator for ReMe V3 - Question Answering
A modular evaluation pipeline that:
1. Loads HaluMem benchmark data
2. Processes user sessions through ReMe V3 (summarization + retrieval)
3. Evaluates question answering performance
4. Generates comprehensive metrics
Usage:
python bench/halumem/eval_reme_simple_v3.py \
--data_path /Users/yuli/workspace/HaluMem/data/HaluMem-Medium.jsonl \
--top_k 20 --user_num 100 --max_concurrency 20
"""
import asyncio
import json
import os
import re
import time
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
from loguru import logger
from eval_tools import evaluation_for_question2
from reme_ai.core.enumeration import MemoryType
from reme_ai.core.schema import MemoryNode
from reme_ai.reme import ReMe
# ==================== Configuration ====================
@dataclass
class EvalConfig:
"""Evaluation configuration parameters."""
data_path: str
top_k: int = 20
user_num: int = 1
max_concurrency: int = 2
batch_size: int = 20
output_dir: str = "bench_results/reme_simple_v3"
# ==================== Utilities ====================
class DataLoader:
"""Handles loading and parsing of HaluMem data."""
@staticmethod
def load_jsonl(file_path: str) -> list[dict]:
"""Load all entries from a JSONL file."""
with open(file_path, "r", encoding="utf-8") as f:
return [json.loads(line.strip()) for line in f if line.strip()]
@staticmethod
def extract_user_name(persona_info: str) -> str:
"""Extract user name from persona info string."""
match = re.search(r"Name:\s*(.*?); Gender:", persona_info)
if not match:
raise ValueError(f"No name found in persona_info: {persona_info}")
return match.group(1).strip()
@staticmethod
def format_dialogue_messages(dialogue: list[dict]) -> list[dict]:
"""Format dialogue into ReMe message format with conversation_time (user messages only)."""
return [
{
"role": turn["role"],
"content": turn["content"],
"time_created": datetime.strptime(
turn["timestamp"], "%b %d, %Y, %H:%M:%S"
)
.replace(tzinfo=timezone.utc)
.strftime("%Y-%m-%d %H:%M:%S"),
}
for turn in dialogue
if turn["role"] == "user" # Only include user messages
]
@staticmethod
def format_dialogue_for_eval(dialogue: list[dict], user_name: str = None) -> str:
"""Format dialogue into string for evaluation."""
formatted_turns = []
for turn in dialogue:
timestamp = datetime.strptime(
turn["timestamp"], "%b %d, %Y, %H:%M:%S"
).replace(tzinfo=timezone.utc).strftime("%Y-%m-%d %H:%M:%S")
# Use user_name if role is 'user' and user_name is provided
role = user_name if turn['role'] == 'user' and user_name else turn['role']
formatted_turns.append(
f"Role: {role}\n"
f"Content: {turn['content']}\n"
f"Time: {timestamp}"
)
return "\n\n".join(formatted_turns)
class FileManager:
"""Manages file I/O operations."""
def __init__(self, base_dir: str):
self.base_dir = Path(base_dir)
self.tmp_dir = self.base_dir / "tmp"
self.tmp_dir.mkdir(parents=True, exist_ok=True)
def get_user_dir(self, user_name: str) -> Path:
"""Get the directory path for a user."""
user_dir = self.tmp_dir / user_name
user_dir.mkdir(parents=True, exist_ok=True)
return user_dir
def get_session_file(self, user_name: str, session_id: int) -> Path:
"""Get the file path for a specific session."""
return self.get_user_dir(user_name) / f"session_{session_id}.json"
def save_session(self, user_name: str, session_id: int, data: dict):
"""Save session data to file."""
file_path = self.get_session_file(user_name, session_id)
with open(file_path, "w", encoding="utf-8") as f:
json.dump(data, f, ensure_ascii=False, indent=2)
logger.info(f"✅ Saved session {session_id} to {file_path}")
def load_session(self, user_name: str, session_id: int) -> dict | None:
"""Load session data from file."""
file_path = self.get_session_file(user_name, session_id)
if not file_path.exists():
return None
with open(file_path, "r", encoding="utf-8") as f:
return json.load(f)
def user_has_cache(self, user_name: str) -> bool:
"""Check if user has cached results."""
user_dir = self.get_user_dir(user_name)
return any(f.name.startswith("session_") and f.suffix == ".json"
for f in user_dir.iterdir())
def combine_results(self, output_file: str):
"""Combine all user session files into a single JSONL file."""
with open(output_file, "w", encoding="utf-8") as f_out:
for user_dir in self.tmp_dir.iterdir():
if not user_dir.is_dir():
continue
session_files = sorted([
f for f in user_dir.iterdir()
if f.name.startswith("session_") and f.suffix == ".json"
])
if not session_files:
continue
# Load first session to get user metadata
with open(session_files[0], "r", encoding="utf-8") as f_in:
first_session = json.load(f_in)
user_data = {
"uuid": first_session["uuid"],
"user_name": first_session["user_name"],
"sessions": []
}
# Load all sessions
for session_file in session_files:
with open(session_file, "r", encoding="utf-8") as f_in:
session_data = json.load(f_in)
# Remove redundant user metadata
session_data.pop("uuid", None)
session_data.pop("user_name", None)
user_data["sessions"].append(session_data)
f_out.write(json.dumps(user_data, ensure_ascii=False) + "\n")
# ==================== Memory Operations ====================
class MemoryProcessor:
"""Handles ReMe V3 memory operations."""
def __init__(self, reme: ReMe):
self.reme = reme
async def add_memories(
self,
user_id: str,
messages: list[dict],
batch_size: int = 10000
) -> tuple[list[str], list[list[dict]], float]:
"""
Add memories in batches using ReMe V3 and return extracted memory contents.
Returns:
tuple: (extracted_memories, agent_messages, total_duration_ms)
"""
added_memories: list[MemoryNode] = []
deleted_memories: list[str] = []
all_agent_messages: list = []
total_duration_ms = 0
for i in range(0, len(messages), batch_size):
batch = messages[i:i + batch_size]
start = time.time()
# Use summary_v3 instead of summary_v2
memory_nodes, agent_messages, success = await self.reme.summary_v3(
messages=batch,
user_id=user_id
)
duration_ms = (time.time() - start) * 1000
total_duration_ms += duration_ms
# Save agent messages for this batch
if agent_messages:
all_agent_messages.extend(agent_messages)
if memory_nodes:
for node in memory_nodes:
if isinstance(node, MemoryNode) and node.memory_type == MemoryType.HISTORY:
continue
if isinstance(node, MemoryNode):
added_memories.append(node)
if isinstance(node, str):
deleted_memories.append(node)
extracted_memories = deleted_memories
extracted_memories += ["[delete]" + n.format_memory() for n in added_memories if n.memory_id in deleted_memories]
extracted_memories += ["[add]" + n.format_memory() for n in added_memories if n.memory_id not in deleted_memories]
return extracted_memories, all_agent_messages, total_duration_ms
async def search_memory(
self,
query: str,
user_id: str,
top_k: int = 20
) -> tuple[str, list, float]:
"""
Search memory using ReMe V3 and return response.
Returns:
tuple: (response, agent_messages, duration_ms)
"""
start = time.time()
# Use retrieve_v3 instead of retrieve_v2
response, agent_messages, success = await self.reme.retrieve_v3(
query=query,
user_id=user_id,
top_k=top_k
)
duration_ms = (time.time() - start) * 1000
return response, agent_messages, duration_ms
# ==================== Evaluation ====================
class QuestionAnsweringEvaluator:
"""Evaluates question answering performance."""
def __init__(self, memory_processor: MemoryProcessor, top_k: int):
self.memory_processor = memory_processor
self.top_k = top_k
async def evaluate_questions(
self,
questions: list[dict],
user_name: str,
uuid: str,
session_id: int,
formatted_dialogue: str
) -> list[dict]:
"""Evaluate all questions for a session."""
results = []
for qa in questions:
# Search memory for answer using V3
response, agent_messages, duration_ms = await self.memory_processor.search_memory(
query=qa["question"],
user_id=user_name,
top_k=self.top_k
)
# Evaluate response
evidence_text = "\n".join([e["memory_content"] for e in qa["evidence"]])
eval_result = await evaluation_for_question2(
qa["question"],
qa["answer"],
evidence_text,
response,
formatted_dialogue
)
# Build result record
qa_result = {
**qa,
"uuid": uuid,
"session_id": session_id,
"system_response": response,
"retrieve_messages": [m.model_dump() for m in agent_messages],
"search_duration_ms": duration_ms,
"result_type": eval_result.get("evaluation_result"),
"question_answering_reasoning": eval_result.get("reasoning", "")
}
results.append(qa_result)
return results
class MetricsAggregator:
"""Aggregates evaluation metrics."""
@staticmethod
def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]:
"""Compute question answering metrics."""
total = len(qa_records)
if total == 0:
return {
"correct_qa_ratio(all)": 0,
"hallucination_qa_ratio(all)": 0,
"omission_qa_ratio(all)": 0,
"correct_qa_ratio(valid)": 0,
"hallucination_qa_ratio(valid)": 0,
"omission_qa_ratio(valid)": 0,
"qa_valid_num": 0,
"qa_num": 0
}
correct = 0
hallucination = 0
omission = 0
valid = 0
for qa in qa_records:
result_type = qa.get("result_type", "")
if result_type in ["Correct", "Hallucination", "Omission"]:
valid += 1
if result_type == "Correct":
correct += 1
elif result_type == "Hallucination":
hallucination += 1
elif result_type == "Omission":
omission += 1
metrics = {
"correct_qa_ratio(all)": correct / total,
"hallucination_qa_ratio(all)": hallucination / total,
"omission_qa_ratio(all)": omission / total,
"qa_valid_num": valid,
"qa_num": total
}
if valid > 0:
metrics.update({
"correct_qa_ratio(valid)": correct / valid,
"hallucination_qa_ratio(valid)": hallucination / valid,
"omission_qa_ratio(valid)": omission / valid
})
else:
metrics.update({
"correct_qa_ratio(valid)": 0,
"hallucination_qa_ratio(valid)": 0,
"omission_qa_ratio(valid)": 0
})
return metrics
@staticmethod
def compute_time_metrics(eval_results_file: str) -> dict[str, float]:
"""Compute timing metrics from evaluation results."""
add_duration = 0
search_duration = 0
with open(eval_results_file, "r", encoding="utf-8") as f:
for line in f:
if not line.strip():
continue
user_data = json.loads(line)
for session in user_data["sessions"]:
add_duration += session.get("add_dialogue_duration_ms", 0)
eval_results = session.get("evaluation_results", {})
for qa in eval_results.get("question_answering_records", []):
search_duration += qa.get("search_duration_ms", 0)
# Convert to minutes
return {
"add_dialogue_duration_time": add_duration / 1000 / 60,
"search_memory_duration_time": search_duration / 1000 / 60,
"total_duration_time": (add_duration + search_duration) / 1000 / 60
}
# ==================== Main Pipeline ====================
class HaluMemEvaluatorV3:
"""Main evaluator orchestrating the entire ReMe V3 pipeline."""
def __init__(self, config: EvalConfig):
self.config = config
self.reme = ReMe()
self.file_manager = FileManager(config.output_dir)
self.memory_processor = MemoryProcessor(self.reme)
self.qa_evaluator = QuestionAnsweringEvaluator(
self.memory_processor,
config.top_k
)
self.data_loader = DataLoader()
async def process_session(
self,
session: dict,
session_id: int,
user_name: str,
uuid: str
) -> dict:
"""Process a single session using ReMe V3."""
session_data = {
"uuid": uuid,
"user_name": user_name,
"session_id": session_id,
"memory_points": session["memory_points"]
}
# Skip generated QA sessions
if session.get("is_generated_qa_session", False):
session_data["is_generated_qa_session"] = True
return session_data
# Format and add dialogue to memory using V3
dialogue = session["dialogue"]
formatted_messages = self.data_loader.format_dialogue_messages(dialogue)
extracted_memories, agent_messages, duration_ms = await self.memory_processor.add_memories(
user_id=user_name,
messages=formatted_messages,
batch_size=self.config.batch_size
)
session_data.update({
"dialogue": dialogue,
"extracted_memories": extracted_memories,
"summary_messages": [m.model_dump() for m in agent_messages],
"add_dialogue_duration_ms": duration_ms
})
# Evaluate questions if present
if "questions" in session:
formatted_dialogue = self.data_loader.format_dialogue_for_eval(dialogue, user_name)
qa_results = await self.qa_evaluator.evaluate_questions(
questions=session["questions"],
user_name=user_name,
uuid=uuid,
session_id=session_id,
formatted_dialogue=formatted_dialogue
)
session_data["evaluation_results"] = {
"question_answering_records": qa_results
}
return session_data
async def process_user(self, user_data: dict) -> dict:
"""Process all sessions for a user."""
user_name = self.data_loader.extract_user_name(user_data["persona_info"])
uuid = user_data["uuid"]
logger.info(f"Processing user: {user_name}")
for idx, session in enumerate(user_data["sessions"]):
logger.info(f" Session {idx + 1}/{len(user_data['sessions'])}")
session_data = await self.process_session(
session=session,
session_id=idx,
user_name=user_name,
uuid=uuid
)
self.file_manager.save_session(user_name, idx, session_data)
return {"uuid": uuid, "user_name": user_name, "status": "ok"}
async def run_evaluation(self):
"""Run the complete evaluation pipeline using ReMe V3."""
start_time = time.time()
# Clear existing data
await self.reme.vector_store.delete_all()
# Load user data
all_users = self.data_loader.load_jsonl(self.config.data_path)
users_to_process = all_users[:self.config.user_num]
print("\n" + "=" * 80)
print("HALUMEM EVALUATION - REME V3 - QUESTION ANSWERING")
print(f"Users: {len(users_to_process)} | Concurrency: {self.config.max_concurrency}")
print("=" * 80 + "\n")
# Process users with concurrency control
semaphore = asyncio.Semaphore(self.config.max_concurrency)
async def process_with_cache_check(idx: int, user_data: dict):
async with semaphore:
user_name = self.data_loader.extract_user_name(user_data["persona_info"])
# Check cache
if self.file_manager.user_has_cache(user_name):
print(f"⚡ [{idx}/{len(users_to_process)}] Skipping {user_name} (cached)")
return {"user_name": user_name, "status": "cached"}
print(f"🔄 [{idx}/{len(users_to_process)}] Processing {user_name}...")
result = await self.process_user(user_data)
print(f"✅ [{idx}/{len(users_to_process)}] Completed {user_name}")
return result
tasks = [
process_with_cache_check(idx, user)
for idx, user in enumerate(users_to_process, 1)
]
await asyncio.gather(*tasks)
# Combine results
output_file = os.path.join(self.config.output_dir, "eval_results.jsonl")
self.file_manager.combine_results(output_file)
elapsed = time.time() - start_time
print(f"\n✅ Processing completed in {elapsed:.2f}s")
print(f"📁 Results: {output_file}\n")
# Aggregate metrics
await self.aggregate_and_report(output_file)
async def aggregate_and_report(self, results_file: str):
"""Aggregate results and generate final report."""
print("=" * 80)
print("AGGREGATING METRICS")
print("=" * 80 + "\n")
# Collect all QA records
qa_records = []
with open(results_file, "r", encoding="utf-8") as f:
for line in f:
if not line.strip():
continue
user_data = json.loads(line)
for session in user_data["sessions"]:
if session.get("is_generated_qa_session"):
continue
eval_results = session.get("evaluation_results", {})
qa_records.extend(
eval_results.get("question_answering_records", [])
)
# Compute metrics
qa_metrics = MetricsAggregator.compute_qa_metrics(qa_records)
time_metrics = MetricsAggregator.compute_time_metrics(results_file)
final_results = {
"overall_score": {
"question_answering": qa_metrics,
"time_consuming": time_metrics
},
"question_answering_records": qa_records
}
# Save final report
report_file = os.path.join(self.config.output_dir, "eval_statistics.json")
with open(report_file, "w", encoding="utf-8") as f:
json.dump(final_results, f, ensure_ascii=False, indent=4)
print(f"📊 Statistics saved to: {report_file}\n")
# Print summary
self._print_summary(qa_metrics, time_metrics)
def _print_summary(self, qa_metrics: dict, time_metrics: dict):
"""Print evaluation summary."""
print("=" * 80)
print("EVALUATION SUMMARY - REME V3")
print("=" * 80 + "\n")
print("📊 Question Answering:")
print(f" Correct (all): {qa_metrics['correct_qa_ratio(all)']:.4f}")
print(f" Hallucination (all): {qa_metrics['hallucination_qa_ratio(all)']:.4f}")
print(f" Omission (all): {qa_metrics['omission_qa_ratio(all)']:.4f}")
print(f" Correct (valid): {qa_metrics['correct_qa_ratio(valid)']:.4f}")
print(f" Hallucination (valid): {qa_metrics['hallucination_qa_ratio(valid)']:.4f}")
print(f" Omission (valid): {qa_metrics['omission_qa_ratio(valid)']:.4f}")
print(f" Valid/Total: {qa_metrics['qa_valid_num']}/{qa_metrics['qa_num']}")
print(f"\n⏱️ Time Metrics:")
print(f" Memory Addition: {time_metrics['add_dialogue_duration_time']:.2f} min")
print(f" Memory Search: {time_metrics['search_memory_duration_time']:.2f} min")
print(f" Total: {time_metrics['total_duration_time']:.2f} min")
print("\n" + "=" * 80)
# ==================== Entry Point ====================
def main(
data_path: str,
top_k: int = 20,
user_num: int = 1,
max_concurrency: int = 2
):
"""Main entry point for ReMe V3 evaluation."""
config = EvalConfig(
data_path=data_path,
top_k=top_k,
user_num=user_num,
max_concurrency=max_concurrency
)
evaluator = HaluMemEvaluatorV3(config)
asyncio.run(evaluator.run_evaluation())
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(
description="Evaluate ReMe V3 on HaluMem benchmark (Question Answering)"
)
parser.add_argument(
"--data_path",
type=str,
required=True,
help="Path to HaluMem JSONL file"
)
parser.add_argument(
"--top_k",
type=int,
default=20,
help="Number of memories to retrieve (default: 20)"
)
parser.add_argument(
"--user_num",
type=int,
default=1,
help="Number of users to evaluate (default: 1)"
)
parser.add_argument(
"--max_concurrency",
type=int,
default=2,
help="Maximum concurrent user processing (default: 2)"
)
args = parser.parse_args()
main(
data_path=args.data_path,
top_k=args.top_k,
user_num=args.user_num,
max_concurrency=args.max_concurrency
)

View file

@ -55,7 +55,7 @@ class PromptHandler(BaseContext):
key += "_" + self.language.strip()
assert key in self, f"prompt_name={key} not found."
return self[key]
return self[key].strip()
def prompt_format(self, prompt_name: str, **kwargs) -> str:
"""Format a prompt by filtering flagged lines and filling template variables."""

View file

@ -41,14 +41,14 @@ class ToolAttr(BaseModel):
if self.enum:
res["enum"] = self.enum
if self.type == "object" and self.properties:
if self.type == "object" and self.properties is not None:
res["properties"] = {
k: v.simple_input_dump() if isinstance(v, ToolAttr) else v for k, v in self.properties.items()
}
if self.required:
if self.required is not None:
res["required"] = self.required
if self.type == "array" and self.items:
if self.type == "array" and self.items is not None:
res["items"] = self.items.simple_input_dump() if isinstance(self.items, ToolAttr) else self.items
return res

View file

@ -9,7 +9,15 @@ from .http_client import HttpClient
from .llm_utils import extract_content, format_messages, deduplicate_memories
from .logger_utils import init_logger
from .logo_utils import print_logo
from .mcp_client import MCPClient
# Make MCPClient import optional to avoid breaking if MCP dependencies are not available
try:
from .mcp_client import MCPClient
_HAS_MCP = True
except ImportError:
MCPClient = None
_HAS_MCP = False
from .pydantic_config_parser import PydanticConfigParser
from .pydantic_utils import create_pydantic_model
from .singleton import singleton

View file

@ -15,7 +15,7 @@ class CacheHandler:
_EXTENSIONS = {
pd.DataFrame: ".csv",
dict: ".json",
list: ".json",
list: ".jsonl",
str: ".txt",
}
@ -76,11 +76,17 @@ class CacheHandler:
data.to_csv(path, index=kwargs.get("index", False), encoding="utf-8")
return {"row_count": len(data), "file_size": path.stat().st_size}
if dtype in (dict, list):
if dtype is dict:
with open(path, "w", encoding="utf-8") as f:
json.dump(data, f, ensure_ascii=False, indent=2)
return {"item_count": len(data), "file_size": path.stat().st_size}
if dtype is list:
with open(path, "w", encoding="utf-8") as f:
for item in data:
f.write(json.dumps(item, ensure_ascii=False) + "\n")
return {"item_count": len(data), "file_size": path.stat().st_size}
if dtype is str:
path.write_text(data, encoding=kwargs.get("encoding", "utf-8"))
return {"char_count": len(data), "file_size": path.stat().st_size}
@ -92,9 +98,17 @@ class CacheHandler:
"""Execute type-specific load operations."""
if type_name == "DataFrame":
return pd.read_csv(path, encoding=kwargs.get("encoding", "utf-8"))
if type_name in ("dict", "list"):
if type_name == "dict":
with open(path, "r", encoding="utf-8") as f:
return json.load(f)
if type_name == "list":
result = []
with open(path, "r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if line:
result.append(json.loads(line))
return result
if type_name == "str":
return path.read_text(encoding=kwargs.get("encoding", "utf-8"))
raise ValueError(f"Unknown data type in metadata: {type_name}")

View file

@ -117,14 +117,33 @@ class ChromaVectorStore(BaseVectorStore):
@staticmethod
def _generate_where_clause(filters: dict | None) -> dict | None:
"""Convert the universal filter format to a ChromaDB-compatible where clause."""
"""Convert the universal filter format to a ChromaDB-compatible where clause.
Supports two filter formats:
1. Range query: {"field": [start_value, end_value]} - filters for field >= start_value AND field <= end_value
2. Exact match: {"field": value} - filters for field == value
"""
if not filters:
return None
def convert_condition(k: str, v: Any) -> dict | None:
"""Convert a single filter condition to ChromaDB operator format."""
def convert_condition(k: str, v: Any) -> dict | list | None:
"""Convert a single filter condition to ChromaDB operator format.
Returns:
- dict for simple conditions
- list of dicts for range queries (which need to be wrapped in $and)
- None for wildcard filters
"""
if v == "*":
return None
# New syntax: [start, end] represents a range query
if isinstance(v, list) and len(v) == 2:
# Range query: field >= v[0] AND field <= v[1]
# ChromaDB requires separate conditions combined with $and
return [
{k: {"$gte": v[0]}},
{k: {"$lte": v[1]}}
]
if isinstance(v, dict):
chroma_condition = {}
for op, val in v.items():
@ -141,8 +160,7 @@ class ChromaVectorStore(BaseVectorStore):
chroma_op = mapping.get(op, "$eq")
chroma_condition[k] = {chroma_op: val}
return chroma_condition
if isinstance(v, list):
return {k: {"$in": v}}
# Exact match for non-list values
return {k: {"$eq": v}}
processed_filters = []
@ -155,7 +173,11 @@ class ChromaVectorStore(BaseVectorStore):
for sub_key, sub_value in condition.items():
converted = convert_condition(sub_key, sub_value)
if converted:
or_condition.update(converted)
if isinstance(converted, list):
# Range query in OR condition - need to wrap in $and
or_conditions.append({"$and": converted})
else:
or_condition.update(converted)
if or_condition:
or_conditions.append(or_condition)
if len(or_conditions) > 1:
@ -168,13 +190,21 @@ class ChromaVectorStore(BaseVectorStore):
for sub_key, sub_value in condition.items():
converted = convert_condition(sub_key, sub_value)
if converted:
processed_filters.append(converted)
if isinstance(converted, list):
# Range query - add each condition separately
processed_filters.extend(converted)
else:
processed_filters.append(converted)
elif key == "$not":
continue
else:
converted = convert_condition(key, value)
if converted:
processed_filters.append(converted)
if isinstance(converted, list):
# Range query - add each condition separately
processed_filters.extend(converted)
else:
processed_filters.append(converted)
if not processed_filters:
return None

View file

@ -262,9 +262,19 @@ class ESVectorStore(BaseVectorStore):
if filters:
filter_conditions = []
for key, value in filters.items():
if isinstance(value, list):
filter_conditions.append({"terms": {f"metadata.{key}": value}})
# New syntax: [start, end] represents a range query
if isinstance(value, list) and len(value) == 2:
# Range query: field >= value[0] AND field <= value[1]
filter_conditions.append({
"range": {
f"metadata.{key}": {
"gte": value[0],
"lte": value[1]
}
}
})
else:
# Exact match
filter_conditions.append({"term": {f"metadata.{key}": value}})
search_query["knn"]["filter"] = {"bool": {"must": filter_conditions}}
@ -448,9 +458,19 @@ class ESVectorStore(BaseVectorStore):
if filters:
filter_conditions = []
for key, value in filters.items():
if isinstance(value, list):
filter_conditions.append({"terms": {f"metadata.{key}": value}})
# New syntax: [start, end] represents a range query
if isinstance(value, list) and len(value) == 2:
# Range query: field >= value[0] AND field <= value[1]
filter_conditions.append({
"range": {
f"metadata.{key}": {
"gte": value[0],
"lte": value[1]
}
}
})
else:
# Exact match
filter_conditions.append({"term": {f"metadata.{key}": value}})
query["query"] = {"bool": {"must": filter_conditions}}

View file

@ -91,17 +91,32 @@ class LocalVectorStore(BaseVectorStore):
@staticmethod
def _match_filters(node: VectorNode, filters: dict | None) -> bool:
"""Check if a vector node matches the provided metadata filters."""
"""Check if a vector node matches the provided metadata filters.
Supports two filter formats:
1. Range query: {"field": [start_value, end_value]} - filters for field >= start_value AND field <= end_value
2. Exact match: {"field": value} - filters for field == value
"""
if not filters:
return True
for key, value in filters.items():
node_value = node.metadata.get(key)
if isinstance(value, list):
if node_value not in value:
# New syntax: [start, end] represents a range query
if isinstance(value, list) and len(value) == 2:
# Range query: field >= value[0] AND field <= value[1]
if node_value is None:
return False
try:
# Try numeric comparison
if not (value[0] <= node_value <= value[1]):
return False
except TypeError:
# If comparison fails, the filter doesn't match
return False
else:
# Exact match
if node_value != value:
return False

View file

@ -1,6 +1,7 @@
"""PostgreSQL pgvector implementation for vector storage and retrieval."""
import json
import re
from typing import Any
from loguru import logger
@ -25,6 +26,25 @@ except ImportError as e:
class PGVectorStore(BaseVectorStore):
"""Vector store implementation using PostgreSQL and pgvector for efficient similarity search."""
@staticmethod
def _validate_table_name(name: str) -> None:
"""Validate table name to prevent SQL injection.
PostgreSQL table names must:
- Contain only alphanumeric characters and underscores
- Not start with a digit
- Be between 1 and 63 characters
"""
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,
@ -47,6 +67,9 @@ class PGVectorStore(BaseVectorStore):
"PGVector requires extra dependencies. Install with `pip install asyncpg pgvector`",
) from _ASYNCPG_IMPORT_ERROR
# Validate collection name to prevent SQL injection
self._validate_table_name(collection_name)
super().__init__(collection_name=collection_name, embedding_model=embedding_model, **kwargs)
self.dsn = dsn
@ -106,6 +129,7 @@ class PGVectorStore(BaseVectorStore):
async def create_collection(self, collection_name: str, **kwargs):
"""Create a new PostgreSQL table with vector support and appropriate indexing."""
self._validate_table_name(collection_name)
pool = await self._get_pool()
dimensions = kwargs.get("dimensions", self.embedding_model_dims)
@ -150,6 +174,7 @@ class PGVectorStore(BaseVectorStore):
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()
async with pool.acquire() as conn:
await conn.execute(f"DROP TABLE IF EXISTS {collection_name}")
@ -157,6 +182,7 @@ class PGVectorStore(BaseVectorStore):
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()
async with pool.acquire() as conn:
@ -252,7 +278,14 @@ class PGVectorStore(BaseVectorStore):
@staticmethod
def _build_filter_clause(filters: dict | None) -> tuple[str, list]:
"""Generate an SQL WHERE clause and parameter list from a filter dictionary."""
"""Generate an SQL WHERE clause and parameter list from a filter dictionary.
Supports two filter formats:
1. Range query: {"field": [start_value, end_value]} - filters for field >= start_value AND field <= end_value
2. Exact match: {"field": value} - filters for field == value
Range queries support both numeric and string (e.g., timestamp strings) comparisons.
"""
if not filters:
return "", []
@ -261,12 +294,28 @@ class PGVectorStore(BaseVectorStore):
param_idx = 1
for key, value in filters.items():
if isinstance(value, list):
placeholders = ", ".join([f"${param_idx + i}" for i in range(len(value))])
conditions.append(f"metadata->>'{key}' IN ({placeholders})")
params.extend([str(v) for v in value])
param_idx += len(value)
# Sanitize key to prevent SQL injection (only allow alphanumeric and underscore)
if not key.replace('_', '').replace('.', '').isalnum():
raise ValueError(f"Invalid metadata key: {key}. Only alphanumeric characters, underscore and dot are allowed.")
# New syntax: [start, end] represents a range query
if isinstance(value, list) and len(value) == 2:
# Range query: field >= value[0] AND field <= value[1]
# Try numeric comparison first, fall back to text comparison if needed
if isinstance(value[0], (int, float)) and isinstance(value[1], (int, float)):
# Numeric range query
conditions.append(
f"(metadata->>'{key}')::numeric >= ${param_idx} AND (metadata->>'{key}')::numeric <= ${param_idx + 1}"
)
else:
# Text range query (works for strings, timestamps, etc.)
conditions.append(
f"metadata->>'{key}' >= ${param_idx} AND metadata->>'{key}' <= ${param_idx + 1}"
)
params.extend([value[0], value[1]])
param_idx += 2
else:
# Exact match
conditions.append(f"metadata->>'{key}' = ${param_idx}")
params.append(str(value))
param_idx += 1
@ -290,11 +339,14 @@ class PGVectorStore(BaseVectorStore):
filter_clause, filter_params = self._build_filter_clause(filters)
# Adjust parameter indices in filter clause to account for $1 being used by vector_str
if filter_clause:
for i in range(len(filter_params)):
old_idx = i + 1
new_idx = i + 2
filter_clause = filter_clause.replace(f"${old_idx}", f"${new_idx}", 1)
# Replace from highest index to lowest to avoid conflicts
for i in range(len(filter_params), 0, -1):
old_placeholder = f"${i}"
new_placeholder = f"${i + 1}"
# Use word boundary to ensure we only replace exact matches (e.g., $1 not $10)
filter_clause = re.sub(rf'\${i}\b', new_placeholder, filter_clause)
async with pool.acquire() as conn:
sql = f"""

View file

@ -246,29 +246,65 @@ class QdrantVectorStore(BaseVectorStore):
@staticmethod
def _create_filter(filters: dict) -> Filter | None:
"""Convert a dictionary of filter conditions into a Qdrant Filter object."""
"""Convert a dictionary of filter conditions into a Qdrant Filter object.
Supports two filter formats:
1. Range query: {"field": [start_value, end_value]} - filters for field >= start_value AND field <= end_value
2. Exact match: {"field": value} - filters for field == value
"""
if not filters:
return None
conditions = []
for key, value in filters.items():
if isinstance(value, dict) and ("gte" in value or "lte" in value):
# New syntax: [start, end] represents a range query
if isinstance(value, list) and len(value) == 2:
# Range query: field >= value[0] AND field <= value[1]
# Qdrant's Range only supports numeric values
if isinstance(value[0], (int, float)) and isinstance(value[1], (int, float)):
conditions.append(
FieldCondition(
key=f"metadata.{key}",
range=Range(gte=value[0], lte=value[1]),
),
)
else:
# For non-numeric values (e.g., string dates), Qdrant doesn't support range queries
# We need to skip this filter with a warning
logger.warning(
f"Qdrant does not support range queries for non-numeric values. "
f"Skipping range filter for key '{key}' with values {value}. "
f"Consider using numeric timestamps instead."
)
elif isinstance(value, dict) and ("gte" in value or "lte" in value):
range_params = {}
# Check if values are numeric
if "gte" in value:
range_params["gte"] = value["gte"]
if isinstance(value["gte"], (int, float)):
range_params["gte"] = value["gte"]
else:
logger.warning(
f"Qdrant range filter for key '{key}' requires numeric gte value, got {type(value['gte']).__name__}. Skipping."
)
continue
if "lte" in value:
range_params["lte"] = value["lte"]
conditions.append(
FieldCondition(
key=f"metadata.{key}",
range=Range(**range_params),
),
)
elif isinstance(value, list):
conditions.append(
FieldCondition(key=f"metadata.{key}", match=MatchValue(value=value[0])),
)
if isinstance(value["lte"], (int, float)):
range_params["lte"] = value["lte"]
else:
logger.warning(
f"Qdrant range filter for key '{key}' requires numeric lte value, got {type(value['lte']).__name__}. Skipping."
)
continue
if range_params: # Only add condition if we have valid numeric parameters
conditions.append(
FieldCondition(
key=f"metadata.{key}",
range=Range(**range_params),
),
)
else:
# Exact match
conditions.append(
FieldCondition(key=f"metadata.{key}", match=MatchValue(value=value)),
)

View file

@ -0,0 +1,9 @@
from .personal_summarizer_v3 import PersonalSummarizerV3
from .reme_retriever_v3 import ReMeRetrieverV3
from .reme_summarizer_v3 import ReMeSummarizerV3
__all__ = [
"PersonalSummarizerV3",
"ReMeRetrieverV3",
"ReMeSummarizerV3",
]

View file

@ -0,0 +1,69 @@
from ..base_memory_agent import BaseMemoryAgent
from ...core.enumeration import Role, MemoryType
from ...core.schema import Message, ToolCall
from ...core.utils import format_messages
class PersonalSummarizerV3(BaseMemoryAgent):
memory_type: MemoryType = MemoryType.PERSONAL
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": self.get_prompt("tool"),
"parameters": {
"type": "object",
"properties": {
"messages": {
"type": "array",
"items": {
"type": "object",
"properties": {
"role": {
"type": "string",
"description": "role",
},
"content": {
"type": "string",
"description": "content",
},
},
"required": ["role", "content"],
},
},
},
"required": ["messages"],
},
},
)
async def build_messages(self) -> list[Message]:
"""Construct messages with context, memory_target, and memory_type information."""
system_prompt = self.prompt_format(
prompt_name="system_prompt",
context=self.description + "\n" + format_messages(self.get_messages()),
memory_type=self.memory_type.value,
memory_target=self.memory_target,
)
messages = [
Message(role=Role.SYSTEM, content=system_prompt),
Message(role=Role.USER, content=self.get_prompt("user_message")),
]
return messages
async def _reasoning_step(self, messages: list[Message], step: int, **kwargs) -> tuple[Message, bool]:
return await super()._reasoning_step(messages, step, **kwargs)
async def _acting_step(self, assistant_message: Message, step: int, **kwargs) -> list[Message]:
"""Execute tool calls with memory_target, memory_type, and author context."""
messages: list[Message] = await super()._acting_step(
assistant_message,
step,
memory_type=self.memory_type.value,
memory_target=self.memory_target,
ref_memory_id=self.ref_memory_id,
author=self.author,
**kwargs,
)
return messages

View file

@ -0,0 +1,38 @@
tool: |
Extract and update personal memories about the user from conversation context.
Analyze dialogues to identify preferences, habits, background, relationships, and key facts.
system_prompt: |
You are a memory agent managing **{memory_type}** memories about **{memory_target}**.
## Latest Conversation:
{context}
Each message format: `round<index> [<timestamp>] <role/name>: <content>` (timestamp: YYYY-MM-DD HH:MM:SS).
**CRITICAL**: Extract ONLY explicitly stated information. DO NOT infer, assume, or fabricate.
## Three-Step Workflow
### Step 1: Extract Conversation Memories
Use `AddMemory` to extract key personal facts from the conversation.
- Extract: preferences, habits, status, personal details, decisions, conclusions
- Keep entries concise and distinct (no duplicates, no omissions)
- Record `conversation_time` for each memory (format: 2020-01-01 00:00:00; use 0000-00-00 00:00:00 if unavailable)
### Step 2: Read User Profile
Use `ReadUserProfile` to retrieve the current user profile.
- Review existing memories to identify conflicts and duplicates
### Step 3: Update User Profile
Use `UpdateUserProfile` to synchronize the profile with new information.
- `profile_ids_to_delete`: Remove outdated or conflicting profiles
- `profiles_to_add`: Add new profiles that are not duplicates
- Use `timestamp` from conversation_time (format: 2020-01-01 00:00:00)
- Keep final profiles concise with no information loss
user_message: |
Execute the three-step workflow:
1. Use `AddMemory` to extract personal memories from the conversation
2. Use `ReadUserProfile` to read existing user profile
3. Use `UpdateUserProfile` to remove outdated entries and add new profiles

View file

@ -0,0 +1,44 @@
"""ReMe retriever v2 that autonomously retrieves memories from multiple angles."""
from typing import List
from ..base_memory_agent import BaseMemoryAgent
from ...core.enumeration import Role
from ...core.schema import Message
from ...core.utils import format_messages
class ReMeRetrieverV3(BaseMemoryAgent):
def __init__(self, meta_memories: list[dict] | None = None, **kwargs):
super().__init__(**kwargs)
self.meta_memories: list[dict] = meta_memories or []
async def _read_meta_memories(self) -> str:
"""Fetch all meta-memory entries that define specialized memory agents."""
from ...mem_tool import ReadMetaMemory
op = ReadMetaMemory(enable_identity_memory=False)
return op.format_memory_metadata(self.meta_memories)
async def build_messages(self) -> List[Message]:
"""Build messages with system prompt and user message."""
if self.context.get("query"):
context = self.context.query
elif self.context.get("messages"):
context = format_messages(self.context.messages)
else:
raise ValueError("input must have either `query` or `messages`")
system_prompt = self.prompt_format(
prompt_name="system_prompt",
meta_memory_info=await self._read_meta_memories(),
context=context,
)
messages = [
Message(role=Role.SYSTEM, content=system_prompt),
Message(role=Role.USER, content=self.get_prompt("user_message")),
]
return messages

View file

@ -0,0 +1,53 @@
tool: |
Autonomously retrieve relevant memories through a three-step strategy to answer user questions.
Steps: read user profile → vector search with multiple angles → read original conversations.
State "I don't know" if information cannot be found after exhaustive searching.
NEVER hallucinate or fabricate information not present in retrieved memories.
system_prompt: |
You are a memory retrieval agent. Search for relevant memories to answer the user's question following this strategy:
## Available Meta Memories
Format: "- <memory_type>(<memory_target>): <description>"
{meta_memory_info}
## User Context
{context}
## Three-Step Retrieval Strategy
**STEP 1: Read User Profile (REQUIRED FIRST)**
- Use `read_user_profile` with memory_type and memory_target from available meta memories
- Check if the user profile directly answers the question
- If sufficient information found, provide the answer and STOP
**STEP 2: Vector Search (If Step 1 insufficient)**
- Use `retrieve_memory` with memory_type, memory_target, and query
- Try multiple retrieval angles (at least 3 different attempts):
* Direct query with user's question
* Reformulated queries with different phrasing/keywords
* Queries focused on specific entities or concepts
- **Time Range Filtering** (when applicable):
* Format: [start_date, end_date] in YYYYMMDD format
* Example: [20200101, 20200102] means 20200101 < time < 20200102
* Single-sided: [0, 20200102] for before, [20200101, 99999999] for after
* If no results, try broader time ranges or remove time constraints
- If no results after multiple attempts, try different memory_type/memory_target combinations
**STEP 3: Read Original Conversations (If Step 2 insufficient)**
- Use `read_history` with history_id from retrieved memories
- Prioritize reading:
* Most recent memories with history_id
* Most relevant memories from Step 2 with history_id
- Try multiple history_id entries if needed
## Response Rules
- Answer ONLY based on retrieved information - NEVER guess or fabricate
- If nothing found after all three steps: State clearly "I don't know. I cannot find relevant information to answer this question."
- Be persistent: try multiple angles in each step before moving to the next
- Once you find sufficient information, provide a direct answer
user_message: |
Retrieve relevant memories and answer the question using the three-step strategy.

View file

@ -0,0 +1,88 @@
from loguru import logger
from ..base_memory_agent import BaseMemoryAgent
from ...core.enumeration import Role, MemoryType
from ...core.schema import Message, MemoryNode, ToolCall
from ...core.utils import format_messages
class ReMeSummarizerV3(BaseMemoryAgent):
def __init__(self, meta_memories: list[dict] | None = None, **kwargs):
"""Initialize with meta memories list."""
super().__init__(**kwargs)
self.meta_memories: list[dict] = meta_memories or []
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": self.get_prompt("tool"),
"parameters": {
"type": "object",
"properties": {
"messages": {
"type": "array",
"items": {
"type": "object",
"properties": {
"role": {
"type": "string",
"description": "role",
},
"content": {
"type": "string",
"description": "content",
},
},
"required": ["role", "content"],
},
},
},
"required": ["messages"],
},
},
)
async def _read_meta_memories(self) -> str:
from ...mem_tool import ReadMetaMemory
return ReadMetaMemory().format_memory_metadata(self.meta_memories)
async def build_messages(self) -> list[Message]:
"""Construct initial messages with context and meta-memory information."""
messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages]
self.context["messages_formated"] = self.description + "\n" + format_messages(messages)
self.context["ref_memory_id"] = MemoryNode(
memory_type=MemoryType.HISTORY,
content=self.context["messages_formated"],
).memory_id
meta_memory_info = await self._read_meta_memories()
logger.info(f"meta_memory_info={meta_memory_info}")
system_prompt = self.prompt_format(
prompt_name="system_prompt",
meta_memory_info=meta_memory_info,
context=self.context["messages_formated"],
)
user_message = self.get_prompt("user_message")
messages = [
Message(role=Role.SYSTEM, content=system_prompt),
Message(role=Role.USER, content=user_message),
]
return messages
async def _acting_step(self, assistant_message: Message, step: int, **kwargs) -> list[Message]:
"""Execute tool calls with ref_memory_id and author context."""
return await super()._acting_step(
assistant_message,
step,
messages=self.context.get("messages", []),
description=self.context.get("description"),
ref_memory_id=self.context["ref_memory_id"],
messages_formated=self.context["messages_formated"],
author=self.author,
**kwargs,
)

View file

@ -0,0 +1,25 @@
tool: |
Orchestrate the complete memory summarization for the agent.
system_prompt: |
You are a Memory Agent responsible for performing necessary updates and summaries of the main Agent's memories based on the **context**.
# Context
{context}
## Main Agent's Meta Memory
Each line of meta memory indicates the existence of a specialized Memory Agent dedicated to deep summarization and updating of memories within a specific dimension (memory_type + memory_target).
Format: "- <memory_type>(<memory_target>): <description>"
{meta_memory_info}
## Your Task
Use `summary_and_hands_off` tool to:
1. Create a concise summary in `summary_content` that captures key points, decisions, or important facts from the context.
2. Identify which memory dimensions need updates and specify them in `memory_tasks` (each with `memory_type` and `memory_target`).
- The `memory_type` and `memory_target` must exactly match existing entries in the "Main Agent's Meta Memory" listed above.
- Multiple tasks can be specified to enable parallel processing by specialized agents.
Note: If the context contains no memorable information (e.g., simple greetings), output `<NO_MEMORY_NEEDED>`.
user_message: |
Please perform your task based on the context.

View file

@ -3,8 +3,6 @@
from abc import ABCMeta
from pathlib import Path
from loguru import logger
from ..core.enumeration import MemoryType
from ..core.op import BaseOp
from ..core.schema import ToolCall, MemoryNode

View file

@ -0,0 +1,54 @@
from loguru import logger
from ..base_memory_tool import BaseMemoryTool
from ...core.context import C
from ...core.schema.memory_node import MemoryNode
@C.register_op()
class ReadLocalMemories(BaseMemoryTool):
def __init__(self, **kwargs):
kwargs["enable_multiple"] = False
super().__init__(**kwargs)
def _build_parameters(self) -> dict:
return {
"type": "object",
"properties": {
"memory_type": {
"type": "string",
"description": self.get_prompt("memory_type"),
},
"memory_target": {
"type": "string",
"description": self.get_prompt("memory_target"),
},
},
"required": ["memory_type", "memory_target"],
}
async def execute(self):
memory_type = self.context.get("memory_type", "")
memory_target = self.context.get("memory_target", "")
if not memory_type or not memory_target:
self.output = "memory_type and memory_target are required."
return
cache_key = f"{memory_type}_{memory_target}"
cached_data = self.meta_memory.load(cache_key, auto_clean=False)
if not cached_data:
self.output = f"Local memory not found: {memory_type}_{memory_target}"
logger.info(self.output)
return
memory_nodes = [MemoryNode(**node_data) for node_data in cached_data]
if not memory_nodes:
self.output = f"No valid memory nodes found in {memory_type}_{memory_target}"
return
self.output = memory_nodes
logger.info(f"Read {len(memory_nodes)} nodes from cache key: {cache_key}")

View file

@ -0,0 +1,8 @@
tool: |
Read memory nodes from local memory files.
memory_type: |
The type of local memory to read.
memory_target: |
The target identifier for the local memory.

View file

@ -0,0 +1,15 @@
from .add_memory import AddMemory
from .read_history import ReadHistory
from .read_user_profile import ReadUserProfile
from .retrieve_memory import RetrieveMemory
from .summary_and_hands_off import SummaryAndHandsOff
from .update_user_profile import UpdateUserProfile
__all__ = [
"AddMemory",
"ReadHistory",
"ReadUserProfile",
"RetrieveMemory",
"SummaryAndHandsOff",
"UpdateUserProfile",
]

View file

@ -0,0 +1,67 @@
from loguru import logger
from ..base_memory_tool import BaseMemoryTool
from ...core.schema import MemoryNode
class AddMemory(BaseMemoryTool):
def __init__(self, **kwargs):
kwargs['enable_multiple'] = True
super().__init__(**kwargs)
def _build_tool_description(self) -> str:
return "Add multiple memories to the vector store for future retrieval."
def _build_multiple_parameters(self) -> dict:
return {
"type": "object",
"properties": {
"memories": {
"type": "array",
"description": "A list of memory objects to store.",
"items": {
"type": "object",
"properties": {
"memory_content": {
"type": "string",
"description": "memory content",
},
"conversation_time": {
"type": "object",
"description": "conversation time, e.g. '2020-01-01 00:00:00'",
}
},
"required": ["memory_content", "conversation_time"],
},
},
},
"required": ["memories"],
}
async def execute(self):
memories: list[dict] = self.context.get("memories", [])
if not memories:
self.output = "No memories provided for addition."
return
memory_nodes: list[MemoryNode] = []
for mem in memories:
memory_content = mem.get("memory_content", "")
conversation_time = mem.get("conversation_time", "")
metadata: dict = {"conversation_time": conversation_time}
try:
metadata["time_int"] = int(conversation_time.split(" ")[0].replace("-", ""))
except Exception:
...
memory_nodes.append(self._build_memory_node(memory_content, metadata=metadata))
vector_nodes = [node.to_vector_node() for node in memory_nodes]
vector_ids: list[str] = [node.vector_id for node in vector_nodes]
await self.vector_store.delete(vector_ids=vector_ids)
await self.vector_store.insert(nodes=vector_nodes)
self.memory_nodes = memory_nodes
self.output = f"Successfully added {len(memory_nodes)} memories to vector_store."
logger.info(self.output)

View file

@ -0,0 +1,38 @@
from loguru import logger
from ..base_memory_tool import BaseMemoryTool
from ...core.schema import MemoryNode
class ReadHistory(BaseMemoryTool):
def __init__(self, **kwargs):
kwargs["enable_multiple"] = False
super().__init__(**kwargs)
def _build_tool_description(self) -> str:
return "Read original history dialogue."
def _build_parameters(self) -> dict:
return {
"type": "object",
"properties": {
"history_id": {
"type": "string",
"description": "history_id",
},
},
"required": ["history_id"],
}
async def execute(self):
history_id = self.context.get("history_id", "")
nodes = await self.vector_store.get(vector_ids=[history_id])
if not nodes:
self.output = f"No history: {history_id}"
logger.warning(self.output)
return
memory = MemoryNode.from_vector_node(nodes[0])
self.output = memory.content
logger.info(f"Successfully read history memory: {history_id}")

View file

@ -0,0 +1,65 @@
from loguru import logger
from ..base_memory_tool import BaseMemoryTool
from ...core.schema.memory_node import MemoryNode
class ReadUserProfile(BaseMemoryTool):
def __init__(self, add_memory_type_target: bool = True, **kwargs):
kwargs["enable_multiple"] = False
self.add_memory_type_target = add_memory_type_target
super().__init__(**kwargs)
def _build_tool_description(self) -> str:
return "Read personal memory profile for the current user."
def _build_parameters(self) -> dict:
if self.add_memory_type_target:
return {
"type": "object",
"properties": {
"memory_type": {
"type": "string",
"description": "memory_type",
},
"memory_target": {
"type": "string",
"description": "memory_target",
},
},
"required": ["memory_type", "memory_target"],
}
else:
return {
"type": "object",
"properties": {},
"required": [],
}
async def execute(self):
cache_key = f"{self.memory_type}_{self.memory_target}"
cached_data = self.meta_memory.load(cache_key, auto_clean=False)
if not cached_data:
self.output = f"Local memory not found: {self.memory_type}_{self.memory_target}"
logger.info(self.output)
return
# Convert to MemoryNode objects and sort by conversation_time (oldest first)
memory_nodes = [MemoryNode(**node_data) for node_data in cached_data]
memory_nodes.sort(
key=lambda node: node.metadata.get("conversation_time", "")
)
memory_formated = []
for node in memory_nodes:
node_formated = f"profile_id={node.memory_id} profile_content={node.content}"
if "conversation_time" in node.metadata:
node_formated += f" conversation_time={node.metadata['conversation_time']}"
if node.ref_memory_id:
node_formated += f" history_id={node.ref_memory_id}"
memory_formated.append(node_formated.strip())
self.output = "\n".join(memory_formated)
logger.info(f"Read {len(memory_formated)} nodes from cache key: {cache_key}")

View file

@ -0,0 +1,84 @@
import json
from loguru import logger
from ..base_memory_tool import BaseMemoryTool
from ...core.schema import MemoryNode
from ...core.utils import deduplicate_memories
class RetrieveMemory(BaseMemoryTool):
def __init__(self, top_k: int = 20, **kwargs):
super().__init__(**kwargs)
self.top_k: int = top_k
def _build_tool_description(self) -> str:
return "Retrieve memories using vector similarity search."
def _build_multiple_parameters(self) -> dict:
return {
"type": "object",
"properties": {
"query_items": {
"type": "array",
"description": "query_items",
"items": {
"type": "object",
"properties": {
"memory_type": {
"type": "string",
"description": "memory_type",
},
"memory_target": {
"type": "string",
"description": "memory_target",
},
"query": {
"type": "string",
"description": "query",
},
"time_range": {
"type": "string",
"description": "time_range(optional), e.g. [20200101, 20200101]",
},
},
"required": ["memory_type", "memory_target", "query"],
},
},
},
"required": ["query_items"],
}
async def execute(self):
query_items: list[dict] = self.context.get("query_items", [])
memory_nodes: list[MemoryNode] = []
for query_item in query_items:
memory_type = query_item.get("memory_type")
memory_target = query_item.get("memory_target")
query = query_item.get("query")
time_range = query_item.get("time_range", "")
filter_dict = {
"memory_type": memory_type,
"memory_target": memory_target,
}
if time_range:
time_range = json.loads(time_range)
filter_dict["time_range"] = [int(time_range[0]), int(time_range[1])]
nodes = await self.vector_store.search(query=query, limit=self.top_k, filters=filter_dict)
memory_nodes.extend([MemoryNode.from_vector_node(n) for n in nodes])
memory_nodes = deduplicate_memories(memory_nodes)
retrieved_memory_ids = {node.memory_id for node in self.retrieved_nodes if node.memory_id}
new_memory_nodes = [node for node in memory_nodes if node.memory_id not in retrieved_memory_ids]
self.retrieved_nodes.extend(new_memory_nodes)
self.memory_nodes = new_memory_nodes
if not new_memory_nodes:
self.output = "No new memory_nodes found matching the query (duplicates removed)."
else:
self.output = "\n".join([f"{m.metadata['conversation_time']} {m.content}" for m in new_memory_nodes])
logger.info(f"Retrieved {len(memory_nodes)} memory_nodes, {len(new_memory_nodes)} new after deduplication")

View file

@ -0,0 +1,140 @@
import json
from typing import TYPE_CHECKING
from loguru import logger
from ..base_memory_tool import BaseMemoryTool
from ...core.enumeration import MemoryType
from ...core.schema import MemoryNode, Message
if TYPE_CHECKING:
from ...mem_agent import BaseMemoryAgent
class SummaryAndHandsOff(BaseMemoryTool):
def __init__(self, memory_agents: list["BaseMemoryAgent"], **kwargs):
kwargs["enable_multiple"] = True
kwargs["sub_ops"] = memory_agents or []
super().__init__(**kwargs)
from ...mem_agent import BaseMemoryAgent
self.sub_ops: list[BaseMemoryAgent] = [a for a in self.sub_ops if isinstance(a, BaseMemoryAgent)]
self.messages: list[Message] = []
@property
def memory_agent_dict(self) -> dict[MemoryType, "BaseMemoryAgent"]:
return {a.memory_type: a for a in self.sub_ops}
def _build_tool_description(self) -> str:
return "Summarize and distribute memory tasks to appropriate agents."
def _build_multiple_parameters(self) -> dict:
return {
"type": "object",
"properties": {
"summary_content": {
"type": "string",
"description": "summary content",
},
"memory_tasks": {
"type": "array",
"description": "memory_tasks",
"items": {
"type": "object",
"properties": {
"memory_type": {
"type": "string",
"description": "memory_type",
"enum": [k.value for k in self.memory_agent_dict],
},
"memory_target": {
"type": "string",
"description": "memory_target",
},
},
"required": ["memory_type", "memory_target"],
},
},
},
"required": ["summary_content", "memory_tasks"],
}
@staticmethod
def _parse_memory_type_target(task: dict):
return {
"memory_type": MemoryType(task.get("memory_type", "")),
"memory_target": task.get("memory_target", ""),
}
def _collect_tasks(self) -> list[dict]:
tasks = []
for task in self.context.get("memory_tasks", []):
tasks.append(self._parse_memory_type_target(task))
return tasks
async def execute(self):
summary_content = self.context.get("summary_content", "")
assert summary_content, "No summary content provided."
summary_node = MemoryNode(
memory_type=MemoryType.HISTORY,
memory_target="",
when_to_use=summary_content,
content=self.messages_formated,
ref_memory_id="",
author=self.author,
metadata={},
)
logger.info(f"Adding summary node: {summary_node.model_dump_json(indent=2, exclude_none=True)}")
self.memory_nodes.append(summary_node)
vector_node = summary_node.to_vector_node()
await self.vector_store.delete(vector_ids=[vector_node.vector_id])
await self.vector_store.insert([vector_node])
tasks = self._collect_tasks()
if not tasks:
self.output = "No valid memory tasks to execute."
return
agent_list = []
for i, task in enumerate(tasks):
memory_type: MemoryType = task["memory_type"]
memory_target: str = task["memory_target"]
if memory_type not in self.memory_agent_dict:
logger.warning(f"No agent found for memory_type={memory_type}")
continue
agent = self.memory_agent_dict[memory_type].copy()
agent_list.append([agent, memory_type, memory_target])
logger.info(f"Task {i}: Submitting {memory_type.value} agent for target={memory_target}")
self.submit_async_task(
agent.call,
query=self.context.get("query", ""),
messages=self.context.get("messages", []),
memory_type=memory_type,
memory_target=memory_target,
description=self.context.get("description"),
ref_memory_id=self.context.get("ref_memory_id", ""),
)
await self.join_async_tasks()
results = []
for i, (agent, memory_type, memory_target) in enumerate(agent_list):
result_str = str(agent.output)
if agent.memory_nodes:
self.memory_nodes.extend(agent.memory_nodes)
if agent.messages:
self.messages.extend(agent.messages)
results.append({
"memory_type": memory_type.value,
"memory_target": memory_target,
"result": result_str[:100] + ("..." if len(result_str) > 100 else ""),
})
logger.info(f"Task {i}: Completed {memory_type.value} agent for target={memory_target}")
results_str = json.dumps(results, ensure_ascii=False, indent=2)
self.output = f"Successfully executed summary and {len(results)} hands-off task(s):\n{results_str}"

View file

@ -0,0 +1,118 @@
from loguru import logger
from ..base_memory_tool import BaseMemoryTool
from ...core.context import C
from ...core.schema.memory_node import MemoryNode
@C.register_op()
class UpdateUserProfile(BaseMemoryTool):
def __init__(self, **kwargs):
kwargs["enable_multiple"] = True
super().__init__(**kwargs)
def _build_multiple_parameters(self) -> dict:
return {
"type": "object",
"properties": {
"profile_ids_to_delete": {
"type": "array",
"description": self.get_prompt("profile_ids_to_delete"),
"items": {"type": "string"},
},
"profiles_to_add": {
"type": "array",
"description": self.get_prompt("profiles_to_add"),
"items": {
"type": "object",
"properties": {
"profile_content": {
"type": "string",
"description": self.get_prompt("profile_content"),
},
"timestamp": {
"type": "string",
"description": self.get_prompt("timestamp"),
},
},
"required": ["profile_content", "timestamp"],
},
},
},
"required": ["profile_ids_to_delete", "profiles_to_add"],
}
async def execute(self):
memory_type = "personal"
memory_target = self.memory_target
assert memory_target, "memory_target is not configured."
cache_key = f"{memory_type}_{memory_target}"
profile_ids_to_delete = self.context.get("profile_ids_to_delete", [])
profile_ids_to_delete = [m for m in profile_ids_to_delete if m]
profile_ids_to_delete = list(dict.fromkeys(profile_ids_to_delete))
profiles_to_add = self.context.get("profiles_to_add", [])
if not profile_ids_to_delete and not profiles_to_add:
self.output = "No memories to remove or add. Operation has been done."
return
cached_data = self.meta_memory.load(cache_key, auto_clean=False)
existing_memory_nodes = []
if cached_data:
existing_memory_nodes = [MemoryNode(**node_data) for node_data in cached_data]
removed_count = 0
added_count = 0
if profile_ids_to_delete:
profile_ids_set = set(profile_ids_to_delete)
existing_memory_nodes = [
node for node in existing_memory_nodes if node.memory_id not in profile_ids_set
]
removed_count = len(profile_ids_to_delete)
logger.info(f"Removed {removed_count} memories from user profile.")
new_memory_nodes = []
if profiles_to_add:
for mem in profiles_to_add:
profile_content = mem.get("profile_content", "")
timestamp = mem.get("timestamp", "")
if not profile_content:
logger.warning("Skipping memory with empty content")
continue
memory_node = self._build_memory_node(
memory_content=profile_content,
when_to_use="",
metadata={"timestamp": timestamp}
)
memory_node.memory_type = MemoryNode.MemoryType.PERSONAL
memory_node.memory_target = memory_target
new_memory_nodes.append(memory_node)
added_count = len(new_memory_nodes)
logger.info(f"Added {added_count} new memories to user profile.")
updated_memory_nodes = existing_memory_nodes + new_memory_nodes
nodes_data = [node.model_dump(exclude_none=True) for node in updated_memory_nodes]
self.meta_memory.save(cache_key, nodes_data)
operations = []
if removed_count > 0:
operations.append(f"removed {removed_count} old memories")
if added_count > 0:
operations.append(f"added {added_count} new memories")
if operations:
self.output = f"Successfully {' and '.join(operations)} in user profile."
else:
self.output = "Operation has been done."
logger.info(self.output)

View file

@ -0,0 +1,57 @@
from loguru import logger
from ..base_memory_tool import BaseMemoryTool
from ...core.context import C
from ...core.schema.memory_node import MemoryNode
@C.register_op()
class WriteLocalMemories(BaseMemoryTool):
def __init__(self, **kwargs):
kwargs["enable_multiple"] = True
super().__init__(**kwargs)
def _build_multiple_parameters(self) -> dict:
return {
"type": "object",
"properties": {
"memory_nodes": {
"type": "array",
"description": self.get_prompt("memory_nodes"),
"items": {
"type": "object",
"description": "Memory node object",
},
},
},
"required": ["memory_nodes"],
}
async def execute(self):
memory_nodes = self.context.get("memory_nodes", [])
if not memory_nodes:
self.output = "No memory nodes provided."
return
memory_nodes = [MemoryNode(**node) if isinstance(node, dict) else node for node in memory_nodes]
grouped = {}
for node in memory_nodes:
key = (node.memory_type.value, node.memory_target)
if key not in grouped:
grouped[key] = []
grouped[key].append(node)
written_keys = []
for (memory_type, memory_target), nodes in grouped.items():
cache_key = f"{memory_type}_{memory_target}"
nodes_data = [node.model_dump() for node in nodes]
self.meta_memory.save(cache_key, nodes_data)
written_keys.append(f"{memory_type}_{memory_target}")
logger.info(f"Saved {len(nodes)} nodes to cache key: {cache_key}")
self.output = f"Successfully written local memories: {', '.join(written_keys)}"

View file

@ -0,0 +1,5 @@
tool_multiple: |
Write memory nodes to local memory files.
memory_nodes: |
List of memory nodes to write to local files.

View file

@ -13,6 +13,11 @@ from .mem_agent.retriever import ReMeRetriever
from .mem_agent.retriever_v2 import ReMeRetrieverV2
from .mem_agent.summarizer import ReMeSummarizer, PersonalSummarizer
from .mem_agent.summarizer_v2 import ReMeSummarizerV2, PersonalSummarizerV2
from .mem_agent.v3 import (
PersonalSummarizerV3,
ReMeRetrieverV3,
ReMeSummarizerV3,
)
from .mem_tool import (
HandsOffTool,
ReadHistoryMemory,
@ -24,12 +29,19 @@ from .mem_tool import (
)
from .mem_tool.v2 import (
AddMemoryDrafts,
ReadHistory,
RetrieveMemories,
RetrieveRecentAndSimilarMemories,
SummaryAndHandsOff,
UpdateMemories,
)
from .mem_tool.v3 import (
AddMemory as AddMemoryV3,
ReadHistory as ReadHistoryV3,
ReadUserProfile,
RetrieveMemory,
SummaryAndHandsOff as SummaryAndHandsOffV3,
UpdateUserProfile,
)
@singleton
@ -314,3 +326,86 @@ class ReMe(Application):
else:
raise NotImplementedError
async def summary_v3(
self,
messages: list[dict],
description: str = "",
user_id: str = "",
assistant_id: str = "",
**kwargs,
):
"""Summarizes messages using V3 workflow with user profile management."""
if user_id:
meta_memories = [
{
"memory_type": "personal",
"memory_target": user_id,
},
]
messages = self._prepare_messages(messages, user_id, assistant_id)
personal_summarizer_v3 = PersonalSummarizerV3(
tools=[
AddMemoryV3(),
ReadUserProfile(add_memory_type_target=False),
UpdateUserProfile(),
],
)
reme_summarizer_v3 = ReMeSummarizerV3(
meta_memories=meta_memories,
tools=[SummaryAndHandsOffV3(memory_agents=[personal_summarizer_v3])],
)
# try:
await reme_summarizer_v3.call(messages=messages, description=description, **kwargs)
return reme_summarizer_v3.memory_nodes, reme_summarizer_v3.messages, reme_summarizer_v3.success
# except Exception as e:
# print(f"Warning: reme_summarizer_v3.call failed: {e}")
# return [], [], False
else:
raise NotImplementedError
async def retrieve_v3(
self,
query: str = "",
messages: list[dict] | None = None,
description: str = "",
user_id: str = "",
assistant_id: str = "",
top_k: int = 20,
**kwargs,
):
"""Retrieves relevant memories using V3 workflow with user profile support."""
if user_id:
messages = self._prepare_messages(messages, user_id, assistant_id)
meta_memories = [
{
"memory_type": "personal",
"memory_target": user_id,
},
]
reme_retriever_v3 = ReMeRetrieverV3(
meta_memories=meta_memories,
tools=[
ReadUserProfile(add_memory_type_target=True),
RetrieveMemory(top_k=top_k),
ReadHistoryV3(),
],
)
# try:
await reme_retriever_v3.call(query=query, messages=messages, description=description, **kwargs)
return reme_retriever_v3.output, reme_retriever_v3.messages, reme_retriever_v3.success
# except Exception as e:
# print(f"Warning: reme_retriever_v3.call failed: {e}")
# return "error, not retrieved", [], False
else:
raise NotImplementedError

View file

@ -350,35 +350,36 @@ async def test_search_with_single_filter(store: BaseVectorStore, _store_name: st
logger.info("✓ Single filter search test passed")
async def test_search_with_list_filter(store: BaseVectorStore, _store_name: str):
"""Test vector search with list filter (IN operation)."""
logger.info("=" * 20 + " LIST FILTER SEARCH TEST " + "=" * 20)
async def test_search_with_exact_match_filter(store: BaseVectorStore, _store_name: str):
"""Test vector search with exact match filter."""
logger.info("=" * 20 + " EXACT MATCH FILTER SEARCH TEST " + "=" * 20)
# Test list filter (IN operation)
filters = {"node_type": ["tech", "tech_new"]}
# Test exact match filter
filters = {"node_type": "tech"}
results = await store.search(
query="What is artificial intelligence?",
limit=5,
filters=filters,
)
logger.info(f"Filtered search (node_type IN [tech, tech_new]) returned {len(results)} results")
logger.info(f"Filtered search (node_type=tech) returned {len(results)} results")
for i, r in enumerate(results, 1):
node_type = r.metadata.get("node_type")
logger.info(f" Result {i}: type={node_type}, content={r.content[:50]}...")
assert node_type in ["tech", "tech_new"], "Result should have node_type in [tech, tech_new]"
assert node_type == "tech", "Result should have node_type='tech'"
logger.info("✓ List filter search test passed")
logger.info("✓ Exact match filter search test passed")
async def test_search_with_multiple_filters(store: BaseVectorStore, _store_name: str):
"""Test vector search with multiple metadata filters (AND operation)."""
logger.info("=" * 20 + " MULTIPLE FILTERS SEARCH TEST " + "=" * 20)
# Test multiple filters (AND operation)
# Test multiple exact match filters (AND operation)
filters = {
"node_type": ["tech", "tech_new"],
"node_type": "tech",
"source": "research",
"priority": "high",
}
results = await store.search(
query="What is artificial intelligence?",
@ -387,14 +388,16 @@ async def test_search_with_multiple_filters(store: BaseVectorStore, _store_name:
)
logger.info(
f"Multi-filter search (node_type IN [tech, tech_new] AND source=research) " f"returned {len(results)} results",
f"Multi-filter search (node_type=tech AND source=research AND priority=high) " f"returned {len(results)} results",
)
for i, r in enumerate(results, 1):
node_type = r.metadata.get("node_type")
source = r.metadata.get("source")
logger.info(f" Result {i}: type={node_type}, source={source}, content={r.content[:40]}...")
assert node_type in ["tech", "tech_new"], "Result should have node_type in [tech, tech_new]"
priority = r.metadata.get("priority")
logger.info(f" Result {i}: type={node_type}, source={source}, priority={priority}")
assert node_type == "tech", "Result should have node_type='tech'"
assert source == "research", "Result should have source='research'"
assert priority == "high", "Result should have priority='high'"
logger.info("✓ Multiple filters search test passed")
@ -789,10 +792,9 @@ async def test_complex_metadata_queries(store: BaseVectorStore, _store_name: str
await store.insert(complex_nodes)
logger.info(f"✓ Inserted {len(complex_nodes)} nodes with complex metadata")
# Test 1: Multiple field filters with list values
# Test 1: Multiple exact match filters
filters_1 = {
"domain": "AI",
"year": ["2023", "2024"],
"impact_factor": "high",
}
results_1 = await store.search(
@ -800,26 +802,25 @@ async def test_complex_metadata_queries(store: BaseVectorStore, _store_name: str
limit=10,
filters=filters_1,
)
logger.info(f"Test 1 - AI + high impact + recent years: {len(results_1)} results")
logger.info(f"Test 1 - AI + high impact: {len(results_1)} results")
for r in results_1:
assert r.metadata.get("domain") == "AI"
assert r.metadata.get("impact_factor") == "high"
assert r.metadata.get("year") in ["2023", "2024"]
# Test 2: List filter with multiple subdomains
# Test 2: Single exact match filter
filters_2 = {
"subdomain": ["nlp", "computer_vision"],
"subdomain": "nlp",
}
results_2 = await store.search(
query="deep learning applications",
limit=10,
filters=filters_2,
)
logger.info(f"Test 2 - NLP or Computer Vision: {len(results_2)} results")
logger.info(f"Test 2 - NLP subdomain: {len(results_2)} results")
for r in results_2:
assert r.metadata.get("subdomain") in ["nlp", "computer_vision"]
assert r.metadata.get("subdomain") == "nlp"
# Test 3: Year-based filtering
# Test 3: Year-based exact match filtering
filters_3 = {
"year": "2024",
}
@ -1119,65 +1120,47 @@ async def test_filter_combinations(store: BaseVectorStore, _store_name: str):
results_1 = await store.search(query="technology", filters={}, limit=10)
logger.info(f"Test 1 - Empty filter: {len(results_1)} results")
# Test 2: Single value filter
# Test 2: Single exact match filter
results_2 = await store.search(
query="technology",
filters={"node_type": "tech"},
limit=10,
)
logger.info(f"Test 2 - Single value filter: {len(results_2)} results")
logger.info(f"Test 2 - Single exact match filter: {len(results_2)} results")
for r in results_2:
assert r.metadata.get("node_type") == "tech"
# Test 3: List filter with single item
# Test 3: Multiple exact match filters (AND operation)
results_3 = await store.search(
query="technology",
filters={"node_type": ["tech"]},
limit=10,
)
logger.info(f"Test 3 - List filter (single item): {len(results_3)} results")
# Test 4: List filter with multiple items
results_4 = await store.search(
query="technology",
filters={"category": ["AI", "ML", "DL"]},
limit=10,
)
logger.info(f"Test 4 - List filter (multiple items): {len(results_4)} results")
for r in results_4:
assert r.metadata.get("category") in ["AI", "ML", "DL"]
# Test 5: Multiple filters (AND operation)
results_5 = await store.search(
query="technology",
filters={
"node_type": ["tech", "tech_new"],
"node_type": "tech",
"source": "research",
"priority": "high",
},
limit=10,
)
logger.info(f"Test 5 - Multiple filters (AND): {len(results_5)} results")
for r in results_5:
assert r.metadata.get("node_type") in ["tech", "tech_new"]
logger.info(f"Test 3 - Multiple exact match filters (AND): {len(results_3)} results")
for r in results_3:
assert r.metadata.get("node_type") == "tech"
assert r.metadata.get("source") == "research"
assert r.metadata.get("priority") == "high"
# Test 6: Filter with non-existent value
results_6 = await store.search(
# Test 4: Filter with non-existent value
results_4 = await store.search(
query="technology",
filters={"category": "NON_EXISTENT_CATEGORY"},
limit=10,
)
logger.info(f"Test 6 - Non-existent filter value: {len(results_6)} results")
assert len(results_6) == 0, "Should return no results for non-existent filter value"
logger.info(f"Test 4 - Non-existent filter value: {len(results_4)} results")
assert len(results_4) == 0, "Should return no results for non-existent filter value"
# Test 7: List operation with filters
# Test 5: List operation with multiple exact match filters
list_results = await store.list(
filters={"node_type": "tech", "priority": "high"},
limit=20,
)
logger.info(f"Test 7 - List with filters: {len(list_results)} results")
logger.info(f"Test 5 - List with multiple filters: {len(list_results)} results")
for r in list_results:
assert r.metadata.get("node_type") == "tech"
assert r.metadata.get("priority") == "high"
@ -1185,6 +1168,329 @@ async def test_filter_combinations(store: BaseVectorStore, _store_name: str):
logger.info("✓ Filter combinations test passed")
async def test_range_query_filters(store: BaseVectorStore, _store_name: str):
"""Test range query filters using the new [start, end] syntax."""
logger.info("=" * 20 + " RANGE QUERY FILTERS TEST " + "=" * 20)
# Clean up any existing test data first
try:
existing_nodes = await store.list(filters={"test_type": "range_query_test"})
if existing_nodes:
await store.delete([node.vector_id for node in existing_nodes])
logger.info(f"Cleaned up {len(existing_nodes)} existing test nodes")
except Exception as e:
logger.warning(f"Failed to clean up existing nodes: {e}")
# Create test nodes with numeric metadata for range queries
import time
base_timestamp = int(time.time())
test_nodes = []
for i in range(20):
node = VectorNode(
vector_id=f"range_node_{i}",
content=f"Test content for range query node {i}",
metadata={
"test_type": "range_query_test",
"timestamp": base_timestamp + i * 1000, # Each node is 1000 seconds apart
"rating": 50 + i * 2, # Ratings from 50 to 88
"priority": i % 3, # 0, 1, or 2
"category": ["tech", "science", "business"][i % 3],
},
)
test_nodes.append(node)
# Insert test nodes
await store.insert(test_nodes)
logger.info(f"Inserted {len(test_nodes)} test nodes with numeric metadata")
# Test 1: Range query on timestamp field
start_time = base_timestamp + 5000
end_time = base_timestamp + 15000
results_1 = await store.search(
query="test content",
limit=20,
filters={
"timestamp": [start_time, end_time], # Range query: >= start_time AND <= end_time
},
)
logger.info(f"Test 1 - Timestamp range [{start_time}, {end_time}]: {len(results_1)} results")
# Verify all results are within range
for r in results_1:
ts = r.metadata.get("timestamp")
assert ts >= start_time, f"Timestamp {ts} should be >= {start_time}"
assert ts <= end_time, f"Timestamp {ts} should be <= {end_time}"
logger.debug(f" Node {r.vector_id}: timestamp={ts}")
# Expected nodes: range_node_5 to range_node_15 (11 nodes)
assert len(results_1) >= 10, f"Expected at least 10 results, got {len(results_1)}"
logger.info("✓ Timestamp range query validated")
# Test 2: Range query on rating field
results_2 = await store.search(
query="test content",
limit=20,
filters={
"rating": [60, 80], # Range query: rating >= 60 AND rating <= 80
},
)
logger.info(f"Test 2 - Rating range [60, 80]: {len(results_2)} results")
# Verify all results are within rating range
for r in results_2:
rating = r.metadata.get("rating")
assert rating >= 60, f"Rating {rating} should be >= 60"
assert rating <= 80, f"Rating {rating} should be <= 80"
logger.debug(f" Node {r.vector_id}: rating={rating}")
# Expected: ratings from 60 to 80 (nodes 5-15)
assert len(results_2) >= 10, f"Expected at least 10 results, got {len(results_2)}"
logger.info("✓ Rating range query validated")
# Test 3: Combine range query with exact match filter
results_3 = await store.search(
query="test content",
limit=20,
filters={
"timestamp": [start_time, end_time],
"category": "tech", # Exact match
},
)
logger.info(
f"Test 3 - Timestamp range + exact match (category=tech): {len(results_3)} results",
)
# Verify filters
for r in results_3:
ts = r.metadata.get("timestamp")
category = r.metadata.get("category")
assert ts >= start_time and ts <= end_time, "Timestamp should be in range"
assert category == "tech", f"Category should be 'tech', got '{category}'"
logger.debug(f" Node {r.vector_id}: timestamp={ts}, category={category}")
# Expected: nodes within range AND category=tech
assert len(results_3) >= 3, f"Expected at least 3 results, got {len(results_3)}"
logger.info("✓ Combined range + exact match query validated")
# Test 4: Multiple range queries
results_4 = await store.search(
query="test content",
limit=20,
filters={
"timestamp": [base_timestamp + 8000, base_timestamp + 12000],
"rating": [65, 75],
},
)
logger.info(f"Test 4 - Multiple range queries: {len(results_4)} results")
# Verify both ranges
for r in results_4:
ts = r.metadata.get("timestamp")
rating = r.metadata.get("rating")
assert ts >= base_timestamp + 8000 and ts <= base_timestamp + 12000, "Timestamp out of range"
assert rating >= 65 and rating <= 75, f"Rating {rating} out of range [65, 75]"
logger.debug(f" Node {r.vector_id}: timestamp={ts}, rating={rating}")
# Expected: nodes 8-12 (5 nodes) with overlapping ranges
assert len(results_4) >= 3, f"Expected at least 3 results, got {len(results_4)}"
logger.info("✓ Multiple range queries validated")
# Test 5: Range query with list operation
results_5 = await store.list(
filters={
"rating": [60, 70],
"test_type": "range_query_test",
},
limit=20,
)
logger.info(f"Test 5 - Range query in list operation: {len(results_5)} results")
# Verify rating range in list results
for r in results_5:
rating = r.metadata.get("rating")
assert rating >= 60 and rating <= 70, f"Rating {rating} should be in range [60, 70]"
logger.info("✓ Range query in list operation validated")
# Test 6: Edge case - exact boundary values
results_6 = await store.list(
filters={
"rating": [60, 60], # Exact match using range syntax
"test_type": "range_query_test",
},
limit=20,
)
logger.info(f"Test 6 - Exact value using range syntax [60, 60]: {len(results_6)} results")
# Should return exactly one node (range_node_5 with rating=60)
for r in results_6:
rating = r.metadata.get("rating")
assert rating == 60, f"Rating should be exactly 60, got {rating}"
logger.info("✓ Boundary value range query validated")
# Test 7: Range query with sorting
results_7 = await store.list(
filters={
"rating": [60, 80],
"test_type": "range_query_test",
},
sort_key="rating",
reverse=True,
limit=5,
)
logger.info(f"Test 7 - Range query with sorting: {len(results_7)} results")
# Verify results are sorted and within range
for i in range(len(results_7) - 1):
rating1 = results_7[i].metadata.get("rating")
rating2 = results_7[i + 1].metadata.get("rating")
assert rating1 >= rating2, f"Results not sorted: {rating1} < {rating2}"
assert rating1 >= 60 and rating1 <= 80, "Rating out of range"
logger.info("✓ Range query with sorting validated")
# Clean up test data
await store.delete([node.vector_id for node in test_nodes])
logger.info("Cleaned up test nodes")
logger.info("✓ Range query filters test passed")
async def test_string_range_queries(store: BaseVectorStore, store_name: str):
"""Test range queries with string values (e.g., date strings, timestamps)."""
logger.info("=" * 20 + " STRING RANGE QUERIES TEST " + "=" * 20)
# Skip this test for stores that don't support string range queries properly
# Qdrant and ChromaDB only support numeric range queries, not string range queries
if store_name not in ["PGVectorStore", "LocalVectorStore", "ESVectorStore"]:
logger.info(f"Skipping string range query test for {store_name}")
return
# Clean up any existing test data first
try:
existing_nodes = await store.list(filters={"test_type": "string_range_test"})
if existing_nodes:
await store.delete([node.vector_id for node in existing_nodes])
logger.info(f"Cleaned up {len(existing_nodes)} existing test nodes")
except Exception as e:
logger.warning(f"Failed to clean up existing nodes: {e}")
# Create test nodes with string date metadata
test_nodes = []
dates = [
"2024-01-01",
"2024-01-15",
"2024-02-01",
"2024-02-15",
"2024-03-01",
"2024-03-15",
"2024-04-01",
]
for i, date in enumerate(dates):
node = VectorNode(
vector_id=f"string_range_node_{i}",
content=f"Test content for date {date}",
metadata={
"test_type": "string_range_test",
"date": date,
"index": i,
},
)
test_nodes.append(node)
# Insert test nodes
await store.insert(test_nodes)
logger.info(f"Inserted {len(test_nodes)} test nodes with string dates")
# Test 1: String range query on date field
try:
results = await store.search(
query="test content",
limit=20,
filters={
"date": ["2024-02-01", "2024-03-15"], # Range query on string dates
},
)
logger.info(f"Test 1 - String date range ['2024-02-01', '2024-03-15']: {len(results)} results")
# Verify all results are within range
expected_dates = ["2024-02-01", "2024-02-15", "2024-03-01", "2024-03-15"]
for r in results:
date = r.metadata.get("date")
assert date >= "2024-02-01", f"Date {date} should be >= '2024-02-01'"
assert date <= "2024-03-15", f"Date {date} should be <= '2024-03-15'"
logger.debug(f" Node {r.vector_id}: date={date}")
assert len(results) >= 3, f"Expected at least 3 results, got {len(results)}"
logger.info("✓ String range query validated")
except Exception as e:
# For PGVector, this might fail on older implementations
if "PGVector" in store_name:
logger.warning(f"String range query failed for PGVector (expected if not updated): {e}")
else:
raise
# Clean up test data
await store.delete([node.vector_id for node in test_nodes])
logger.info("Cleaned up test nodes")
logger.info("✓ String range queries test passed")
async def test_sql_injection_protection(store: BaseVectorStore, store_name: str):
"""Test SQL injection protection in filter keys and collection names."""
logger.info("=" * 20 + " SQL INJECTION PROTECTION TEST " + "=" * 20)
# This test is only relevant for SQL-based stores
if store_name not in ["PGVectorStore"]:
logger.info(f"Skipping SQL injection test for {store_name}")
return
# Test 1: Invalid collection name (SQL injection attempt)
try:
from reme_ai.core.vector_store import PGVectorStore
from reme_ai.core.embedding import OpenAIEmbeddingModel
embedding_model = OpenAIEmbeddingModel()
# This should raise ValueError due to invalid table name
try:
invalid_store = PGVectorStore(
collection_name="test'; DROP TABLE users; --",
embedding_model=embedding_model,
)
logger.error("❌ FAILED: Invalid collection name was accepted (SQL injection risk!)")
assert False, "Should have raised ValueError for invalid collection name"
except ValueError as e:
logger.info(f"✓ Invalid collection name rejected: {e}")
# Test 2: Invalid metadata key in filters
try:
results = await store.search(
query="test",
filters={
"normal_key": "value",
"bad'; DROP TABLE users; --": "value",
},
)
logger.error("❌ FAILED: Invalid metadata key was accepted (SQL injection risk!)")
assert False, "Should have raised ValueError for invalid metadata key"
except ValueError as e:
logger.info(f"✓ Invalid metadata key rejected: {e}")
logger.info("✓ SQL injection protection validated")
except Exception as e:
logger.error(f"SQL injection protection test failed: {e}")
raise
logger.info("✓ SQL injection protection test passed")
async def test_list_with_sorting(store: BaseVectorStore, _store_name: str):
"""Test list operation with sorting by timestamp to get most recent top 10 items."""
logger.info("=" * 20 + " LIST WITH SORTING TEST " + "=" * 20)
@ -1353,7 +1659,7 @@ async def run_all_tests_for_store(store_type: str, store_name: str):
await test_insert(store, store_name)
await test_search(store, store_name)
await test_search_with_single_filter(store, store_name)
await test_search_with_list_filter(store, store_name)
await test_search_with_exact_match_filter(store, store_name)
await test_search_with_multiple_filters(store, store_name)
await test_get_by_id(store, store_name)
await test_list_all(store, store_name)
@ -1374,6 +1680,9 @@ async def run_all_tests_for_store(store_type: str, store_name: str):
await test_metadata_statistics(store, store_name)
await test_update_metadata_only(store, store_name)
await test_filter_combinations(store, store_name)
await test_range_query_filters(store, store_name)
await test_string_range_queries(store, store_name)
await test_sql_injection_protection(store, store_name)
await test_list_with_sorting(store, store_name)
# ========== Collection Management Tests ==========