mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-09 03:20:54 +00:00
feat(core): add ReMe V3 implementation with optional MCP client and enhanced filtering
This commit is contained in:
parent
4d312ea682
commit
e6ad682ede
32 changed files with 2369 additions and 106 deletions
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
668
bench/halumem/eval_reme_simple_v3.py
Normal file
668
bench/halumem/eval_reme_simple_v3.py
Normal 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
|
||||
)
|
||||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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)),
|
||||
)
|
||||
|
|
|
|||
9
reme_ai/mem_agent/v3/__init__.py
Normal file
9
reme_ai/mem_agent/v3/__init__.py
Normal 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",
|
||||
]
|
||||
69
reme_ai/mem_agent/v3/personal_summarizer_v3.py
Normal file
69
reme_ai/mem_agent/v3/personal_summarizer_v3.py
Normal 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
|
||||
38
reme_ai/mem_agent/v3/personal_summarizer_v3.yaml
Normal file
38
reme_ai/mem_agent/v3/personal_summarizer_v3.yaml
Normal 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
|
||||
44
reme_ai/mem_agent/v3/reme_retriever_v3.py
Normal file
44
reme_ai/mem_agent/v3/reme_retriever_v3.py
Normal 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
|
||||
53
reme_ai/mem_agent/v3/reme_retriever_v3.yaml
Normal file
53
reme_ai/mem_agent/v3/reme_retriever_v3.yaml
Normal 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.
|
||||
88
reme_ai/mem_agent/v3/reme_summarizer_v3.py
Normal file
88
reme_ai/mem_agent/v3/reme_summarizer_v3.py
Normal 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,
|
||||
)
|
||||
25
reme_ai/mem_agent/v3/reme_summarizer_v3.yaml
Normal file
25
reme_ai/mem_agent/v3/reme_summarizer_v3.yaml
Normal 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.
|
||||
|
|
@ -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
|
||||
|
|
|
|||
54
reme_ai/mem_tool/read_local_memories.py
Normal file
54
reme_ai/mem_tool/read_local_memories.py
Normal 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}")
|
||||
8
reme_ai/mem_tool/read_local_memories.yaml
Normal file
8
reme_ai/mem_tool/read_local_memories.yaml
Normal 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.
|
||||
15
reme_ai/mem_tool/v3/__init__.py
Normal file
15
reme_ai/mem_tool/v3/__init__.py
Normal 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",
|
||||
]
|
||||
67
reme_ai/mem_tool/v3/add_memory.py
Normal file
67
reme_ai/mem_tool/v3/add_memory.py
Normal 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)
|
||||
38
reme_ai/mem_tool/v3/read_history.py
Normal file
38
reme_ai/mem_tool/v3/read_history.py
Normal 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}")
|
||||
65
reme_ai/mem_tool/v3/read_user_profile.py
Normal file
65
reme_ai/mem_tool/v3/read_user_profile.py
Normal 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}")
|
||||
84
reme_ai/mem_tool/v3/retrieve_memory.py
Normal file
84
reme_ai/mem_tool/v3/retrieve_memory.py
Normal 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")
|
||||
140
reme_ai/mem_tool/v3/summary_and_hands_off.py
Normal file
140
reme_ai/mem_tool/v3/summary_and_hands_off.py
Normal 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}"
|
||||
118
reme_ai/mem_tool/v3/update_user_profile.py
Normal file
118
reme_ai/mem_tool/v3/update_user_profile.py
Normal 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)
|
||||
57
reme_ai/mem_tool/write_local_memories.py
Normal file
57
reme_ai/mem_tool/write_local_memories.py
Normal 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)}"
|
||||
5
reme_ai/mem_tool/write_local_memories.yaml
Normal file
5
reme_ai/mem_tool/write_local_memories.yaml
Normal 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.
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 ==========
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue