diff --git a/bench/halumem/analyze_dataset_stats.py b/bench/halumem/analyze_dataset_stats.py index 6695ad7f..d5a213b4 100644 --- a/bench/halumem/analyze_dataset_stats.py +++ b/bench/halumem/analyze_dataset_stats.py @@ -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() diff --git a/bench/halumem/eval_reme_simple_v3.py b/bench/halumem/eval_reme_simple_v3.py new file mode 100644 index 00000000..62b1fe64 --- /dev/null +++ b/bench/halumem/eval_reme_simple_v3.py @@ -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 + ) diff --git a/reme_ai/core/context/prompt_handler.py b/reme_ai/core/context/prompt_handler.py index e48f98ac..b428b163 100644 --- a/reme_ai/core/context/prompt_handler.py +++ b/reme_ai/core/context/prompt_handler.py @@ -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.""" diff --git a/reme_ai/core/schema/tool_call.py b/reme_ai/core/schema/tool_call.py index 0a96f8c4..035c63f1 100644 --- a/reme_ai/core/schema/tool_call.py +++ b/reme_ai/core/schema/tool_call.py @@ -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 diff --git a/reme_ai/core/utils/__init__.py b/reme_ai/core/utils/__init__.py index 242f7be9..23396f97 100644 --- a/reme_ai/core/utils/__init__.py +++ b/reme_ai/core/utils/__init__.py @@ -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 diff --git a/reme_ai/core/utils/cache_handler.py b/reme_ai/core/utils/cache_handler.py index 70c8585a..f3b0072f 100644 --- a/reme_ai/core/utils/cache_handler.py +++ b/reme_ai/core/utils/cache_handler.py @@ -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}") diff --git a/reme_ai/core/vector_store/chroma_vector_store.py b/reme_ai/core/vector_store/chroma_vector_store.py index 7baa5eab..3a483c23 100644 --- a/reme_ai/core/vector_store/chroma_vector_store.py +++ b/reme_ai/core/vector_store/chroma_vector_store.py @@ -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 diff --git a/reme_ai/core/vector_store/es_vector_store.py b/reme_ai/core/vector_store/es_vector_store.py index 16226749..39989277 100644 --- a/reme_ai/core/vector_store/es_vector_store.py +++ b/reme_ai/core/vector_store/es_vector_store.py @@ -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}} diff --git a/reme_ai/core/vector_store/local_vector_store.py b/reme_ai/core/vector_store/local_vector_store.py index cce3cae2..cfeec74e 100644 --- a/reme_ai/core/vector_store/local_vector_store.py +++ b/reme_ai/core/vector_store/local_vector_store.py @@ -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 diff --git a/reme_ai/core/vector_store/pgvector_store.py b/reme_ai/core/vector_store/pgvector_store.py index a23c84b9..576d8e6c 100644 --- a/reme_ai/core/vector_store/pgvector_store.py +++ b/reme_ai/core/vector_store/pgvector_store.py @@ -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""" diff --git a/reme_ai/core/vector_store/qdrant_vector_store.py b/reme_ai/core/vector_store/qdrant_vector_store.py index 1ac4db64..5d9fa4f5 100644 --- a/reme_ai/core/vector_store/qdrant_vector_store.py +++ b/reme_ai/core/vector_store/qdrant_vector_store.py @@ -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)), ) diff --git a/reme_ai/mem_agent/v3/__init__.py b/reme_ai/mem_agent/v3/__init__.py new file mode 100644 index 00000000..0fd9c86d --- /dev/null +++ b/reme_ai/mem_agent/v3/__init__.py @@ -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", +] diff --git a/reme_ai/mem_agent/v3/personal_summarizer_v3.py b/reme_ai/mem_agent/v3/personal_summarizer_v3.py new file mode 100644 index 00000000..0093884d --- /dev/null +++ b/reme_ai/mem_agent/v3/personal_summarizer_v3.py @@ -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 diff --git a/reme_ai/mem_agent/v3/personal_summarizer_v3.yaml b/reme_ai/mem_agent/v3/personal_summarizer_v3.yaml new file mode 100644 index 00000000..97d18afc --- /dev/null +++ b/reme_ai/mem_agent/v3/personal_summarizer_v3.yaml @@ -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 [] : ` (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 diff --git a/reme_ai/mem_agent/v3/reme_retriever_v3.py b/reme_ai/mem_agent/v3/reme_retriever_v3.py new file mode 100644 index 00000000..8f5c62dc --- /dev/null +++ b/reme_ai/mem_agent/v3/reme_retriever_v3.py @@ -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 diff --git a/reme_ai/mem_agent/v3/reme_retriever_v3.yaml b/reme_ai/mem_agent/v3/reme_retriever_v3.yaml new file mode 100644 index 00000000..8ff005ee --- /dev/null +++ b/reme_ai/mem_agent/v3/reme_retriever_v3.yaml @@ -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: "- (): " + {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. diff --git a/reme_ai/mem_agent/v3/reme_summarizer_v3.py b/reme_ai/mem_agent/v3/reme_summarizer_v3.py new file mode 100644 index 00000000..a0f466b9 --- /dev/null +++ b/reme_ai/mem_agent/v3/reme_summarizer_v3.py @@ -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, + ) diff --git a/reme_ai/mem_agent/v3/reme_summarizer_v3.yaml b/reme_ai/mem_agent/v3/reme_summarizer_v3.yaml new file mode 100644 index 00000000..30792a08 --- /dev/null +++ b/reme_ai/mem_agent/v3/reme_summarizer_v3.yaml @@ -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: "- (): " + {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 ``. + +user_message: | + Please perform your task based on the context. diff --git a/reme_ai/mem_tool/base_memory_tool.py b/reme_ai/mem_tool/base_memory_tool.py index 6accfe2a..0d92671c 100644 --- a/reme_ai/mem_tool/base_memory_tool.py +++ b/reme_ai/mem_tool/base_memory_tool.py @@ -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 diff --git a/reme_ai/mem_tool/read_local_memories.py b/reme_ai/mem_tool/read_local_memories.py new file mode 100644 index 00000000..98d96187 --- /dev/null +++ b/reme_ai/mem_tool/read_local_memories.py @@ -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}") diff --git a/reme_ai/mem_tool/read_local_memories.yaml b/reme_ai/mem_tool/read_local_memories.yaml new file mode 100644 index 00000000..c8e155f6 --- /dev/null +++ b/reme_ai/mem_tool/read_local_memories.yaml @@ -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. diff --git a/reme_ai/mem_tool/v3/__init__.py b/reme_ai/mem_tool/v3/__init__.py new file mode 100644 index 00000000..9d06f4c8 --- /dev/null +++ b/reme_ai/mem_tool/v3/__init__.py @@ -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", +] diff --git a/reme_ai/mem_tool/v3/add_memory.py b/reme_ai/mem_tool/v3/add_memory.py new file mode 100644 index 00000000..a013db42 --- /dev/null +++ b/reme_ai/mem_tool/v3/add_memory.py @@ -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) diff --git a/reme_ai/mem_tool/v3/read_history.py b/reme_ai/mem_tool/v3/read_history.py new file mode 100644 index 00000000..e9ab2a15 --- /dev/null +++ b/reme_ai/mem_tool/v3/read_history.py @@ -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}") diff --git a/reme_ai/mem_tool/v3/read_user_profile.py b/reme_ai/mem_tool/v3/read_user_profile.py new file mode 100644 index 00000000..57cff042 --- /dev/null +++ b/reme_ai/mem_tool/v3/read_user_profile.py @@ -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}") diff --git a/reme_ai/mem_tool/v3/retrieve_memory.py b/reme_ai/mem_tool/v3/retrieve_memory.py new file mode 100644 index 00000000..32e526d7 --- /dev/null +++ b/reme_ai/mem_tool/v3/retrieve_memory.py @@ -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") diff --git a/reme_ai/mem_tool/v3/summary_and_hands_off.py b/reme_ai/mem_tool/v3/summary_and_hands_off.py new file mode 100644 index 00000000..19d88744 --- /dev/null +++ b/reme_ai/mem_tool/v3/summary_and_hands_off.py @@ -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}" diff --git a/reme_ai/mem_tool/v3/update_user_profile.py b/reme_ai/mem_tool/v3/update_user_profile.py new file mode 100644 index 00000000..56ad4284 --- /dev/null +++ b/reme_ai/mem_tool/v3/update_user_profile.py @@ -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) diff --git a/reme_ai/mem_tool/write_local_memories.py b/reme_ai/mem_tool/write_local_memories.py new file mode 100644 index 00000000..ac1c5aae --- /dev/null +++ b/reme_ai/mem_tool/write_local_memories.py @@ -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)}" diff --git a/reme_ai/mem_tool/write_local_memories.yaml b/reme_ai/mem_tool/write_local_memories.yaml new file mode 100644 index 00000000..81615b1b --- /dev/null +++ b/reme_ai/mem_tool/write_local_memories.yaml @@ -0,0 +1,5 @@ +tool_multiple: | + Write memory nodes to local memory files. + +memory_nodes: | + List of memory nodes to write to local files. diff --git a/reme_ai/reme.py b/reme_ai/reme.py index e3aaf54d..ff9d8822 100644 --- a/reme_ai/reme.py +++ b/reme_ai/reme.py @@ -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 diff --git a/tests/test_vector_store.py b/tests/test_vector_store.py index 2f4c1f4d..6ef9ab0e 100644 --- a/tests/test_vector_store.py +++ b/tests/test_vector_store.py @@ -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 ==========