diff --git a/.gitignore b/.gitignore index 8dd24c0c..51164732 100644 --- a/.gitignore +++ b/.gitignore @@ -36,4 +36,6 @@ test_working_memory/* local_vector_store/* chroma_vector_store/* bench_results/* -meta_memory/* \ No newline at end of file +meta_memory/* +*.sqlite3 +**/data/*.json \ No newline at end of file diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 1795bc17..917a5cd8 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -3,7 +3,7 @@ repos: rev: v6.0.0 hooks: - id: check-ast - exclude: ^(test/|cookbook/) + exclude: ^(test/|cookbook/|reme_ai/|bench) - id: check-yaml - id: check-xml - id: check-toml @@ -14,18 +14,18 @@ repos: rev: v4.0.0 hooks: - id: add-trailing-comma - exclude: ^(test/|cookbook/) + exclude: ^(test/|cookbook/|reme_ai/|bench) - repo: https://github.com/psf/black rev: 25.9.0 hooks: - id: black - exclude: ^(test/|cookbook/) + exclude: ^(test/|cookbook/|reme_ai/|bench) args: [--line-length=120] - repo: https://github.com/PyCQA/flake8 rev: 7.3.0 hooks: - id: flake8 - exclude: ^(test/|cookbook/) + exclude: ^(test/|cookbook/|reme_ai/|bench) args: [ "--extend-ignore=E203", "--max-line-length=120" @@ -44,6 +44,8 @@ repos: | \.demo$ | \.md$ | \.html$ + | reme_ai/ + | bench ) args: [ --disable=W0511, @@ -76,6 +78,7 @@ repos: --disable=C3001, --disable=R1702, --disable=R0912, + --max-statements=75, --max-line-length=120, ] - repo: https://github.com/regebro/pyroma diff --git a/bench/halumem/analyze_dataset_stats.py b/bench/halumem/analyze_dataset_stats.py index 6695ad7f..061f96ae 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 @@ -37,29 +38,29 @@ class DatasetStats: total_users: int total_sessions: int total_dialogues: int - + avg_sessions_per_user: float avg_dialogues_per_session: float avg_dialogue_length_per_session: float - + # 详细分布 sessions_per_user_list: list[int] dialogues_per_session_list: list[int] dialogue_lengths_per_session_list: list[int] - + # Content 统计 total_contents: int # 所有对话回合的 content 总数 content_sizes: list[int] # 每个 content 的大小(字符数) min_content_size: int max_content_size: int percentiles: dict[str, float] # 分位点统计(全部) - + # 按 role 分类的 Content 统计 total_user_contents: int total_assistant_contents: int user_percentiles: dict[str, float] # user 角色的分位点 assistant_percentiles: dict[str, float] # assistant 角色的分位点 - + # Session 分割统计 total_chunks_after_split: int # 按 5000 字符分割后的总 chunk 数 chunks_per_user_list: list[int] # 每个用户分割后的 chunk 数量 @@ -68,14 +69,14 @@ class DatasetStats: class DatasetAnalyzer: """数据集分析器""" - + def __init__(self, data_path: str): self.data_path = data_path self.user_stats_list: list[UserStats] = [] self.all_content_sizes: list[int] = [] # 收集所有 content 的大小 self.user_content_sizes: list[int] = [] # user 角色的 content 大小 self.assistant_content_sizes: list[int] = [] # assistant 角色的 content 大小 - + @staticmethod def extract_user_name(persona_info: str) -> str: """从 persona_info 中提取用户名""" @@ -83,7 +84,7 @@ class DatasetAnalyzer: if not match: return "Unknown" return match.group(1).strip() - + @staticmethod def calculate_dialogue_length(dialogue: list[dict]) -> int: """计算对话的总长度(字符数)""" @@ -92,7 +93,7 @@ class DatasetAnalyzer: content = turn.get("content", "") total_length += len(content) return total_length - + @staticmethod def split_session_into_chunks(dialogue: list[dict], max_length: int = 5000) -> int: """ @@ -101,23 +102,23 @@ class DatasetAnalyzer: 1. 每次添加 2 个对话回合(user-assistant 对) 2. 如果添加后超过 max_length,就开始新的 chunk 3. 但是每个 chunk 至少包含 2 个对话回合 - + 返回分割后的 chunk 数量 """ if not dialogue: return 0 - + chunks = [] current_chunk = [] current_length = 0 - + # 每次处理 2 个对话回合 i = 0 while i < len(dialogue): # 取 2 个对话回合(如果不足 2 个,取剩余的) pair = dialogue[i:i+2] pair_length = sum(len(turn.get("content", "")) for turn in pair) - + # 如果当前 chunk 为空,直接添加(保证至少 2 个) if not current_chunk: current_chunk.extend(pair) @@ -136,93 +137,100 @@ class DatasetAnalyzer: current_chunk.extend(pair) current_length += pair_length i += len(pair) - + # 添加最后一个 chunk if current_chunk: chunks.append(current_chunk) - + return len(chunks) - + def load_and_analyze(self): """加载并分析数据集""" logger.info(f"Loading data from: {self.data_path}") - + with open(self.data_path, "r", encoding="utf-8") as f: for line_num, line in enumerate(f, 1): if not line.strip(): continue - + try: user_data = json.loads(line) self._analyze_user(user_data) except json.JSONDecodeError as e: logger.error(f"Error parsing line {line_num}: {e}") continue - + logger.info(f"Analyzed {len(self.user_stats_list)} users") - + def _analyze_user(self, user_data: dict): """分析单个用户的数据""" user_name = self.extract_user_name(user_data.get("persona_info", "")) uuid = user_data.get("uuid", "") sessions = user_data.get("sessions", []) - + dialogues_per_session = [] dialogue_lengths_per_session = [] + session_time_ranges = [] total_chunks = 0 - + for session in sessions: dialogue = session.get("dialogue", []) num_dialogues = len(dialogue) dialogue_length = self.calculate_dialogue_length(dialogue) - + 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 - + # 收集每个 content 的大小,并按 role 分类 for turn in dialogue: content = turn.get("content", "") content_size = len(content) role = turn.get("role", "") - + self.all_content_sizes.append(content_size) - + if role == "user": self.user_content_sizes.append(content_size) elif role == "assistant": self.assistant_content_sizes.append(content_size) - + user_stats = UserStats( user_name=user_name, uuid=uuid, 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) - + def compute_dataset_stats(self) -> DatasetStats: """计算整体数据集统计""" total_users = len(self.user_stats_list) - + sessions_per_user_list = [u.num_sessions for u in self.user_stats_list] total_sessions = sum(sessions_per_user_list) - + dialogues_per_session_list = [] dialogue_lengths_per_session_list = [] - + for user in self.user_stats_list: dialogues_per_session_list.extend(user.dialogues_per_session) dialogue_lengths_per_session_list.extend(user.dialogue_lengths_per_session) - + total_dialogues = sum(dialogues_per_session_list) - + # 计算平均值 avg_sessions_per_user = total_sessions / total_users if total_users > 0 else 0 avg_dialogues_per_session = ( @@ -232,43 +240,43 @@ class DatasetAnalyzer: sum(dialogue_lengths_per_session_list) / len(dialogue_lengths_per_session_list) if dialogue_lengths_per_session_list else 0 ) - + # Content 统计 total_contents = len(self.all_content_sizes) min_content_size = min(self.all_content_sizes) if self.all_content_sizes else 0 max_content_size = max(self.all_content_sizes) if self.all_content_sizes else 0 - + # 计算分位点 (10%, 15%, 20%, ..., 95%) percentile_points = list(range(10, 100, 5)) # 10, 15, 20, ..., 95 - + # 全部 content 的分位点 percentiles = {} if self.all_content_sizes: content_array = np.array(self.all_content_sizes) for p in percentile_points: percentiles[f"p{p}"] = float(np.percentile(content_array, p)) - + # user 角色的分位点 user_percentiles = {} if self.user_content_sizes: user_array = np.array(self.user_content_sizes) for p in percentile_points: user_percentiles[f"p{p}"] = float(np.percentile(user_array, p)) - + # assistant 角色的分位点 assistant_percentiles = {} if self.assistant_content_sizes: assistant_array = np.array(self.assistant_content_sizes) for p in percentile_points: assistant_percentiles[f"p{p}"] = float(np.percentile(assistant_array, p)) - + # Session 分割统计 chunks_per_user_list = [u.num_chunks_after_split for u in self.user_stats_list] total_chunks_after_split = sum(chunks_per_user_list) avg_chunks_per_user = ( total_chunks_after_split / total_users if total_users > 0 else 0 ) - + return DatasetStats( total_users=total_users, total_sessions=total_sessions, @@ -292,16 +300,16 @@ class DatasetAnalyzer: chunks_per_user_list=chunks_per_user_list, avg_chunks_per_user=avg_chunks_per_user ) - + @staticmethod def _print_percentiles(percentiles: dict[str, float]): """打印分位点统计(辅助函数)""" if not percentiles: print(" (无数据)") return - + sorted_percentiles = sorted(percentiles.keys(), key=lambda x: int(x[1:])) - + # 每行显示 5 个分位点,让输出更紧凑 for i in range(0, len(sorted_percentiles), 5): line_items = [] @@ -310,36 +318,36 @@ class DatasetAnalyzer: p_num = percentile_key[1:] # 去掉 'p' 前缀 line_items.append(f"{p_num}%: {percentile_value:.0f}") print(f" {' | '.join(line_items)}") - + def print_summary(self, stats: DatasetStats): """打印统计摘要""" print("\n" + "=" * 80) print("HALUMEM DATASET STATISTICS") print("=" * 80 + "\n") - + print("📊 总体统计:") print(f" 总用户数: {stats.total_users}") print(f" 总 Session 数: {stats.total_sessions}") print(f" 总对话数: {stats.total_dialogues}") - + print(f"\n📈 平均值:") print(f" 每个用户的平均 Session 数: {stats.avg_sessions_per_user:.2f}") print(f" 每个 Session 的平均对话数: {stats.avg_dialogues_per_session:.2f}") print(f" 每个 Session 的平均对话长度(字符): {stats.avg_dialogue_length_per_session:.2f}") - + print(f"\n📊 分布统计:") if stats.sessions_per_user_list: print(f" 每用户 Session 数 - 最小: {min(stats.sessions_per_user_list)}, " f"最大: {max(stats.sessions_per_user_list)}") - + if stats.dialogues_per_session_list: print(f" 每 Session 对话数 - 最小: {min(stats.dialogues_per_session_list)}, " f"最大: {max(stats.dialogues_per_session_list)}") - + if stats.dialogue_lengths_per_session_list: print(f" 每 Session 对话长度 - 最小: {min(stats.dialogue_lengths_per_session_list)}, " f"最大: {max(stats.dialogue_lengths_per_session_list)}") - + print(f"\n💬 Content 详细统计:") print(f" 总 Content 数量: {stats.total_contents}") print(f" User 消息数: {stats.total_user_contents}") @@ -347,34 +355,34 @@ class DatasetAnalyzer: print(f" Content 大小(字符数):") print(f" 最小值: {stats.min_content_size}") print(f" 最大值: {stats.max_content_size}") - + if stats.content_sizes: avg_content_size = sum(stats.content_sizes) / len(stats.content_sizes) print(f" 平均值: {avg_content_size:.2f}") - + print(f"\n📈 Content 大小分位点 (全部):") self._print_percentiles(stats.percentiles) - + print(f"\n📈 Content 大小分位点 (User 角色):") self._print_percentiles(stats.user_percentiles) - + print(f"\n📈 Content 大小分位点 (Assistant 角色):") self._print_percentiles(stats.assistant_percentiles) - + print(f"\n✂️ Session 分割统计 (按 5000 字符分割):") print(f" 原始 Session 总数: {stats.total_sessions}") print(f" 分割后 Chunk 总数: {stats.total_chunks_after_split}") print(f" 每个用户平均 Chunk 数: {stats.avg_chunks_per_user:.2f}") print(f" Chunk/Session 比例: {stats.total_chunks_after_split / stats.total_sessions:.2f}x") - + print("\n" + "=" * 80) - + def print_per_user_stats(self): """打印每个用户的详细统计""" print("\n" + "=" * 80) print("PER-USER STATISTICS") print("=" * 80 + "\n") - + for idx, user_stats in enumerate(self.user_stats_list, 1): avg_dialogues = ( sum(user_stats.dialogues_per_session) / len(user_stats.dialogues_per_session) @@ -384,24 +392,50 @@ class DatasetAnalyzer: sum(user_stats.dialogue_lengths_per_session) / len(user_stats.dialogue_lengths_per_session) if user_stats.dialogue_lengths_per_session else 0 ) - + print(f"[{idx}] {user_stats.user_name} (UUID: {user_stats.uuid[:8]}...)") print(f" Session 数: {user_stats.num_sessions}") print(f" 分割后 Chunk 数: {user_stats.num_chunks_after_split}") print(f" 平均每 Session 对话数: {avg_dialogues:.2f}") 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) print("PER-USER SESSION SPLIT SUMMARY (按 5000 字符分割)") print("=" * 80 + "\n") - + # 表头 print(f"{'序号':<6} {'用户名':<25} {'原始Sessions':<15} {'分割后Chunks':<15} {'比例':<10}") print("-" * 80) - + # 每个用户的数据 for idx, user_stats in enumerate(self.user_stats_list, 1): ratio = ( @@ -410,17 +444,17 @@ class DatasetAnalyzer: ) print(f"{idx:<6} {user_stats.user_name[:24]:<25} {user_stats.num_sessions:<15} " f"{user_stats.num_chunks_after_split:<15} {ratio:.2f}x") - + print("-" * 80) - + # 总计 total_sessions = sum(u.num_sessions for u in self.user_stats_list) total_chunks = sum(u.num_chunks_after_split for u in self.user_stats_list) overall_ratio = total_chunks / total_sessions if total_sessions > 0 else 0 - + print(f"{'总计':<6} {'':<25} {total_sessions:<15} {total_chunks:<15} {overall_ratio:.2f}x") print("=" * 80) - + def save_results(self, output_path: str, stats: DatasetStats): """保存统计结果到 JSON 文件""" results = { @@ -469,15 +503,19 @@ 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 ] } - + with open(output_path, "w", encoding="utf-8") as f: json.dump(results, f, ensure_ascii=False, indent=2) - + logger.info(f"Results saved to: {output_path}") @@ -487,24 +525,27 @@ def main(data_path: str, output_path: str = None, show_per_user: bool = False): if not Path(data_path).exists(): logger.error(f"File not found: {data_path}") return - + # 创建分析器并执行分析 analyzer = DatasetAnalyzer(data_path) analyzer.load_and_analyze() - + # 计算统计数据 stats = analyzer.compute_dataset_stats() - + # 打印摘要 analyzer.print_summary(stats) - + + # 打印第一个用户的 session 时间统计 + analyzer.print_first_user_session_times() + # 打印每个用户的分割统计摘要(始终显示) analyzer.print_user_split_summary() - + # 打印每个用户的详细统计(可选) if show_per_user: analyzer.print_per_user_stats() - + # 保存结果到文件 if output_path: analyzer.save_results(output_path, stats) @@ -516,7 +557,7 @@ def main(data_path: str, output_path: str = None, show_per_user: bool = False): if __name__ == "__main__": import argparse - + parser = argparse.ArgumentParser( description="Analyze HaluMem dataset statistics" ) @@ -537,9 +578,9 @@ if __name__ == "__main__": action="store_true", help="Show detailed statistics for each user" ) - + args = parser.parse_args() - + main( data_path=args.data_path, output_path=args.output_path, diff --git a/bench/halumem/analyze_results.py b/bench/halumem/analyze_results.py index 4814001e..630ae5e3 100644 --- a/bench/halumem/analyze_results.py +++ b/bench/halumem/analyze_results.py @@ -13,73 +13,73 @@ from typing import Dict, List, Tuple def analyze_results(tmp_dir: str = "bench_results/reme_simple/tmp"): """ 分析评估结果目录。 - + Args: tmp_dir: 临时结果目录路径 """ tmp_path = Path(tmp_dir) - + if not tmp_path.exists(): print(f"❌ 目录不存在: {tmp_dir}") return - + # 统计数据 result_counter = Counter() non_correct_results = [] # 存储非Correct结果的详细信息 - + # 遍历所有用户目录 user_dirs = sorted([d for d in tmp_path.iterdir() if d.is_dir()]) - + if not user_dirs: print(f"❌ {tmp_dir} 下没有用户目录") return - + print(f"📁 找到 {len(user_dirs)} 个用户目录\n") print("=" * 80) print("开始分析...") print("=" * 80 + "\n") - + total_sessions = 0 total_questions = 0 - + # 遍历每个用户目录 for user_dir in user_dirs: user_name = user_dir.name - + # 获取该用户的所有session文件 session_files = sorted([ - f for f in user_dir.iterdir() + f for f in user_dir.iterdir() if f.name.startswith("session_") and f.suffix == ".json" ]) - + if not session_files: continue - + # 遍历每个session for session_file in session_files: try: with open(session_file, "r", encoding="utf-8") as f: session_data = json.load(f) - + session_id = session_data.get("session_id", -1) total_sessions += 1 - + # 跳过生成的QA session if session_data.get("is_generated_qa_session", False): continue - + # 获取评估结果 eval_results = session_data.get("evaluation_results", {}) qa_records = eval_results.get("question_answering_records", []) - + # 分析每个问题的结果 for qa_idx, qa_record in enumerate(qa_records): result_type = qa_record.get("result_type", "Unknown") - + # 统计result_type result_counter[result_type] += 1 total_questions += 1 - + # 如果不是Correct,记录详细信息 if result_type != "Correct": non_correct_results.append({ @@ -91,42 +91,42 @@ def analyze_results(tmp_dir: str = "bench_results/reme_simple/tmp"): "answer": qa_record.get("answer", ""), "system_response": qa_record.get("system_response", "") }) - + except Exception as e: print(f"⚠️ 读取文件失败: {session_file}, 错误: {e}") continue - + # 输出统计结果 print("\n" + "=" * 80) print("统计结果") print("=" * 80 + "\n") - + print(f"📊 总用户数: {len(user_dirs)}") print(f"📊 总Session数: {total_sessions}") print(f"📊 总问题数: {total_questions}\n") - + if total_questions == 0: print("❌ 没有找到任何问题数据") return - + # 输出result_type分布 print("=" * 80) print("Result Type 分布") print("=" * 80 + "\n") - + # 按数量降序排列 sorted_results = sorted(result_counter.items(), key=lambda x: x[1], reverse=True) - + for result_type, count in sorted_results: ratio = count / total_questions * 100 print(f" {result_type:20s}: {count:5d} ({ratio:6.2f}%)") - + # 输出非Correct结果的详细信息 if non_correct_results: print("\n" + "=" * 80) print(f"非 Correct 结果详情 (共 {len(non_correct_results)} 条)") print("=" * 80 + "\n") - + for idx, result in enumerate(non_correct_results, 1): print(f"[{idx}] {result['result_type']}") print(f" 用户: {result['user_name']}") @@ -135,10 +135,10 @@ def analyze_results(tmp_dir: str = "bench_results/reme_simple/tmp"): print(f" 正确答案: {result['answer']}") print(f" 系统回答: {result['system_response'][:200]}{'...' if len(result['system_response']) > 200 else ''}") print() - + else: print("\n🎉 所有问题都是 Correct!") - + # 保存详细报告到文件 report_file = Path(tmp_dir).parent / "analysis_report.json" report_data = { @@ -148,16 +148,16 @@ def analyze_results(tmp_dir: str = "bench_results/reme_simple/tmp"): "total_questions": total_questions, "result_type_distribution": dict(result_counter), "result_type_ratio": { - result_type: count / total_questions + result_type: count / total_questions for result_type, count in result_counter.items() } }, "non_correct_results": non_correct_results } - + with open(report_file, "w", encoding="utf-8") as f: json.dump(report_data, f, ensure_ascii=False, indent=2) - + print("=" * 80) print(f"📄 详细报告已保存到: {report_file}") print("=" * 80) @@ -165,7 +165,7 @@ def analyze_results(tmp_dir: str = "bench_results/reme_simple/tmp"): if __name__ == "__main__": import argparse - + parser = argparse.ArgumentParser( description="分析 ReMe 评估结果中的 result_type 分布" ) @@ -175,6 +175,6 @@ if __name__ == "__main__": default="bench_results/reme_simple/tmp", help="临时结果目录路径 (默认: bench_results/reme_simple/tmp)" ) - + args = parser.parse_args() analyze_results(args.tmp_dir) diff --git a/bench/halumem/compute_qa_stats_v4.py b/bench/halumem/compute_qa_stats_v4.py new file mode 100644 index 00000000..afd07e4b --- /dev/null +++ b/bench/halumem/compute_qa_stats_v4.py @@ -0,0 +1,285 @@ +""" +Compute Question Answering statistics from eval_reme_simple_v4.py results. + +Usage: + python bench/halumem/compute_qa_stats_v4.py --results_file bench_results/reme_simple_v4/eval_results.jsonl + python bench/halumem/compute_qa_stats_v4.py --tmp_dir bench_results/reme_simple_v4/tmp +""" + +import json +import os +from pathlib import Path +from typing import Any + + +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 = hallucination = omission = valid = 0 + + for qa in qa_records: + result_type = qa.get("result_type", "") + if result_type == "Correct": + correct += 1 + valid += 1 + elif result_type == "Hallucination": + hallucination += 1 + valid += 1 + elif result_type == "Omission": + omission += 1 + valid += 1 + + metrics = { + "correct_qa_ratio(all)": correct / total, + "hallucination_qa_ratio(all)": hallucination / total, + "omission_qa_ratio(all)": omission / total, + "correct_qa_ratio(valid)": correct / valid if valid > 0 else 0, + "hallucination_qa_ratio(valid)": hallucination / valid if valid > 0 else 0, + "omission_qa_ratio(valid)": omission / valid if valid > 0 else 0, + "qa_valid_num": valid, + "qa_num": total + } + + return metrics + + +def compute_time_metrics(results_file: str) -> dict[str, float]: + """Compute timing metrics from evaluation results.""" + add_duration = search_duration = 0 + + 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.get("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) + + 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 + } + + +def load_from_tmp_dir(tmp_dir: str) -> str: + """Load data from tmp directory and generate eval_results.jsonl file.""" + tmp_path = Path(tmp_dir) + eval_results_file = tmp_path.parent / "eval_results.jsonl" + + print(f"\n📁 Loading from: {tmp_dir}") + print(f"📝 Generating: {eval_results_file}") + + user_dirs = [d for d in tmp_path.iterdir() if d.is_dir()] + print(f" Found {len(user_dirs)} users") + + users_data = [] + for user_dir in user_dirs: + session_files = sorted( + [f for f in user_dir.iterdir() if f.name.startswith("session_") and f.suffix == ".json"], + key=lambda f: int(f.stem.split("_")[1]) + ) + + if not session_files: + continue + + with open(session_files[0], "r", encoding="utf-8") as f: + first_session = json.load(f) + + user_data = { + "uuid": first_session["uuid"], + "user_name": first_session["user_name"], + "sessions": [] + } + + for session_file in session_files: + with open(session_file, "r", encoding="utf-8") as f: + session_data = json.load(f) + session_data.pop("uuid", None) + session_data.pop("user_name", None) + user_data["sessions"].append(session_data) + + users_data.append(user_data) + print(f" ✓ {user_dir.name}: {len(session_files)} sessions") + + with open(eval_results_file, "w", encoding="utf-8") as f: + for user_data in users_data: + f.write(json.dumps(user_data, ensure_ascii=False) + "\n") + + print(f" ✅ Generated: {eval_results_file}") + return str(eval_results_file) + + +def main(input_path: str): + """Main function to compute statistics from eval results.""" + + if not os.path.exists(input_path): + print(f"❌ Error: Path not found: {input_path}") + return + + print("\n" + "=" * 80) + print("REME V4 - QUESTION ANSWERING STATISTICS") + print("=" * 80) + + # Load or generate eval_results.jsonl + if os.path.isdir(input_path): + results_file = load_from_tmp_dir(input_path) + else: + results_file = input_path + print(f"\n📁 Using: {results_file}") + + # Collect QA records with metadata + qa_records = [] + qa_with_metadata = [] + user_count = session_count = 0 + + with open(results_file, "r", encoding="utf-8") as f: + for line in f: + if not line.strip(): + continue + user_data = json.loads(line) + user_count += 1 + user_name = user_data.get("user_name", "Unknown") + + valid_session_idx = 0 + for original_idx, session in enumerate(user_data.get("sessions", [])): + if session.get("is_generated_qa_session"): + continue + + session_count += 1 + eval_results = session.get("evaluation_results", {}) + + for qa_idx, qa in enumerate(eval_results.get("question_answering_records", [])): + qa_records.append(qa) + qa_with_metadata.append({ + "user_name": user_name, + "session_idx": valid_session_idx, + "question_idx": qa_idx, + "qa_record": qa + }) + + valid_session_idx += 1 + + print(f"\n📊 Data Summary:") + print(f" Users: {user_count}") + print(f" Sessions: {session_count}") + print(f" QA Records: {len(qa_records)}") + + # Compute metrics + qa_metrics = compute_qa_metrics(qa_records) + time_metrics = compute_time_metrics(results_file) + + # Save results + output_dir = Path(results_file).parent + report_file = output_dir / "reme_eval_stat_result.json" + + final_results = { + "overall_score": { + "question_answering": qa_metrics, + "time_consuming": time_metrics + }, + "question_answering_records": qa_records + } + + with open(report_file, "w", encoding="utf-8") as f: + json.dump(final_results, f, ensure_ascii=False, indent=4) + + print(f"\n✅ Results saved to: {report_file}") + + # Print metrics + print("\n" + "=" * 80) + print("📊 QUESTION ANSWERING METRICS") + print("=" * 80) + print(f"\n 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 error records + print("\n" + "=" * 80) + print("❌ ERROR RECORDS (Non-Correct)") + print("=" * 80) + + error_records = [r for r in qa_with_metadata if r["qa_record"].get("result_type") not in ["Correct", ""]] + + if not error_records: + print("\n✅ All QA records are correct!") + else: + print(f"\nFound {len(error_records)} error records:\n") + + for idx, record in enumerate(error_records, 1): + qa = record["qa_record"] + + print(f"\n{'━' * 80}") + print(f"❌ ERROR #{idx}") + print(f"{'━' * 80}") + print(f"👤 User: {record['user_name']}") + print(f"📅 Session: {record['session_idx']} | Question: {record['question_idx']}") + print(f"🏷️ Result Type: {qa.get('result_type', 'Unknown')}") + print(f"\n❓ Question:") + print(f" {qa.get('question', 'N/A')}") + print(f"\n✅ Expected Answer:") + print(f" {qa.get('answer', 'N/A')}") + print(f"\n🤖 System Response:") + print(f" {qa.get('system_response', 'N/A')}") + print(f"\n💭 Reasoning:") + reason = qa.get('question_answering_reasoning', 'N/A') + # Wrap long reasoning text + if len(reason) > 80: + words = reason.split() + lines = [] + current_line = " " + for word in words: + if len(current_line) + len(word) + 1 <= 80: + current_line += word + " " + else: + lines.append(current_line.rstrip()) + current_line = " " + word + " " + if current_line.strip(): + lines.append(current_line.rstrip()) + print("\n".join(lines)) + else: + print(f" {reason}") + + print("\n" + "=" * 80) + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser(description="Compute QA statistics from eval_reme_simple_v4.py results") + parser.add_argument("--results_file", type=str, help="Path to eval_results.jsonl file") + parser.add_argument("--tmp_dir", type=str, help="Path to tmp directory") + + args = parser.parse_args() + + if args.tmp_dir: + main(input_path=args.tmp_dir) + elif args.results_file: + main(input_path=args.results_file) + else: + parser.error("Either --results_file or --tmp_dir must be provided") diff --git a/bench/halumem/compute_stats_from_tmp.py b/bench/halumem/compute_stats_from_tmp.py index 3891c21a..2d0201ec 100644 --- a/bench/halumem/compute_stats_from_tmp.py +++ b/bench/halumem/compute_stats_from_tmp.py @@ -9,7 +9,7 @@ are available in the tmp directory. It will: 4. Aggregate results and compute metrics Usage: - python bench/halumem/compute_stats_from_tmp.py --tmp_dir bench_results/reme/tmp + python bench/halumem/compute_stats_from_tmp.py --tmp_dir bench_results/reme_simple_v4/tmp """ import asyncio @@ -358,7 +358,7 @@ async def main_async(tmp_dir: str): # Determine paths parent_dir = os.path.dirname(tmp_dir) frame = "reme" - + output_file_stage1 = os.path.join(parent_dir, f"{frame}_eval_results.jsonl") output_file_stage2 = os.path.join(parent_dir, f"{frame}_eval_stat_result.json") @@ -392,7 +392,7 @@ async def main_async(tmp_dir: str): # Load all users and process user_data_list = list(enumerate(iter_jsonl(output_file_stage1), 1)) - + for idx, user_data in user_data_list: uuid = user_data["uuid"] tmp_file = os.path.join(tmp_dir2, f"{uuid}.json") diff --git a/bench/halumem/eval_baseline_simple.py b/bench/halumem/eval_baseline_simple.py index 6da0bbed..34be3d0a 100644 --- a/bench/halumem/eval_baseline_simple.py +++ b/bench/halumem/eval_baseline_simple.py @@ -44,13 +44,13 @@ class EvalConfig: 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.""" @@ -58,7 +58,7 @@ class DataLoader: 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.""" @@ -74,7 +74,7 @@ class DataLoader: } for turn in dialogue ] - + @staticmethod def format_dialogue_for_eval(dialogue: list[dict], user_name: str = None) -> str: """Format dialogue into string for evaluation (only user messages).""" @@ -83,14 +83,14 @@ class DataLoader: # Skip assistant messages - only include user messages if turn['role'] != 'user': continue - + 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 provided role = user_name if user_name else 'user' - + formatted_turns.append( f"Role: {role}\n" f"Content: {turn['content']}\n" @@ -101,29 +101,29 @@ class DataLoader: 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) @@ -131,38 +131,38 @@ class FileManager: 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" + 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() + 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: @@ -171,7 +171,7 @@ class FileManager: 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") @@ -206,10 +206,10 @@ Please respond in JSON format with the following structure: class BaselineQuestionAnsweringEvaluator: """Evaluates question answering performance using direct LLM inference (no memory system).""" - + def __init__(self): pass - + async def answer_question( self, question: str, @@ -217,18 +217,18 @@ class BaselineQuestionAnsweringEvaluator: ) -> tuple[str, str, float]: """ Answer a question using the dialogue history directly. - + Returns: tuple: (answer, reasoning, duration_ms) """ start = time.time() - + # Format prompt prompt = BASELINE_QA_PROMPT.format( dialogue=formatted_dialogue, question=question ) - + # Get answer from LLM try: # model_name = "qwen3-max" @@ -240,10 +240,10 @@ class BaselineQuestionAnsweringEvaluator: logger.error(f"Error getting answer from LLM: {e}") answer = "Error: Failed to get answer" reasoning = str(e) - + duration_ms = (time.time() - start) * 1000 return answer, reasoning, duration_ms - + async def evaluate_questions( self, questions: list[dict], @@ -254,14 +254,14 @@ class BaselineQuestionAnsweringEvaluator: ) -> list[dict]: """Evaluate all questions for a session.""" results = [] - + for qa in questions: # Get answer directly from LLM answer, reasoning, duration_ms = await self.answer_question( question=qa["question"], formatted_dialogue=formatted_dialogue ) - + # Evaluate response evidence_text = "\n".join([e["memory_content"] for e in qa["evidence"]]) eval_result = await evaluation_for_question2( @@ -271,7 +271,7 @@ class BaselineQuestionAnsweringEvaluator: answer, formatted_dialogue ) - + # Build result record qa_result = { **qa, @@ -284,13 +284,13 @@ class BaselineQuestionAnsweringEvaluator: "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.""" @@ -306,15 +306,15 @@ class MetricsAggregator: "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": @@ -323,7 +323,7 @@ class MetricsAggregator: hallucination += 1 elif result_type == "Omission": omission += 1 - + metrics = { "correct_qa_ratio(all)": correct / total, "hallucination_qa_ratio(all)": hallucination / total, @@ -331,7 +331,7 @@ class MetricsAggregator: "qa_valid_num": valid, "qa_num": total } - + if valid > 0: metrics.update({ "correct_qa_ratio(valid)": correct / valid, @@ -344,25 +344,25 @@ class MetricsAggregator: "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.""" answer_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"]: eval_results = session.get("evaluation_results", {}) for qa in eval_results.get("question_answering_records", []): answer_duration += qa.get("answer_duration_ms", 0) - + # Convert to minutes return { "answer_duration_time": answer_duration / 1000 / 60, @@ -374,13 +374,13 @@ class MetricsAggregator: class HaluMemBaselineEvaluator: """Main evaluator orchestrating the baseline evaluation pipeline.""" - + def __init__(self, config: EvalConfig): self.config = config self.file_manager = FileManager(config.output_dir) self.qa_evaluator = BaselineQuestionAnsweringEvaluator() self.data_loader = DataLoader() - + async def process_session( self, session: dict, @@ -395,16 +395,16 @@ class HaluMemBaselineEvaluator: "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 - + # Store dialogue dialogue = session["dialogue"] session_data["dialogue"] = dialogue - + # Evaluate questions if present if "questions" in session: formatted_dialogue = self.data_loader.format_dialogue_for_eval(dialogue, user_name) @@ -415,25 +415,25 @@ class HaluMemBaselineEvaluator: 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"] - + total_sessions = len(user_data["sessions"]) logger.info(f"Processing user: {user_name} ({total_sessions} sessions)") - + # Semaphore for concurrency control within user sessions semaphore = asyncio.Semaphore(self.config.max_concurrency) completed_count = [0] # Use list to allow modification in nested async function - + async def process_session_with_log(idx: int, session: dict): async with semaphore: session_data = await self.process_session( @@ -442,65 +442,65 @@ class HaluMemBaselineEvaluator: user_name=user_name, uuid=uuid ) - + self.file_manager.save_session(user_name, idx, session_data) - + # Update and log completion completed_count[0] += 1 print(f"✅ {user_name} complete {completed_count[0]}/{total_sessions}") - + # Process all sessions in parallel tasks = [ process_session_with_log(idx, session) for idx, session in enumerate(user_data["sessions"]) ] await asyncio.gather(*tasks) - + return {"uuid": uuid, "user_name": user_name, "status": "ok"} - + async def run_evaluation(self): """Run the complete evaluation pipeline.""" start_time = time.time() - + # 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 BASELINE EVALUATION - DIRECT QA WITHOUT MEMORY SYSTEM") print(f"Users: {len(users_to_process)} | Session Concurrency: {self.config.max_concurrency}") print("=" * 80 + "\n") - + # Process users sequentially (for loop) for idx, user_data in enumerate(users_to_process, 1): 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)") continue - + print(f"🔄 [{idx}/{len(users_to_process)}] Processing {user_name}...") await self.process_user(user_data) print(f"✅ [{idx}/{len(users_to_process)}] User {user_name} completed\n") - + # 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: @@ -508,20 +508,20 @@ class HaluMemBaselineEvaluator: 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, @@ -529,23 +529,23 @@ class HaluMemBaselineEvaluator: }, "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") 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}") @@ -554,7 +554,7 @@ class HaluMemBaselineEvaluator: 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" Answer Duration: {time_metrics['answer_duration_time']:.2f} min") print(f" Total: {time_metrics['total_duration_time']:.2f} min") @@ -574,14 +574,14 @@ def main( user_num=user_num, max_concurrency=max_concurrency ) - + evaluator = HaluMemBaselineEvaluator(config) asyncio.run(evaluator.run_evaluation()) if __name__ == "__main__": import argparse - + parser = argparse.ArgumentParser( description="Evaluate Baseline (Direct QA) on HaluMem benchmark" ) @@ -603,9 +603,9 @@ if __name__ == "__main__": default=2, help="Maximum concurrent user processing (default: 2)" ) - + args = parser.parse_args() - + main( data_path=args.data_path, user_num=args.user_num, diff --git a/bench/halumem/eval_reme.py b/bench/halumem/eval_reme.py index d3044535..2b943143 100644 --- a/bench/halumem/eval_reme.py +++ b/bench/halumem/eval_reme.py @@ -685,7 +685,7 @@ async def main_async( user_data_list = user_data_list[:total_users] print(f"Processing {total_users} users with max concurrency {max_concurrency}...") - + # Create semaphore to limit concurrency for Stage 1 semaphore_stage1 = asyncio.Semaphore(max_concurrency) @@ -694,11 +694,11 @@ async def main_async( async with semaphore_stage1: uuid = user_data['uuid'] tmp_file = os.path.join(tmp_dir, f"{uuid}.json") - + if os.path.exists(tmp_file): print(f"⚡ Skipping user {uuid} ({idx}/{total_users}) — cached result found.") return {"uuid": uuid, "status": "cached", "path": tmp_file} - + print(f"[{idx}/{total_users}] Processing user {uuid}...") result = await process_user_stage1(user_data, top_k, save_path) print(f"[{idx}/{total_users}] ✅ Finished {uuid} ({result['status']})") @@ -733,7 +733,7 @@ async def main_async( # Load all users and process sequentially user_data_list = list(enumerate(iter_jsonl(output_file_stage1), 1)) - + for idx, user_data in user_data_list: uuid = user_data["uuid"] tmp_file = os.path.join(tmp_dir2, f"{uuid}.json") diff --git a/bench/halumem/eval_reme_simple.py b/bench/halumem/eval_reme_simple.py index aff0934d..17671def 100644 --- a/bench/halumem/eval_reme_simple.py +++ b/bench/halumem/eval_reme_simple.py @@ -48,13 +48,13 @@ class EvalConfig: 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.""" @@ -62,7 +62,7 @@ class DataLoader: 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.""" @@ -78,7 +78,7 @@ class DataLoader: } for turn in dialogue ] - + @staticmethod def format_dialogue_for_eval(dialogue: list[dict], user_name: str = None) -> str: """Format dialogue into string for evaluation.""" @@ -87,10 +87,10 @@ class DataLoader: 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" @@ -101,29 +101,29 @@ class DataLoader: 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) @@ -131,38 +131,38 @@ class FileManager: 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" + 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() + 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: @@ -171,7 +171,7 @@ class FileManager: 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") @@ -179,19 +179,19 @@ class FileManager: class MemoryProcessor: """Handles ReMe memory operations.""" - + def __init__(self, reme: ReMe): self.reme = reme - + async def add_memories( - self, - user_id: str, + self, + user_id: str, messages: list[dict], batch_size: int = 20 ) -> tuple[list[str], list[list[dict]], float]: """ Add memories in batches and return extracted memory contents. - + Returns: tuple: (extracted_memories, agent_messages, total_duration_ms) """ @@ -202,19 +202,19 @@ class MemoryProcessor: for i in range(0, len(messages), batch_size): batch = messages[i:i + batch_size] start = time.time() - + memory_nodes, agent_messages, success = await self.reme.summary_v2( - messages=batch, + 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: @@ -230,23 +230,23 @@ class MemoryProcessor: 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, + self, + query: str, + user_id: str, top_k: int = 20 ) -> tuple[str, list, float]: """ Search memory and return response. - + Returns: tuple: (response, agent_messages, duration_ms) """ start = time.time() response, agent_messages, success = await self.reme.retrieve_v2( - query=query, - user_id=user_id, + query=query, + user_id=user_id, top_k=top_k ) duration_ms = (time.time() - start) * 1000 @@ -257,11 +257,11 @@ class MemoryProcessor: 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], @@ -272,7 +272,7 @@ class QuestionAnsweringEvaluator: ) -> list[dict]: """Evaluate all questions for a session.""" results = [] - + for qa in questions: # Search memory for answer response, agent_messages, duration_ms = await self.memory_processor.search_memory( @@ -280,7 +280,7 @@ class QuestionAnsweringEvaluator: 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( @@ -290,7 +290,7 @@ class QuestionAnsweringEvaluator: response, formatted_dialogue ) - + # Build result record qa_result = { **qa, @@ -303,13 +303,13 @@ class QuestionAnsweringEvaluator: "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.""" @@ -325,15 +325,15 @@ class MetricsAggregator: "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": @@ -342,7 +342,7 @@ class MetricsAggregator: hallucination += 1 elif result_type == "Omission": omission += 1 - + metrics = { "correct_qa_ratio(all)": correct / total, "hallucination_qa_ratio(all)": hallucination / total, @@ -350,7 +350,7 @@ class MetricsAggregator: "qa_valid_num": valid, "qa_num": total } - + if valid > 0: metrics.update({ "correct_qa_ratio(valid)": correct / valid, @@ -363,28 +363,28 @@ class MetricsAggregator: "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, @@ -397,18 +397,18 @@ class MetricsAggregator: class HaluMemEvaluator: """Main evaluator orchestrating the entire 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, + self.memory_processor, config.top_k ) self.data_loader = DataLoader() - + async def process_session( self, session: dict, @@ -423,29 +423,29 @@ class HaluMemEvaluator: "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 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) @@ -456,90 +456,90 @@ class HaluMemEvaluator: 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.""" 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 - 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) + 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: @@ -547,20 +547,20 @@ class HaluMemEvaluator: 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, @@ -568,23 +568,23 @@ class HaluMemEvaluator: }, "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") 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}") @@ -593,7 +593,7 @@ class HaluMemEvaluator: 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") @@ -616,14 +616,14 @@ def main( user_num=user_num, max_concurrency=max_concurrency ) - + evaluator = HaluMemEvaluator(config) asyncio.run(evaluator.run_evaluation()) if __name__ == "__main__": import argparse - + parser = argparse.ArgumentParser( description="Evaluate ReMe on HaluMem benchmark (Question Answering)" ) @@ -651,9 +651,9 @@ if __name__ == "__main__": default=2, help="Maximum concurrent user processing (default: 2)" ) - + args = parser.parse_args() - + main( data_path=args.data_path, top_k=args.top_k, diff --git a/bench/halumem/eval_reme_simple_v3.py b/bench/halumem/eval_reme_simple_v3.py new file mode 100644 index 00000000..469102ba --- /dev/null +++ b/bench/halumem/eval_reme_simple_v3.py @@ -0,0 +1,676 @@ +""" +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 shutil +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() + + # Clear meta_memory directory + meta_memory_path = Path(f"meta_memory/{self.reme.vector_store.collection_name}") + if meta_memory_path.exists(): + shutil.rmtree(meta_memory_path) + logger.info(f"Cleared meta_memory directory: {meta_memory_path}") + meta_memory_path.mkdir(parents=True, exist_ok=True) + + # 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/bench/halumem/eval_reme_simple_v4.py b/bench/halumem/eval_reme_simple_v4.py new file mode 100644 index 00000000..c68ee4d1 --- /dev/null +++ b/bench/halumem/eval_reme_simple_v4.py @@ -0,0 +1,690 @@ +""" +HaluMem Benchmark Evaluator for ReMe - Question Answering + +A modular evaluation pipeline that: +1. Loads HaluMem benchmark data +2. Processes user sessions through ReMe (summarization + retrieval) +3. Evaluates question answering performance +4. Generates comprehensive metrics + +Usage: + python bench/halumem/eval_reme_simple_v4.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 shutil +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, answer_question_with_memories +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_v4" + + +# ==================== 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 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 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() + + memory_nodes, agent_messages, success = await self.reme.summary_v4( + 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[dict, list, float]: + """ + Search memory using ReMe and return structured answer with reasoning. + + Returns: + tuple: (answer_dict, agent_messages, duration_ms) + answer_dict contains: {"reasoning": str, "answer": str, "memories": str} + """ + start = time.time() + + # Retrieve memories from ReMe + memories_response, agent_messages, success = await self.reme.retrieve_v4( + query=query, + user_id=user_id, + top_k=top_k + ) + + # Use LLM to generate structured answer from memories + answer_result = await answer_question_with_memories( + question=query, + memories=memories_response, + user_id=user_id + ) + + # Add original memories to the result + answer_result["memories"] = memories_response + + duration_ms = (time.time() - start) * 1000 + return answer_result, 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: + answer_dict, agent_messages, duration_ms = await self.memory_processor.search_memory( + query=qa["question"], + user_id=user_name, + top_k=self.top_k + ) + + # Extract answer and reasoning from the structured response + system_answer = answer_dict.get("answer", "") + system_reasoning = answer_dict.get("reasoning", "") + retrieved_memories = answer_dict.get("memories", "") + + # 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, + system_answer, + formatted_dialogue + ) + + # Build result record + qa_result = { + **qa, + "uuid": uuid, + "session_id": session_id, + "system_response": system_answer, + "system_reasoning": system_reasoning, + "retrieved_memories": retrieved_memories, + "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 HaluMemEvaluatorV4: + + 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.""" + 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 + + 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.""" + start_time = time.time() + + # Clear existing data + await self.reme.vector_store.delete_all() + + # Clear meta_memory directory + meta_memory_path = Path(f"meta_memory/{self.reme.vector_store.collection_name}") + if meta_memory_path.exists(): + shutil.rmtree(meta_memory_path) + logger.info(f"Cleared meta_memory directory: {meta_memory_path}") + meta_memory_path.mkdir(parents=True, exist_ok=True) + + # 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 - 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") + 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 evaluation.""" + config = EvalConfig( + data_path=data_path, + top_k=top_k, + user_num=user_num, + max_concurrency=max_concurrency + ) + + evaluator = HaluMemEvaluatorV4(config) + asyncio.run(evaluator.run_evaluation()) + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser( + description="Evaluate ReMe 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/bench/halumem/eval_tools.py b/bench/halumem/eval_tools.py index bf68ce2d..e0fae45d 100644 --- a/bench/halumem/eval_tools.py +++ b/bench/halumem/eval_tools.py @@ -120,7 +120,8 @@ async def evaluation_for_question2( dialogue: The formatted dialogue history (role, content, time_created). """ - prompt = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION2"].format( + # prompt = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION2"].format( + prompt = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION"].format( question=question, reference_answer=reference_answer, key_memory_points=key_memory_points, @@ -128,6 +129,43 @@ async def evaluation_for_question2( dialogue=dialogue, ) - result = await llm_request_for_json(prompt) + result = await llm_request_for_json(prompt, model_name="qwen3-max") + + return result + + +async def answer_question_with_memories( + question: str, + memories: str, + user_id: str = None, +): + """ + Answer a question using retrieved memories with PROMPT_MEMZERO_JSON template. + + Args: + question: The question to answer + memories: The retrieved memories (formatted as context) + user_id: Optional user ID for context formatting + + Returns: + dict with 'reasoning' and 'answer' fields + """ + # Format context with memories + if user_id: + context = _PROMPTS["TEMPLATE_MEMOS"].format( + user_id=user_id, + memories=memories + ) + else: + context = f"Memories:\n{memories}" + + # Use PROMPT_MEMZERO_JSON template for structured JSON response + prompt = _PROMPTS["PROMPT_MEMZERO_JSON"].format( + context=context, + question=question + ) + + # result = await llm_request_for_json(prompt, model_name="qwen3-max") + result = await llm_request_for_json(prompt, model_name="qwen3-30b-a3b-instruct-2507") return result diff --git a/bench/halumem/halumem.yaml b/bench/halumem/halumem.yaml index 60f7cbda..56dde2db 100644 --- a/bench/halumem/halumem.yaml +++ b/bench/halumem/halumem.yaml @@ -2,6 +2,54 @@ TEMPLATE_MEMOS: | Memories for user {user_id}: {memories} +PROMPT_MEMZERO_JSON: | + # CONTEXT: + {context} + + # CONTEXT PRIORITY: + When the context contains information from multiple sources, follow this strict priority order: + 1. **Historical Dialogue** (highest priority) - Direct conversation content + 2. **Extracted Memories** (medium priority) - Summarized memory points + 3. **User Profile** (lowest priority) - General user information + + # Question: + {question} + + # OUTPUT FORMAT: + Do not hallucinate; strictly answer the user's question based on the content of the CONTEXT. + Please provide your response in the following JSON format: + + ```json + {{ + "reasoning": "reasoning content", + "answer": "Provide a detailed answer" + }} + ``` + +PROMPT_MEMZERO_JSON2: | + # CONTEXT: + {context} + + # CONTEXT PRIORITY: + When the context contains information from multiple sources, follow this strict priority order: + 1. **Historical Dialogue** (highest priority) - Direct conversation content + 2. **Extracted Memories** (medium priority) - Summarized memory points + 3. **User Profile** (lowest priority) - General user information + + # Question: + {question} + + # OUTPUT FORMAT: + Do not hallucinate; strictly answer the user's question based on the content of the CONTEXT. + Please provide your response in the following JSON format: + + ```json + {{ + "reasoning": "reasoning content", + "answer": "Provide a detailed answer" + }} + ``` + PROMPT_MEMZERO: | You are an intelligent memory assistant tasked with retrieving accurate information from conversation memories. @@ -419,15 +467,12 @@ EVALUATION_PROMPT_FOR_QUESTION: | "evaluation_result": "Correct | Hallucination | Omission" }} ``` - + EVALUATION_PROMPT_FOR_QUESTION2: | You are an **evaluation expert for AI memory system question answering**. - **Dialogue:** - {dialogue} - - Based **only** on the provided **“Question”**, **“Reference Answer”**, and **“Key Memory Points”** (the essential facts needed to derive the reference answer), strictly evaluate the **accuracy** of the **“Memory System Response.”** Classify it as one of **“Correct”**, **“Hallucination”**, or **“Omission.”** Do **not** use any external knowledge or subjective inference. Finally, output your judgment **strictly** in the specified JSON format. + Based **only** on the provided **"Question"**, **"Reference Answer"**, and **"Key Memory Points"** (the essential facts needed to derive the reference answer), strictly evaluate the **accuracy** of the **"Memory System Response."** Classify it as one of **"Correct"**, **"Hallucination"**, or **"Omission."** Do **not** use any external knowledge or subjective inference. Finally, output your judgment **strictly** in the specified JSON format. # Evaluation Criteria @@ -435,36 +480,45 @@ EVALUATION_PROMPT_FOR_QUESTION2: | ### 1. Correct - * The “Memory System Response” accurately answers the “Question,” and its content is **semantically equivalent** to the “Reference Answer.” - * It contains **no contradictions** with the “Key Memory Points” or “Reference Answer.” - * It introduces **no unsupported details** beyond the “Key Memory Points” that could alter the conclusion. + * The "Memory System Response" accurately answers the "Question," and its content is **semantically equivalent** to the "Reference Answer." + * It contains **no contradictions** with the "Key Memory Points" or "Reference Answer." + * **Extra details not present in the Key Memory Points are allowed and should not be penalized**, as long as they: + - Do not contradict the Key Memory Points or Reference Answer + - Do not change or mislead the core conclusion + - Are reasonable additional context that the memory system may have retained from the conversation + * The memory system may have stored additional information beyond the Key Memory Points. Such extra information should be treated as **supplementary context** rather than hallucination, provided it does not conflict with the core answer. * Synonyms, paraphrasing, and reasonable summarization are acceptable. ### 2. Hallucination - * The “Memory System Response” includes information or facts that **contradict or are inconsistent** with the “Reference Answer” or the “Key Memory Points.” - * When the “Reference Answer” is labeled as *unknown/uncertain*, yet the response provides a specific verifiable fact or conclusion. - * Extra irrelevant information that does **not change** the conclusion is **not** considered hallucination by itself; however, if it **changes or misleads** the conclusion, or **contradicts** the “Key Memory Points,” it should be judged as a **Hallucination**. + * The "Memory System Response" includes information or facts that **contradict or are inconsistent** with the "Reference Answer" or the "Key Memory Points." + * The response provides information that **directly contradicts** known facts from the Key Memory Points. + * When the "Reference Answer" is labeled as *unknown/uncertain*, yet the response provides a specific verifiable fact or conclusion. + * **Important:** Extra information that is NOT in Key Memory Points is **NOT automatically a hallucination**. Only classify as hallucination if the extra information: + - Directly contradicts the Key Memory Points or Reference Answer + - Changes or misleads the core conclusion in a way that makes the answer incorrect + - Provides a definitive answer when the Reference Answer indicates uncertainty ### 3. Omission - * The response is **incomplete** compared to the “Reference Answer.” - * It explicitly states “don’t know,” “can’t remember,” or “no related memory,” even though relevant information exists in the “Key Memory Points.” + * The response is **incomplete** compared to the "Reference Answer." + * It explicitly states "don't know," "can't remember," or "no related memory," even though relevant information exists in the "Key Memory Points." * For multi-element questions, **all elements must be correct and present**; omission of **any** element is considered an **Omission**. ## Priority Rules (Conflict Handling) * If the response contains **both missing necessary information** and **fabricated/contradictory information**, classify it as **Hallucination**. * If there is **no fabrication/contradiction** but some necessary information is missing, classify it as **Omission**. - * Only when the meaning is **fully equivalent** to the reference answer should it be classified as **Correct**. + * If the core answer is correct and complete, classify as **Correct** even if there are extra details not in Key Memory Points (as long as they don't contradict or mislead). ## Detailed Guidelines and Tolerance * Equivalent expressions of numbers, times, and units are acceptable, but the **numerical values themselves must not differ**. * For multi-element questions, **all elements must be complete and accurate**; missing any element counts as **Omission**. - * If the reference answer is *“unknown / cannot be determined”* and the system provides a definite fact, that is a **Hallucination**. - If the system also answers *“unknown”* (without guessing), it may be **Correct**. - * The evaluation must rely **only** on the *Reference Answer*, *Key Memory Points*, and *System Response* — no external context, world knowledge, or speculative reasoning is allowed. + * If the reference answer is *"unknown / cannot be determined"* and the system provides a definite fact, that is a **Hallucination**. + If the system also answers *"unknown"* (without guessing), it may be **Correct**. + * **Focus on evaluating whether the core answer to the question is correct**, not whether the response is limited to only the Key Memory Points. + * Extra contextual information (e.g., additional preferences, related details) should be viewed as enrichment, not as errors, unless they contradict or mislead. # Information for Evaluation @@ -487,7 +541,7 @@ EVALUATION_PROMPT_FOR_QUESTION2: | ```json {{ - "reasoning": "Provide a concise and traceable evaluation rationale: first compare the system’s response with the Key Memory Points (which were correctly used, which were missing, and whether there was any fabrication/contradiction), then assess its consistency with the Reference Answer, and finally state the classification basis.", + "reasoning": "Provide a concise and traceable evaluation rationale: first verify that the system's response correctly includes all required elements from the Reference Answer, then check if any information contradicts the Key Memory Points or Reference Answer. Extra details not in Key Memory Points should be noted but not penalized unless they contradict or mislead. Finally state the classification basis.", "evaluation_result": "Correct | Hallucination | Omission" }} ``` diff --git a/bench/halumem/llms.py b/bench/halumem/llms.py index a6a06218..73577575 100644 --- a/bench/halumem/llms.py +++ b/bench/halumem/llms.py @@ -28,12 +28,12 @@ reme = ReMe() ) async def llm_request(prompt, model_name: str = "qwen3-max", **kwargs) -> str: """Make an LLM request using ReMe's LLM with optional model override. - + Args: prompt: The prompt to send to the LLM model_name: Optional model name to override the default model (default: "qwen3-max") **kwargs: Additional arguments to pass to the chat method - + Returns: The assistant's response content """ @@ -58,17 +58,18 @@ async def llm_request(prompt, model_name: str = "qwen3-max", **kwargs) -> str: reraise=True, before_sleep=before_sleep_log(logger, logging.WARNING), ) -async def llm_request_for_json(prompt, model_name: str = "qwen3-max", **kwargs): +async def llm_request_for_json(prompt, model_name: str = "qwen-flash", **kwargs): + # async def llm_request_for_json(prompt, model_name: str = "qwen3-max", **kwargs): """Make an LLM request expecting JSON response using ReMe's LLM. - + Args: prompt: The prompt to send to the LLM model_name: Optional model name to override the default model (default: "qwen3-max") **kwargs: Additional arguments to pass to the chat method - + Returns: Parsed JSON object from the LLM response - + Raises: ValueError: If no JSON block is found in the model output """ diff --git a/reme_ai/mem_tool/history/__init__.py b/bench/human_in_the_loop/__init__.py similarity index 100% rename from reme_ai/mem_tool/history/__init__.py rename to bench/human_in_the_loop/__init__.py diff --git a/bench/human_in_the_loop/compute_qa_stats.py b/bench/human_in_the_loop/compute_qa_stats.py new file mode 100644 index 00000000..69d65e55 --- /dev/null +++ b/bench/human_in_the_loop/compute_qa_stats.py @@ -0,0 +1,237 @@ +""" +Compute Question Answering statistics from evaluation results in tmp directory. +""" + +import json +from collections import defaultdict +from pathlib import Path +from typing import Any + + +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 = hallucination = omission = valid = 0 + + for qa in qa_records: + result_type = qa.get("result_type", "") + if result_type == "Correct": + correct += 1 + valid += 1 + elif result_type == "Hallucination": + hallucination += 1 + valid += 1 + elif result_type == "Omission": + omission += 1 + valid += 1 + + metrics = { + "correct_qa_ratio(all)": correct / total, + "hallucination_qa_ratio(all)": hallucination / total, + "omission_qa_ratio(all)": omission / total, + "correct_qa_ratio(valid)": correct / valid if valid > 0 else 0, + "hallucination_qa_ratio(valid)": hallucination / valid if valid > 0 else 0, + "omission_qa_ratio(valid)": omission / valid if valid > 0 else 0, + "qa_valid_num": valid, + "qa_num": total + } + + return metrics + + +def compute_time_metrics(users_data: list[dict]) -> dict[str, float]: + """Compute timing metrics from evaluation results.""" + add_duration = search_duration = 0 + + for user_data in users_data: + for session in user_data.get("sessions", []): + add_duration += session.get("add_dialogue_duration_ms", 0) + eval_results = session.get("session", {}).get("evaluation_results", {}) + for qa in eval_results.get("question_answering_records", []): + search_duration += qa.get("search_duration_ms", 0) + + 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 + } + + +def load_from_tmp_dir(tmp_dir: str) -> list[dict]: + """Load data from tmp directory.""" + tmp_path = Path(tmp_dir) + + # Try flat file structure first (conversation_{user}_session_{idx}.json) + json_files = [f for f in tmp_path.iterdir() if f.is_file() and f.suffix == ".json"] + + if json_files: + # Group files by user + users_dict = defaultdict(list) + + for json_file in json_files: + with open(json_file, "r", encoding="utf-8") as f: + session_data = json.load(f) + user_name = session_data.get("user_name") + if user_name: + users_dict[user_name].append(session_data) + + # Sort sessions by session_idx for each user + users_data = [] + for user_name, sessions in users_dict.items(): + sessions_sorted = sorted(sessions, key=lambda s: s.get("session_idx", 0)) + if sessions_sorted: + user_data = { + "uuid": sessions_sorted[0].get("uuid"), + "user_name": user_name, + "sessions": [] + } + for session_data in sessions_sorted: + session_copy = session_data.copy() + session_copy.pop("uuid", None) + session_copy.pop("user_name", None) + user_data["sessions"].append(session_copy) + users_data.append(user_data) + + return users_data + + # Fallback to directory structure (user_name/session_{idx}.json) + user_dirs = [d for d in tmp_path.iterdir() if d.is_dir()] + + users_data = [] + for user_dir in user_dirs: + session_files = sorted( + [f for f in user_dir.iterdir() if "session_" in f.name and f.suffix == ".json"], + key=lambda f: int(f.stem.split("_")[-1]) + ) + + if not session_files: + continue + + with open(session_files[0], "r", encoding="utf-8") as f: + first_session = json.load(f) + + user_data = { + "uuid": first_session["uuid"], + "user_name": first_session["user_name"], + "sessions": [] + } + + for session_file in session_files: + with open(session_file, "r", encoding="utf-8") as f: + session_data = json.load(f) + session_data.pop("uuid", None) + session_data.pop("user_name", None) + user_data["sessions"].append(session_data) + + users_data.append(user_data) + + return users_data + + +def main(tmp_dir: str): + """Main function to compute statistics from tmp directory.""" + tmp_path = Path(tmp_dir) + + if not tmp_path.exists() or not tmp_path.is_dir(): + print(f"❌ Error: Directory not found: {tmp_dir}") + return + + # Load data from tmp directory + users_data = load_from_tmp_dir(tmp_dir) + + # Collect QA records with metadata + qa_records = [] + qa_with_metadata = [] + user_count = session_count = 0 + + for user_data in users_data: + user_count += 1 + user_name = user_data.get("user_name", "Unknown") + + valid_session_idx = 0 + for session in user_data.get("sessions", []): + if session.get("is_generated_qa_session"): + continue + + session_count += 1 + eval_results = session.get("session", {}).get("evaluation_results", {}) + + for qa_idx, qa in enumerate(eval_results.get("question_answering_records", [])): + qa_records.append(qa) + qa_with_metadata.append({ + "user_name": user_name, + "session_idx": valid_session_idx, + "question_idx": qa_idx, + "qa_record": qa + }) + + valid_session_idx += 1 + + # Compute metrics + qa_metrics = compute_qa_metrics(qa_records) + time_metrics = compute_time_metrics(users_data) + + # Save results + output_dir = tmp_path.parent + report_file = output_dir / "reme_eval_stat_result.json" + + final_results = { + "overall_score": { + "question_answering": qa_metrics, + "time_consuming": time_metrics + }, + "question_answering_records": qa_records + } + + with open(report_file, "w", encoding="utf-8") as f: + json.dump(final_results, f, ensure_ascii=False, indent=4) + + # Print summary + print(f"\n📊 Data: {user_count} users, {session_count} sessions, {len(qa_records)} QA records") + print(f"\n✅ Metrics:") + print(f" Correct: {qa_metrics['correct_qa_ratio(all)']:.4f} | Hallucination: {qa_metrics['hallucination_qa_ratio(all)']:.4f} | Omission: {qa_metrics['omission_qa_ratio(all)']:.4f}") + print(f" Valid: {qa_metrics['qa_valid_num']}/{qa_metrics['qa_num']}") + print(f"\n⏱️ Time: {time_metrics['total_duration_time']:.2f} min (Add: {time_metrics['add_dialogue_duration_time']:.2f} | Search: {time_metrics['search_memory_duration_time']:.2f})") + print(f"\n💾 Results saved: {report_file}") + + # Print error records + print(f"\n{'='*80}\n❌ ERROR RECORDS ({len([r for r in qa_with_metadata if r['qa_record'].get('result_type') not in ['Correct', '']])} errors)\n{'='*80}") + + error_records = [r for r in qa_with_metadata if r["qa_record"].get("result_type") not in ["Correct", ""]] + + if error_records: + for idx, record in enumerate(error_records, 1): + qa = record["qa_record"] + print(f"\n[{idx}] {qa.get('result_type', 'Unknown')} - {record['user_name']} (S{record['session_idx']}/Q{record['question_idx']})") + print(f" Q: {qa.get('question', 'N/A')}") + print(f" Expected: {qa.get('answer', 'N/A')}") + print(f" Got: {qa.get('system_response', 'N/A')}") + + print() + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser(description="Compute QA statistics from tmp directory") + parser.add_argument( + "tmp_dir", + nargs='?', + default="./data", + type=str, + help="Path to tmp directory containing user session data (default: ./data)") + + args = parser.parse_args() + main(tmp_dir=args.tmp_dir) diff --git a/bench/human_in_the_loop/eval.yaml b/bench/human_in_the_loop/eval.yaml new file mode 100644 index 00000000..3e685bc7 --- /dev/null +++ b/bench/human_in_the_loop/eval.yaml @@ -0,0 +1,145 @@ +EVALUATION_PROMPT_FOR_QUESTION: | + You are an **evaluation expert for AI memory system question answering**. + Based **only** on the provided **“Question”**, **“Reference Answer”**, and **“Key Memory Points”** (the essential facts needed to derive the reference answer), strictly evaluate the **accuracy** of the **“Memory System Response.”** Classify it as one of **“Correct”**, **“Hallucination”**, or **“Omission.”** Do **not** use any external knowledge or subjective inference. Finally, output your judgment **strictly** in the specified JSON format. + + # Evaluation Criteria + + ## Answer Type Classification + + ### 1. Correct + + * The “Memory System Response” accurately answers the “Question,” and its content is **semantically equivalent** to the “Reference Answer.” + * It contains **no contradictions** with the “Key Memory Points” or “Reference Answer.” + * It introduces **no unsupported details** beyond the “Key Memory Points” that could alter the conclusion. + * Synonyms, paraphrasing, and reasonable summarization are acceptable. + + ### 2. Hallucination + + * The “Memory System Response” includes information or facts that **contradict or are inconsistent** with the “Reference Answer” or the “Key Memory Points.” + * When the “Reference Answer” is labeled as *unknown/uncertain*, yet the response provides a specific verifiable fact or conclusion. + * Extra irrelevant information that does **not change** the conclusion is **not** considered hallucination by itself; however, if it **changes or misleads** the conclusion, or **contradicts** the “Key Memory Points,” it should be judged as a **Hallucination**. + + ### 3. Omission + + * The response is **incomplete** compared to the “Reference Answer.” + * It explicitly states “don’t know,” “can’t remember,” or “no related memory,” even though relevant information exists in the “Key Memory Points.” + * For multi-element questions, **all elements must be correct and present**; omission of **any** element is considered an **Omission**. + + ## Priority Rules (Conflict Handling) + + * If the response contains **both missing necessary information** and **fabricated/contradictory information**, classify it as **Hallucination**. + * If there is **no fabrication/contradiction** but some necessary information is missing, classify it as **Omission**. + * Only when the meaning is **fully equivalent** to the reference answer should it be classified as **Correct**. + + ## Detailed Guidelines and Tolerance + + * Equivalent expressions of numbers, times, and units are acceptable, but the **numerical values themselves must not differ**. + * For multi-element questions, **all elements must be complete and accurate**; missing any element counts as **Omission**. + * If the reference answer is *“unknown / cannot be determined”* and the system provides a definite fact, that is a **Hallucination**. + If the system also answers *“unknown”* (without guessing), it may be **Correct**. + * The evaluation must rely **only** on the *Reference Answer*, *Key Memory Points*, and *System Response* — no external context, world knowledge, or speculative reasoning is allowed. + + # Information for Evaluation + + * **Question:** + {question} + + * **Reference Answer:** + {reference_answer} + + * **Key Memory Points:** + {key_memory_points} + + * **Memory System Response:** + {response} + + # Output Requirements + + Please provide your evaluation result **strictly** in the JSON format below. + Do **not** add any extra explanation or comments outside the JSON block. + + ```json + {{ + "reasoning": "Provide a concise and traceable evaluation rationale: first compare the system’s response with the Key Memory Points (which were correctly used, which were missing, and whether there was any fabrication/contradiction), then assess its consistency with the Reference Answer, and finally state the classification basis.", + "evaluation_result": "Correct | Hallucination | Omission" + }} + ``` + + +EVALUATION_PROMPT_FOR_QUESTION2: | + You are an **evaluation expert for AI memory system question answering**. + + Based **only** on the provided **"Question"**, **"Reference Answer"**, and **"Key Memory Points"** (the essential facts needed to derive the reference answer), strictly evaluate the **accuracy** of the **"Memory System Response."** Classify it as one of **"Correct"**, **"Hallucination"**, or **"Omission."** Do **not** use any external knowledge or subjective inference. Finally, output your judgment **strictly** in the specified JSON format. + + # Evaluation Criteria + + ## Answer Type Classification + + ### 1. Correct + + * The "Memory System Response" accurately answers the "Question," and its content is **semantically equivalent** to the "Reference Answer." + * It contains **no contradictions** with the "Key Memory Points" or "Reference Answer." + * **Extra details not present in the Key Memory Points are allowed and should not be penalized**, as long as they: + - Do not contradict the Key Memory Points or Reference Answer + - Do not change or mislead the core conclusion + - Are reasonable additional context that the memory system may have retained from the conversation + * The memory system may have stored additional information beyond the Key Memory Points. Such extra information should be treated as **supplementary context** rather than hallucination, provided it does not conflict with the core answer. + * Synonyms, paraphrasing, and reasonable summarization are acceptable. + + ### 2. Hallucination + + * The "Memory System Response" includes information or facts that **contradict or are inconsistent** with the "Reference Answer" or the "Key Memory Points." + * The response provides information that **directly contradicts** known facts from the Key Memory Points. + * When the "Reference Answer" is labeled as *unknown/uncertain*, yet the response provides a specific verifiable fact or conclusion. + * **Important:** Extra information that is NOT in Key Memory Points is **NOT automatically a hallucination**. Only classify as hallucination if the extra information: + - Directly contradicts the Key Memory Points or Reference Answer + - Changes or misleads the core conclusion in a way that makes the answer incorrect + - Provides a definitive answer when the Reference Answer indicates uncertainty + + ### 3. Omission + + * The response is **incomplete** compared to the "Reference Answer." + * It explicitly states "don't know," "can't remember," or "no related memory," even though relevant information exists in the "Key Memory Points." + * For multi-element questions, **all elements must be correct and present**; omission of **any** element is considered an **Omission**. + + ## Priority Rules (Conflict Handling) + + * If the response contains **both missing necessary information** and **fabricated/contradictory information**, classify it as **Hallucination**. + * If there is **no fabrication/contradiction** but some necessary information is missing, classify it as **Omission**. + * If the core answer is correct and complete, classify as **Correct** even if there are extra details not in Key Memory Points (as long as they don't contradict or mislead). + + ## Detailed Guidelines and Tolerance + + * Equivalent expressions of numbers, times, and units are acceptable, but the **numerical values themselves must not differ**. + * For multi-element questions, **all elements must be complete and accurate**; missing any element counts as **Omission**. + * If the reference answer is *"unknown / cannot be determined"* and the system provides a definite fact, that is a **Hallucination**. + If the system also answers *"unknown"* (without guessing), it may be **Correct**. + * **Focus on evaluating whether the core answer to the question is correct**, not whether the response is limited to only the Key Memory Points. + * Extra contextual information (e.g., additional preferences, related details) should be viewed as enrichment, not as errors, unless they contradict or mislead. + + # Information for Evaluation + + * **Question:** + {question} + + * **Reference Answer:** + {reference_answer} + + * **Key Memory Points:** + {key_memory_points} + + * **Memory System Response:** + {response} + + # Output Requirements + + Please provide your evaluation result **strictly** in the JSON format below. + Do **not** add any extra explanation or comments outside the JSON block. + + ```json + {{ + "reasoning": "Provide a concise and traceable evaluation rationale: first verify that the system's response correctly includes all required elements from the Reference Answer, then check if any information contradicts the Key Memory Points or Reference Answer. Extra details not in Key Memory Points should be noted but not penalized unless they contradict or mislead. Finally state the classification basis.", + "evaluation_result": "Correct | Hallucination | Omission" + }} + ``` + """ \ No newline at end of file diff --git a/bench/human_in_the_loop/reevaluate_qa.py b/bench/human_in_the_loop/reevaluate_qa.py new file mode 100644 index 00000000..6b9cac8f --- /dev/null +++ b/bench/human_in_the_loop/reevaluate_qa.py @@ -0,0 +1,550 @@ +""" +Re-evaluate Question Answering results from data directory using LLM. + +This script: +1. Loads existing QA records from data directory +2. Re-evaluates each system_response using multiple models in parallel +3. Uses both EVALUATION_PROMPT_FOR_QUESTION and EVALUATION_PROMPT_FOR_QUESTION2 +4. Saves updated results with new evaluation metrics +""" + +import asyncio +import json +import re +import yaml +from collections import defaultdict +from pathlib import Path +from typing import Any + +from reme_ai.core.schema import Message +from reme_ai.core.utils import load_env +from reme_ai.reme import ReMe +from tenacity import retry, stop_after_attempt, wait_random_exponential + +# Load environment +load_env() + +# Initialize ReMe singleton +reme = ReMe() + +# Load prompts from YAML file +_YAML_PATH = Path(__file__).parent / "eval.yaml" +with open(_YAML_PATH, "r", encoding="utf-8") as f: + _PROMPTS = yaml.safe_load(f) + + +@retry( + wait=wait_random_exponential(min=1, max=60), + stop=stop_after_attempt(3), + reraise=True, +) +async def llm_request(prompt: str, model_name: str = "qwen3-max", **kwargs) -> str: + """Make an LLM request using ReMe's LLM.""" + assistant_message = await reme.llm.chat( + messages=[ + Message( + **{ + "role": "user", + "content": prompt, + }, + ), + ], + model_name=model_name, + **kwargs, + ) + return assistant_message.content + + +@retry( + wait=wait_random_exponential(min=1, max=60), + stop=stop_after_attempt(5), + reraise=True, +) +async def llm_request_for_json(prompt: str, model_name: str = "qwen3-max", **kwargs) -> dict: + """Make an LLM request expecting JSON response.""" + content = await llm_request(prompt, model_name=model_name, **kwargs) + + match = re.search(r"```json\s*(\{.*?\})\s*```", content, re.DOTALL) + if not match: + raise ValueError(f"No JSON block found in model output: {content}") + + json_str = match.group(1).strip() + return json.loads(json_str) + + +async def evaluate_qa_record( + question: str, + reference_answer: str, + key_memory_points: str, + response: str, + dialogue: str = "", + model_name: str = "qwen3-max", + prompt_version: str = "v1" +) -> dict: + """Evaluate a single QA record using LLM with specified prompt version. + + Args: + question: The question to evaluate + reference_answer: The reference answer + key_memory_points: Key memory points + response: System response to evaluate + dialogue: Dialogue context (optional) + model_name: LLM model name + prompt_version: "v1" for EVALUATION_PROMPT_FOR_QUESTION, + "v2" for EVALUATION_PROMPT_FOR_QUESTION2 + + Returns: + dict with evaluation_result and reasoning + """ + # Select prompt template + if prompt_version == "v2": + prompt_template = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION2"] + else: + prompt_template = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION"] + + # Format prompt + prompt = prompt_template.format( + question=question, + reference_answer=reference_answer, + key_memory_points=key_memory_points, + response=response, + dialogue=dialogue or "N/A" + ) + + result = await llm_request_for_json(prompt, model_name=model_name) + return result + + +def load_from_data_dir(data_dir: str) -> list[dict]: + """Load data from data directory (same as compute_qa_stats.py).""" + data_path = Path(data_dir) + + # Try flat file structure first (conversation_{user}_session_{idx}.json) + json_files = [f for f in data_path.iterdir() if f.is_file() and f.suffix == ".json"] + + if json_files: + # Group files by user + users_dict = defaultdict(list) + + for json_file in json_files: + with open(json_file, "r", encoding="utf-8") as f: + session_data = json.load(f) + user_name = session_data.get("user_name") + if user_name: + users_dict[user_name].append({ + "file": json_file, + "data": session_data + }) + + # Sort sessions by session_idx for each user + users_data = [] + for user_name, sessions in users_dict.items(): + sessions_sorted = sorted( + sessions, + key=lambda s: s["data"].get("session_idx", 0) + ) + users_data.extend(sessions_sorted) + + return users_data + + return [] + + +def format_dialogue_context(session_data: dict) -> str: + """Format dialogue context from session data.""" + dialogue = session_data.get("session", {}).get("dialogue", []) + if not dialogue: + return "N/A" + + formatted_turns = [] + for turn in dialogue: + role = turn.get("role", "unknown") + content = turn.get("content", "") + timestamp = turn.get("timestamp", "") + formatted_turns.append( + f"Role: {role}\nContent: {content}\nTime: {timestamp}" + ) + return "\n\n".join(formatted_turns) + + +async def reevaluate_session( + session_file: Path, + session_data: dict, + models: list[str], + prompt_versions: list[str], + parallel: bool = True +) -> dict: + """Re-evaluate all QA records in a session using multiple models and prompts. + + Args: + session_file: Path to session file + session_data: Session data dict + models: List of model names to use for evaluation + prompt_versions: List of prompt versions ("v1", "v2") + parallel: If True, use asyncio.gather for parallel execution; + if False, execute sequentially + + Returns: + Updated session data with evaluation results for each model+prompt combination + + Note: + Request rate limiting is handled by base_llm.py's request_interval mechanism. + """ + eval_results = session_data.get("session", {}).get("evaluation_results", {}) + qa_records = eval_results.get("question_answering_records", []) + + if not qa_records: + print(f" ⏭️ No QA records found") + return session_data + + total_evals = len(models) * len(prompt_versions) * len(qa_records) + print(f" 🔍 Re-evaluating {len(qa_records)} QA records with {len(models)} models × {len(prompt_versions)} prompts = {total_evals} evaluations...") + + # Format dialogue context once + dialogue_context = format_dialogue_context(session_data) + + async def evaluate_single_combination( + idx: int, + qa: dict, + model_name: str, + prompt_version: str + ) -> tuple[int, str, str, dict]: + """Evaluate a single QA record with specific model and prompt. + + Note: Rate limiting is handled by BaseLLM's request_interval mechanism. + """ + question = qa.get("question", "") + reference_answer = qa.get("answer", "") + + # Get key memory points from evidence + evidence = qa.get("evidence", []) + key_memory_points = "\n".join([e.get("memory_content", "") for e in evidence]) + + # Get system response + system_response = qa.get("system_response", "") + + try: + # Call LLM for evaluation + eval_result = await evaluate_qa_record( + question=question, + reference_answer=reference_answer, + key_memory_points=key_memory_points, + response=system_response, + dialogue=dialogue_context, + model_name=model_name, + prompt_version=prompt_version + ) + + result = { + "result_type": eval_result.get("evaluation_result", "Invalid"), + "reasoning": eval_result.get("reasoning", "") + } + + return idx, model_name, prompt_version, result + + except Exception as e: + print(f" ❌ QA[{idx+1}] {model_name}/{prompt_version}: Error: {e}") + return idx, model_name, prompt_version, { + "result_type": "Error", + "reasoning": f"Evaluation error: {str(e)}" + } + + # Create all evaluation tasks (all combinations of models, prompts, and QA records) + tasks = [] + for idx, qa in enumerate(qa_records): + for model_name in models: + for prompt_version in prompt_versions: + tasks.append(evaluate_single_combination(idx, qa, model_name, prompt_version)) + + # Execute evaluations based on parallel mode + if parallel: + print(f" ⚡ Starting {len(tasks)} parallel evaluations (rate limited by LLM layer)...") + results = await asyncio.gather(*tasks) + else: + print(f" 🔄 Starting {len(tasks)} sequential evaluations...") + results = [] + for i, task in enumerate(tasks, 1): + result = await task + results.append(result) + if i % 10 == 0 or i == len(tasks): + print(f" ⏳ Progress: {i}/{len(tasks)} evaluations completed") + + # Organize results by QA index, then by model and prompt + # Structure: qa_records[idx]["evaluations"][model][prompt_version] = {result_type, reasoning} + for idx, qa in enumerate(qa_records): + if "evaluations" not in qa: + qa["evaluations"] = {} + + # Initialize evaluations structure + for model_name in models: + if model_name not in qa["evaluations"]: + qa["evaluations"][model_name] = {} + + # Fill in results + completed_count = 0 + for qa_idx, model_name, prompt_version, result in results: + qa_records[qa_idx]["evaluations"][model_name][prompt_version] = result + completed_count += 1 + if completed_count % 10 == 0 or completed_count == len(results): + print(f" ✅ Completed {completed_count}/{len(results)} evaluations") + + # Set default result_type to first model's v1 result for compatibility + if models and prompt_versions: + default_model = models[0] + default_prompt = prompt_versions[0] + for qa in qa_records: + default_eval = qa["evaluations"].get(default_model, {}).get(default_prompt, {}) + qa["result_type"] = default_eval.get("result_type", "Invalid") + qa["question_answering_reasoning"] = default_eval.get("reasoning", "") + + # Update session data + if "session" not in session_data: + session_data["session"] = {} + if "evaluation_results" not in session_data["session"]: + session_data["session"]["evaluation_results"] = {} + + session_data["session"]["evaluation_results"]["question_answering_records"] = qa_records + + # Save updated session data + with open(session_file, "w", encoding="utf-8") as f: + json.dump(session_data, f, ensure_ascii=False, indent=2) + + print(f" 💾 Updated session saved with all evaluations") + + return session_data + + +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 = hallucination = omission = valid = 0 + + for qa in qa_records: + result_type = qa.get("result_type", "") + if result_type == "Correct": + correct += 1 + valid += 1 + elif result_type == "Hallucination": + hallucination += 1 + valid += 1 + elif result_type == "Omission": + omission += 1 + valid += 1 + + metrics = { + "correct_qa_ratio(all)": correct / total, + "hallucination_qa_ratio(all)": hallucination / total, + "omission_qa_ratio(all)": omission / total, + "correct_qa_ratio(valid)": correct / valid if valid > 0 else 0, + "hallucination_qa_ratio(valid)": hallucination / valid if valid > 0 else 0, + "omission_qa_ratio(valid)": omission / valid if valid > 0 else 0, + "qa_valid_num": valid, + "qa_num": total + } + + return metrics + + +async def main( + data_dir: str = "./data", + models: list[str] = None, + prompt_versions: list[str] = None, + parallel: bool = True +): + """Main function to re-evaluate QA records from data directory with multiple models and prompts. + + Args: + data_dir: Path to data directory + models: List of model names (e.g., ["qwen3-max", "qwen-flash"]) + prompt_versions: List of prompt versions (e.g., ["v1", "v2"]) + parallel: If True, use parallel execution; if False, use sequential execution + + Note: + Request rate limiting is automatically handled by base_llm.py's request_interval mechanism. + """ + data_path = Path(data_dir) + + if not data_path.exists() or not data_path.is_dir(): + print(f"❌ Error: Directory not found: {data_dir}") + return + + # Default values + if models is None: + models = ["qwen3-max"] + if prompt_versions is None: + prompt_versions = ["v1"] + + print("=" * 80) + print("RE-EVALUATING QA RECORDS WITH MULTIPLE MODELS & PROMPTS") + print(f"Models: {', '.join(models)}") + print(f"Prompts: {', '.join(['EVALUATION_PROMPT_FOR_QUESTION' if v == 'v1' else 'EVALUATION_PROMPT_FOR_QUESTION2' for v in prompt_versions])}") + print(f"Execution mode: {'Parallel' if parallel else 'Sequential'}") + print("Note: Request rate limiting handled by LLM layer (base_llm.py)") + print("=" * 80 + "\n") + + # Load data from directory + sessions = load_from_data_dir(data_dir) + + if not sessions: + print(f"❌ No session files found in {data_dir}") + return + + print(f"📂 Found {len(sessions)} session files\n") + + # Process each session + all_qa_records = [] + + for idx, session_info in enumerate(sessions, 1): + session_file = session_info["file"] + session_data = session_info["data"] + user_name = session_data.get("user_name", "Unknown") + session_idx = session_data.get("session_idx", 0) + + print(f"[{idx}/{len(sessions)}] {user_name} - Session {session_idx}") + + updated_session = await reevaluate_session( + session_file=session_file, + session_data=session_data, + models=models, + prompt_versions=prompt_versions, + parallel=parallel + ) + + # Collect QA records for metrics + eval_results = updated_session.get("session", {}).get("evaluation_results", {}) + qa_records = eval_results.get("question_answering_records", []) + all_qa_records.extend(qa_records) + + print() + + # Compute and display metrics for each model+prompt combination + print("=" * 80) + print("UPDATED METRICS (BY MODEL & PROMPT)") + print("=" * 80 + "\n") + + for model_name in models: + for prompt_version in prompt_versions: + prompt_name = "EVALUATION_PROMPT_FOR_QUESTION" if prompt_version == "v1" else "EVALUATION_PROMPT_FOR_QUESTION2" + print(f"\n📊 {model_name} / {prompt_name}:") + print("─" * 80) + + # Extract QA records for this model+prompt combination + model_qa_records = [] + for qa in all_qa_records: + eval_data = qa.get("evaluations", {}).get(model_name, {}).get(prompt_version, {}) + if eval_data: + # Create a copy with the specific evaluation result + qa_copy = { + **qa, + "result_type": eval_data.get("result_type", "Invalid"), + "question_answering_reasoning": eval_data.get("reasoning", "") + } + model_qa_records.append(qa_copy) + + if model_qa_records: + metrics = compute_qa_metrics(model_qa_records) + + print(f" Correct (all): {metrics['correct_qa_ratio(all)']:.4f}") + print(f" Hallucination (all): {metrics['hallucination_qa_ratio(all)']:.4f}") + print(f" Omission (all): {metrics['omission_qa_ratio(all)']:.4f}") + print(f" Correct (valid): {metrics['correct_qa_ratio(valid)']:.4f}") + print(f" Hallucination (valid): {metrics['hallucination_qa_ratio(valid)']:.4f}") + print(f" Omission (valid): {metrics['omission_qa_ratio(valid)']:.4f}") + print(f" Valid/Total: {metrics['qa_valid_num']}/{metrics['qa_num']}") + + # Save detailed results with all evaluations + report_file = data_path.parent / "reme_eval_stat_result_detailed.json" + + # Create summary for each model+prompt combination + evaluation_summary = {} + for model_name in models: + evaluation_summary[model_name] = {} + for prompt_version in prompt_versions: + prompt_name = "EVALUATION_PROMPT_FOR_QUESTION" if prompt_version == "v1" else "EVALUATION_PROMPT_FOR_QUESTION2" + + # Extract QA records for this combination + model_qa_records = [] + for qa in all_qa_records: + eval_data = qa.get("evaluations", {}).get(model_name, {}).get(prompt_version, {}) + if eval_data: + qa_copy = { + **qa, + "result_type": eval_data.get("result_type", "Invalid"), + "question_answering_reasoning": eval_data.get("reasoning", "") + } + model_qa_records.append(qa_copy) + + metrics = compute_qa_metrics(model_qa_records) + evaluation_summary[model_name][prompt_name] = { + "metrics": metrics, + "qa_records": model_qa_records + } + + final_results = { + "evaluation_summary": evaluation_summary, + "all_qa_records_with_evaluations": all_qa_records + } + + with open(report_file, "w", encoding="utf-8") as f: + json.dump(final_results, f, ensure_ascii=False, indent=2) + + print(f"\n💾 Detailed results saved: {report_file}") + print("\n" + "=" * 80) + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser( + description="Re-evaluate QA records from data directory using multiple models and prompts in parallel. " + "Request rate limiting is automatically handled by base_llm.py's request_interval mechanism." + ) + parser.add_argument( + "data_dir", + nargs='?', + default="./data", + type=str, + help="Path to data directory containing user session files (default: ./data)" + ) + parser.add_argument( + "--models", + type=str, + nargs='+', + default=["gpt-5.1-2025-11-13", "gemini-3-pro-preview"], + help="LLM model names for evaluation (space-separated, default: qwen3-max)" + ) + # ["gpt-4o-2024-11-20"], ["qwen3-max", "qwen-flash", "qwen3-30b-a3b-instruct-2507", "qwen3-30b-a3b-instruct-2507"] + parser.add_argument( + "--prompts", + type=str, + nargs='+', + choices=["v1", "v2"], + default=["v1", "v2"], + help="Prompt versions to use: v1=EVALUATION_PROMPT_FOR_QUESTION, v2=EVALUATION_PROMPT_FOR_QUESTION2 (default: v1)" + ) + parser.add_argument( + "--serial", + action="store_true", + help="Use sequential execution instead of parallel (default: parallel)" + ) + + args = parser.parse_args() + + asyncio.run(main( + data_dir=args.data_dir, + models=args.models, + prompt_versions=args.prompts, + parallel=not args.serial + )) diff --git a/reme_ai/mem_tool/identity/__init__.py b/bench/human_in_the_loop2/__init__.py similarity index 100% rename from reme_ai/mem_tool/identity/__init__.py rename to bench/human_in_the_loop2/__init__.py diff --git a/bench/human_in_the_loop2/compute_qa_stats.py b/bench/human_in_the_loop2/compute_qa_stats.py new file mode 100644 index 00000000..69d65e55 --- /dev/null +++ b/bench/human_in_the_loop2/compute_qa_stats.py @@ -0,0 +1,237 @@ +""" +Compute Question Answering statistics from evaluation results in tmp directory. +""" + +import json +from collections import defaultdict +from pathlib import Path +from typing import Any + + +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 = hallucination = omission = valid = 0 + + for qa in qa_records: + result_type = qa.get("result_type", "") + if result_type == "Correct": + correct += 1 + valid += 1 + elif result_type == "Hallucination": + hallucination += 1 + valid += 1 + elif result_type == "Omission": + omission += 1 + valid += 1 + + metrics = { + "correct_qa_ratio(all)": correct / total, + "hallucination_qa_ratio(all)": hallucination / total, + "omission_qa_ratio(all)": omission / total, + "correct_qa_ratio(valid)": correct / valid if valid > 0 else 0, + "hallucination_qa_ratio(valid)": hallucination / valid if valid > 0 else 0, + "omission_qa_ratio(valid)": omission / valid if valid > 0 else 0, + "qa_valid_num": valid, + "qa_num": total + } + + return metrics + + +def compute_time_metrics(users_data: list[dict]) -> dict[str, float]: + """Compute timing metrics from evaluation results.""" + add_duration = search_duration = 0 + + for user_data in users_data: + for session in user_data.get("sessions", []): + add_duration += session.get("add_dialogue_duration_ms", 0) + eval_results = session.get("session", {}).get("evaluation_results", {}) + for qa in eval_results.get("question_answering_records", []): + search_duration += qa.get("search_duration_ms", 0) + + 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 + } + + +def load_from_tmp_dir(tmp_dir: str) -> list[dict]: + """Load data from tmp directory.""" + tmp_path = Path(tmp_dir) + + # Try flat file structure first (conversation_{user}_session_{idx}.json) + json_files = [f for f in tmp_path.iterdir() if f.is_file() and f.suffix == ".json"] + + if json_files: + # Group files by user + users_dict = defaultdict(list) + + for json_file in json_files: + with open(json_file, "r", encoding="utf-8") as f: + session_data = json.load(f) + user_name = session_data.get("user_name") + if user_name: + users_dict[user_name].append(session_data) + + # Sort sessions by session_idx for each user + users_data = [] + for user_name, sessions in users_dict.items(): + sessions_sorted = sorted(sessions, key=lambda s: s.get("session_idx", 0)) + if sessions_sorted: + user_data = { + "uuid": sessions_sorted[0].get("uuid"), + "user_name": user_name, + "sessions": [] + } + for session_data in sessions_sorted: + session_copy = session_data.copy() + session_copy.pop("uuid", None) + session_copy.pop("user_name", None) + user_data["sessions"].append(session_copy) + users_data.append(user_data) + + return users_data + + # Fallback to directory structure (user_name/session_{idx}.json) + user_dirs = [d for d in tmp_path.iterdir() if d.is_dir()] + + users_data = [] + for user_dir in user_dirs: + session_files = sorted( + [f for f in user_dir.iterdir() if "session_" in f.name and f.suffix == ".json"], + key=lambda f: int(f.stem.split("_")[-1]) + ) + + if not session_files: + continue + + with open(session_files[0], "r", encoding="utf-8") as f: + first_session = json.load(f) + + user_data = { + "uuid": first_session["uuid"], + "user_name": first_session["user_name"], + "sessions": [] + } + + for session_file in session_files: + with open(session_file, "r", encoding="utf-8") as f: + session_data = json.load(f) + session_data.pop("uuid", None) + session_data.pop("user_name", None) + user_data["sessions"].append(session_data) + + users_data.append(user_data) + + return users_data + + +def main(tmp_dir: str): + """Main function to compute statistics from tmp directory.""" + tmp_path = Path(tmp_dir) + + if not tmp_path.exists() or not tmp_path.is_dir(): + print(f"❌ Error: Directory not found: {tmp_dir}") + return + + # Load data from tmp directory + users_data = load_from_tmp_dir(tmp_dir) + + # Collect QA records with metadata + qa_records = [] + qa_with_metadata = [] + user_count = session_count = 0 + + for user_data in users_data: + user_count += 1 + user_name = user_data.get("user_name", "Unknown") + + valid_session_idx = 0 + for session in user_data.get("sessions", []): + if session.get("is_generated_qa_session"): + continue + + session_count += 1 + eval_results = session.get("session", {}).get("evaluation_results", {}) + + for qa_idx, qa in enumerate(eval_results.get("question_answering_records", [])): + qa_records.append(qa) + qa_with_metadata.append({ + "user_name": user_name, + "session_idx": valid_session_idx, + "question_idx": qa_idx, + "qa_record": qa + }) + + valid_session_idx += 1 + + # Compute metrics + qa_metrics = compute_qa_metrics(qa_records) + time_metrics = compute_time_metrics(users_data) + + # Save results + output_dir = tmp_path.parent + report_file = output_dir / "reme_eval_stat_result.json" + + final_results = { + "overall_score": { + "question_answering": qa_metrics, + "time_consuming": time_metrics + }, + "question_answering_records": qa_records + } + + with open(report_file, "w", encoding="utf-8") as f: + json.dump(final_results, f, ensure_ascii=False, indent=4) + + # Print summary + print(f"\n📊 Data: {user_count} users, {session_count} sessions, {len(qa_records)} QA records") + print(f"\n✅ Metrics:") + print(f" Correct: {qa_metrics['correct_qa_ratio(all)']:.4f} | Hallucination: {qa_metrics['hallucination_qa_ratio(all)']:.4f} | Omission: {qa_metrics['omission_qa_ratio(all)']:.4f}") + print(f" Valid: {qa_metrics['qa_valid_num']}/{qa_metrics['qa_num']}") + print(f"\n⏱️ Time: {time_metrics['total_duration_time']:.2f} min (Add: {time_metrics['add_dialogue_duration_time']:.2f} | Search: {time_metrics['search_memory_duration_time']:.2f})") + print(f"\n💾 Results saved: {report_file}") + + # Print error records + print(f"\n{'='*80}\n❌ ERROR RECORDS ({len([r for r in qa_with_metadata if r['qa_record'].get('result_type') not in ['Correct', '']])} errors)\n{'='*80}") + + error_records = [r for r in qa_with_metadata if r["qa_record"].get("result_type") not in ["Correct", ""]] + + if error_records: + for idx, record in enumerate(error_records, 1): + qa = record["qa_record"] + print(f"\n[{idx}] {qa.get('result_type', 'Unknown')} - {record['user_name']} (S{record['session_idx']}/Q{record['question_idx']})") + print(f" Q: {qa.get('question', 'N/A')}") + print(f" Expected: {qa.get('answer', 'N/A')}") + print(f" Got: {qa.get('system_response', 'N/A')}") + + print() + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser(description="Compute QA statistics from tmp directory") + parser.add_argument( + "tmp_dir", + nargs='?', + default="./data", + type=str, + help="Path to tmp directory containing user session data (default: ./data)") + + args = parser.parse_args() + main(tmp_dir=args.tmp_dir) diff --git a/bench/human_in_the_loop2/eval.yaml b/bench/human_in_the_loop2/eval.yaml new file mode 100644 index 00000000..3e685bc7 --- /dev/null +++ b/bench/human_in_the_loop2/eval.yaml @@ -0,0 +1,145 @@ +EVALUATION_PROMPT_FOR_QUESTION: | + You are an **evaluation expert for AI memory system question answering**. + Based **only** on the provided **“Question”**, **“Reference Answer”**, and **“Key Memory Points”** (the essential facts needed to derive the reference answer), strictly evaluate the **accuracy** of the **“Memory System Response.”** Classify it as one of **“Correct”**, **“Hallucination”**, or **“Omission.”** Do **not** use any external knowledge or subjective inference. Finally, output your judgment **strictly** in the specified JSON format. + + # Evaluation Criteria + + ## Answer Type Classification + + ### 1. Correct + + * The “Memory System Response” accurately answers the “Question,” and its content is **semantically equivalent** to the “Reference Answer.” + * It contains **no contradictions** with the “Key Memory Points” or “Reference Answer.” + * It introduces **no unsupported details** beyond the “Key Memory Points” that could alter the conclusion. + * Synonyms, paraphrasing, and reasonable summarization are acceptable. + + ### 2. Hallucination + + * The “Memory System Response” includes information or facts that **contradict or are inconsistent** with the “Reference Answer” or the “Key Memory Points.” + * When the “Reference Answer” is labeled as *unknown/uncertain*, yet the response provides a specific verifiable fact or conclusion. + * Extra irrelevant information that does **not change** the conclusion is **not** considered hallucination by itself; however, if it **changes or misleads** the conclusion, or **contradicts** the “Key Memory Points,” it should be judged as a **Hallucination**. + + ### 3. Omission + + * The response is **incomplete** compared to the “Reference Answer.” + * It explicitly states “don’t know,” “can’t remember,” or “no related memory,” even though relevant information exists in the “Key Memory Points.” + * For multi-element questions, **all elements must be correct and present**; omission of **any** element is considered an **Omission**. + + ## Priority Rules (Conflict Handling) + + * If the response contains **both missing necessary information** and **fabricated/contradictory information**, classify it as **Hallucination**. + * If there is **no fabrication/contradiction** but some necessary information is missing, classify it as **Omission**. + * Only when the meaning is **fully equivalent** to the reference answer should it be classified as **Correct**. + + ## Detailed Guidelines and Tolerance + + * Equivalent expressions of numbers, times, and units are acceptable, but the **numerical values themselves must not differ**. + * For multi-element questions, **all elements must be complete and accurate**; missing any element counts as **Omission**. + * If the reference answer is *“unknown / cannot be determined”* and the system provides a definite fact, that is a **Hallucination**. + If the system also answers *“unknown”* (without guessing), it may be **Correct**. + * The evaluation must rely **only** on the *Reference Answer*, *Key Memory Points*, and *System Response* — no external context, world knowledge, or speculative reasoning is allowed. + + # Information for Evaluation + + * **Question:** + {question} + + * **Reference Answer:** + {reference_answer} + + * **Key Memory Points:** + {key_memory_points} + + * **Memory System Response:** + {response} + + # Output Requirements + + Please provide your evaluation result **strictly** in the JSON format below. + Do **not** add any extra explanation or comments outside the JSON block. + + ```json + {{ + "reasoning": "Provide a concise and traceable evaluation rationale: first compare the system’s response with the Key Memory Points (which were correctly used, which were missing, and whether there was any fabrication/contradiction), then assess its consistency with the Reference Answer, and finally state the classification basis.", + "evaluation_result": "Correct | Hallucination | Omission" + }} + ``` + + +EVALUATION_PROMPT_FOR_QUESTION2: | + You are an **evaluation expert for AI memory system question answering**. + + Based **only** on the provided **"Question"**, **"Reference Answer"**, and **"Key Memory Points"** (the essential facts needed to derive the reference answer), strictly evaluate the **accuracy** of the **"Memory System Response."** Classify it as one of **"Correct"**, **"Hallucination"**, or **"Omission."** Do **not** use any external knowledge or subjective inference. Finally, output your judgment **strictly** in the specified JSON format. + + # Evaluation Criteria + + ## Answer Type Classification + + ### 1. Correct + + * The "Memory System Response" accurately answers the "Question," and its content is **semantically equivalent** to the "Reference Answer." + * It contains **no contradictions** with the "Key Memory Points" or "Reference Answer." + * **Extra details not present in the Key Memory Points are allowed and should not be penalized**, as long as they: + - Do not contradict the Key Memory Points or Reference Answer + - Do not change or mislead the core conclusion + - Are reasonable additional context that the memory system may have retained from the conversation + * The memory system may have stored additional information beyond the Key Memory Points. Such extra information should be treated as **supplementary context** rather than hallucination, provided it does not conflict with the core answer. + * Synonyms, paraphrasing, and reasonable summarization are acceptable. + + ### 2. Hallucination + + * The "Memory System Response" includes information or facts that **contradict or are inconsistent** with the "Reference Answer" or the "Key Memory Points." + * The response provides information that **directly contradicts** known facts from the Key Memory Points. + * When the "Reference Answer" is labeled as *unknown/uncertain*, yet the response provides a specific verifiable fact or conclusion. + * **Important:** Extra information that is NOT in Key Memory Points is **NOT automatically a hallucination**. Only classify as hallucination if the extra information: + - Directly contradicts the Key Memory Points or Reference Answer + - Changes or misleads the core conclusion in a way that makes the answer incorrect + - Provides a definitive answer when the Reference Answer indicates uncertainty + + ### 3. Omission + + * The response is **incomplete** compared to the "Reference Answer." + * It explicitly states "don't know," "can't remember," or "no related memory," even though relevant information exists in the "Key Memory Points." + * For multi-element questions, **all elements must be correct and present**; omission of **any** element is considered an **Omission**. + + ## Priority Rules (Conflict Handling) + + * If the response contains **both missing necessary information** and **fabricated/contradictory information**, classify it as **Hallucination**. + * If there is **no fabrication/contradiction** but some necessary information is missing, classify it as **Omission**. + * If the core answer is correct and complete, classify as **Correct** even if there are extra details not in Key Memory Points (as long as they don't contradict or mislead). + + ## Detailed Guidelines and Tolerance + + * Equivalent expressions of numbers, times, and units are acceptable, but the **numerical values themselves must not differ**. + * For multi-element questions, **all elements must be complete and accurate**; missing any element counts as **Omission**. + * If the reference answer is *"unknown / cannot be determined"* and the system provides a definite fact, that is a **Hallucination**. + If the system also answers *"unknown"* (without guessing), it may be **Correct**. + * **Focus on evaluating whether the core answer to the question is correct**, not whether the response is limited to only the Key Memory Points. + * Extra contextual information (e.g., additional preferences, related details) should be viewed as enrichment, not as errors, unless they contradict or mislead. + + # Information for Evaluation + + * **Question:** + {question} + + * **Reference Answer:** + {reference_answer} + + * **Key Memory Points:** + {key_memory_points} + + * **Memory System Response:** + {response} + + # Output Requirements + + Please provide your evaluation result **strictly** in the JSON format below. + Do **not** add any extra explanation or comments outside the JSON block. + + ```json + {{ + "reasoning": "Provide a concise and traceable evaluation rationale: first verify that the system's response correctly includes all required elements from the Reference Answer, then check if any information contradicts the Key Memory Points or Reference Answer. Extra details not in Key Memory Points should be noted but not penalized unless they contradict or mislead. Finally state the classification basis.", + "evaluation_result": "Correct | Hallucination | Omission" + }} + ``` + """ \ No newline at end of file diff --git a/bench/human_in_the_loop2/reevaluate_qa.py b/bench/human_in_the_loop2/reevaluate_qa.py new file mode 100644 index 00000000..fbc44015 --- /dev/null +++ b/bench/human_in_the_loop2/reevaluate_qa.py @@ -0,0 +1,550 @@ +""" +Re-evaluate Question Answering results from data directory using LLM. + +This script: +1. Loads existing QA records from data directory +2. Re-evaluates each system_response using multiple models in parallel +3. Uses both EVALUATION_PROMPT_FOR_QUESTION and EVALUATION_PROMPT_FOR_QUESTION2 +4. Saves updated results with new evaluation metrics +""" + +import asyncio +import json +import re +import yaml +from collections import defaultdict +from pathlib import Path +from typing import Any + +from reme_ai.core.schema import Message +from reme_ai.core.utils import load_env +from reme_ai.reme import ReMe +from tenacity import retry, stop_after_attempt, wait_random_exponential + +# Load environment +load_env() + +# Initialize ReMe singleton +reme = ReMe() + +# Load prompts from YAML file +_YAML_PATH = Path(__file__).parent / "eval.yaml" +with open(_YAML_PATH, "r", encoding="utf-8") as f: + _PROMPTS = yaml.safe_load(f) + + +@retry( + wait=wait_random_exponential(min=1, max=60), + stop=stop_after_attempt(3), + reraise=True, +) +async def llm_request(prompt: str, model_name: str = "qwen3-max", **kwargs) -> str: + """Make an LLM request using ReMe's LLM.""" + assistant_message = await reme.llm.chat( + messages=[ + Message( + **{ + "role": "user", + "content": prompt, + }, + ), + ], + model_name=model_name, + **kwargs, + ) + return assistant_message.content + + +@retry( + wait=wait_random_exponential(min=1, max=60), + stop=stop_after_attempt(5), + reraise=True, +) +async def llm_request_for_json(prompt: str, model_name: str = "qwen3-max", **kwargs) -> dict: + """Make an LLM request expecting JSON response.""" + content = await llm_request(prompt, model_name=model_name, **kwargs) + + match = re.search(r"```json\s*(\{.*?\})\s*```", content, re.DOTALL) + if not match: + raise ValueError(f"No JSON block found in model output: {content}") + + json_str = match.group(1).strip() + return json.loads(json_str) + + +async def evaluate_qa_record( + question: str, + reference_answer: str, + key_memory_points: str, + response: str, + dialogue: str = "", + model_name: str = "qwen3-max", + prompt_version: str = "v1" +) -> dict: + """Evaluate a single QA record using LLM with specified prompt version. + + Args: + question: The question to evaluate + reference_answer: The reference answer + key_memory_points: Key memory points + response: System response to evaluate + dialogue: Dialogue context (optional) + model_name: LLM model name + prompt_version: "v1" for EVALUATION_PROMPT_FOR_QUESTION, + "v2" for EVALUATION_PROMPT_FOR_QUESTION2 + + Returns: + dict with evaluation_result and reasoning + """ + # Select prompt template + if prompt_version == "v2": + prompt_template = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION2"] + else: + prompt_template = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION"] + + # Format prompt + prompt = prompt_template.format( + question=question, + reference_answer=reference_answer, + key_memory_points=key_memory_points, + response=response, + dialogue=dialogue or "N/A" + ) + + result = await llm_request_for_json(prompt, model_name=model_name) + return result + + +def load_from_data_dir(data_dir: str) -> list[dict]: + """Load data from data directory (same as compute_qa_stats.py).""" + data_path = Path(data_dir) + + # Try flat file structure first (conversation_{user}_session_{idx}.json) + json_files = [f for f in data_path.iterdir() if f.is_file() and f.suffix == ".json"] + + if json_files: + # Group files by user + users_dict = defaultdict(list) + + for json_file in json_files: + with open(json_file, "r", encoding="utf-8") as f: + session_data = json.load(f) + user_name = session_data.get("user_name") + if user_name: + users_dict[user_name].append({ + "file": json_file, + "data": session_data + }) + + # Sort sessions by session_idx for each user + users_data = [] + for user_name, sessions in users_dict.items(): + sessions_sorted = sorted( + sessions, + key=lambda s: s["data"].get("session_idx", 0) + ) + users_data.extend(sessions_sorted) + + return users_data + + return [] + + +def format_dialogue_context(session_data: dict) -> str: + """Format dialogue context from session data.""" + dialogue = session_data.get("session", {}).get("dialogue", []) + if not dialogue: + return "N/A" + + formatted_turns = [] + for turn in dialogue: + role = turn.get("role", "unknown") + content = turn.get("content", "") + timestamp = turn.get("timestamp", "") + formatted_turns.append( + f"Role: {role}\nContent: {content}\nTime: {timestamp}" + ) + return "\n\n".join(formatted_turns) + + +async def reevaluate_session( + session_file: Path, + session_data: dict, + models: list[str], + prompt_versions: list[str], + parallel: bool = True +) -> dict: + """Re-evaluate all QA records in a session using multiple models and prompts. + + Args: + session_file: Path to session file + session_data: Session data dict + models: List of model names to use for evaluation + prompt_versions: List of prompt versions ("v1", "v2") + parallel: If True, use asyncio.gather for parallel execution; + if False, execute sequentially + + Returns: + Updated session data with evaluation results for each model+prompt combination + + Note: + Request rate limiting is handled by base_llm.py's request_interval mechanism. + """ + eval_results = session_data.get("session", {}).get("evaluation_results", {}) + qa_records = eval_results.get("question_answering_records", []) + + if not qa_records: + print(f" ⏭️ No QA records found") + return session_data + + total_evals = len(models) * len(prompt_versions) * len(qa_records) + print(f" 🔍 Re-evaluating {len(qa_records)} QA records with {len(models)} models × {len(prompt_versions)} prompts = {total_evals} evaluations...") + + # Format dialogue context once + dialogue_context = format_dialogue_context(session_data) + + async def evaluate_single_combination( + idx: int, + qa: dict, + model_name: str, + prompt_version: str + ) -> tuple[int, str, str, dict]: + """Evaluate a single QA record with specific model and prompt. + + Note: Rate limiting is handled by BaseLLM's request_interval mechanism. + """ + question = qa.get("question", "") + reference_answer = qa.get("answer", "") + + # Get key memory points from evidence + evidence = qa.get("evidence", []) + key_memory_points = "\n".join([e.get("memory_content", "") for e in evidence]) + + # Get system response + system_response = qa.get("system_response", "") + + try: + # Call LLM for evaluation + eval_result = await evaluate_qa_record( + question=question, + reference_answer=reference_answer, + key_memory_points=key_memory_points, + response=system_response, + dialogue=dialogue_context, + model_name=model_name, + prompt_version=prompt_version + ) + + result = { + "result_type": eval_result.get("evaluation_result", "Invalid"), + "reasoning": eval_result.get("reasoning", "") + } + + return idx, model_name, prompt_version, result + + except Exception as e: + print(f" ❌ QA[{idx+1}] {model_name}/{prompt_version}: Error: {e}") + return idx, model_name, prompt_version, { + "result_type": "Error", + "reasoning": f"Evaluation error: {str(e)}" + } + + # Create all evaluation tasks (all combinations of models, prompts, and QA records) + tasks = [] + for idx, qa in enumerate(qa_records): + for model_name in models: + for prompt_version in prompt_versions: + tasks.append(evaluate_single_combination(idx, qa, model_name, prompt_version)) + + # Execute evaluations based on parallel mode + if parallel: + print(f" ⚡ Starting {len(tasks)} parallel evaluations (rate limited by LLM layer)...") + results = await asyncio.gather(*tasks) + else: + print(f" 🔄 Starting {len(tasks)} sequential evaluations...") + results = [] + for i, task in enumerate(tasks, 1): + result = await task + results.append(result) + if i % 10 == 0 or i == len(tasks): + print(f" ⏳ Progress: {i}/{len(tasks)} evaluations completed") + + # Organize results by QA index, then by model and prompt + # Structure: qa_records[idx]["evaluations"][model][prompt_version] = {result_type, reasoning} + for idx, qa in enumerate(qa_records): + if "evaluations" not in qa: + qa["evaluations"] = {} + + # Initialize evaluations structure + for model_name in models: + if model_name not in qa["evaluations"]: + qa["evaluations"][model_name] = {} + + # Fill in results + completed_count = 0 + for qa_idx, model_name, prompt_version, result in results: + qa_records[qa_idx]["evaluations"][model_name][prompt_version] = result + completed_count += 1 + if completed_count % 10 == 0 or completed_count == len(results): + print(f" ✅ Completed {completed_count}/{len(results)} evaluations") + + # Set default result_type to first model's v1 result for compatibility + if models and prompt_versions: + default_model = models[0] + default_prompt = prompt_versions[0] + for qa in qa_records: + default_eval = qa["evaluations"].get(default_model, {}).get(default_prompt, {}) + qa["result_type"] = default_eval.get("result_type", "Invalid") + qa["question_answering_reasoning"] = default_eval.get("reasoning", "") + + # Update session data + if "session" not in session_data: + session_data["session"] = {} + if "evaluation_results" not in session_data["session"]: + session_data["session"]["evaluation_results"] = {} + + session_data["session"]["evaluation_results"]["question_answering_records"] = qa_records + + # Save updated session data + with open(session_file, "w", encoding="utf-8") as f: + json.dump(session_data, f, ensure_ascii=False, indent=2) + + print(f" 💾 Updated session saved with all evaluations") + + return session_data + + +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 = hallucination = omission = valid = 0 + + for qa in qa_records: + result_type = qa.get("result_type", "") + if result_type == "Correct": + correct += 1 + valid += 1 + elif result_type == "Hallucination": + hallucination += 1 + valid += 1 + elif result_type == "Omission": + omission += 1 + valid += 1 + + metrics = { + "correct_qa_ratio(all)": correct / total, + "hallucination_qa_ratio(all)": hallucination / total, + "omission_qa_ratio(all)": omission / total, + "correct_qa_ratio(valid)": correct / valid if valid > 0 else 0, + "hallucination_qa_ratio(valid)": hallucination / valid if valid > 0 else 0, + "omission_qa_ratio(valid)": omission / valid if valid > 0 else 0, + "qa_valid_num": valid, + "qa_num": total + } + + return metrics + + +async def main( + data_dir: str = "./data", + models: list[str] = None, + prompt_versions: list[str] = None, + parallel: bool = True +): + """Main function to re-evaluate QA records from data directory with multiple models and prompts. + + Args: + data_dir: Path to data directory + models: List of model names (e.g., ["qwen3-max", "qwen-flash"]) + prompt_versions: List of prompt versions (e.g., ["v1", "v2"]) + parallel: If True, use parallel execution; if False, use sequential execution + + Note: + Request rate limiting is automatically handled by base_llm.py's request_interval mechanism. + """ + data_path = Path(data_dir) + + if not data_path.exists() or not data_path.is_dir(): + print(f"❌ Error: Directory not found: {data_dir}") + return + + # Default values + if models is None: + models = ["qwen3-max"] + if prompt_versions is None: + prompt_versions = ["v1"] + + print("=" * 80) + print("RE-EVALUATING QA RECORDS WITH MULTIPLE MODELS & PROMPTS") + print(f"Models: {', '.join(models)}") + print(f"Prompts: {', '.join(['EVALUATION_PROMPT_FOR_QUESTION' if v == 'v1' else 'EVALUATION_PROMPT_FOR_QUESTION2' for v in prompt_versions])}") + print(f"Execution mode: {'Parallel' if parallel else 'Sequential'}") + print("Note: Request rate limiting handled by LLM layer (base_llm.py)") + print("=" * 80 + "\n") + + # Load data from directory + sessions = load_from_data_dir(data_dir) + + if not sessions: + print(f"❌ No session files found in {data_dir}") + return + + print(f"📂 Found {len(sessions)} session files\n") + + # Process each session + all_qa_records = [] + + for idx, session_info in enumerate(sessions, 1): + session_file = session_info["file"] + session_data = session_info["data"] + user_name = session_data.get("user_name", "Unknown") + session_idx = session_data.get("session_idx", 0) + + print(f"[{idx}/{len(sessions)}] {user_name} - Session {session_idx}") + + updated_session = await reevaluate_session( + session_file=session_file, + session_data=session_data, + models=models, + prompt_versions=prompt_versions, + parallel=parallel + ) + + # Collect QA records for metrics + eval_results = updated_session.get("session", {}).get("evaluation_results", {}) + qa_records = eval_results.get("question_answering_records", []) + all_qa_records.extend(qa_records) + + print() + + # Compute and display metrics for each model+prompt combination + print("=" * 80) + print("UPDATED METRICS (BY MODEL & PROMPT)") + print("=" * 80 + "\n") + + for model_name in models: + for prompt_version in prompt_versions: + prompt_name = "EVALUATION_PROMPT_FOR_QUESTION" if prompt_version == "v1" else "EVALUATION_PROMPT_FOR_QUESTION2" + print(f"\n📊 {model_name} / {prompt_name}:") + print("─" * 80) + + # Extract QA records for this model+prompt combination + model_qa_records = [] + for qa in all_qa_records: + eval_data = qa.get("evaluations", {}).get(model_name, {}).get(prompt_version, {}) + if eval_data: + # Create a copy with the specific evaluation result + qa_copy = { + **qa, + "result_type": eval_data.get("result_type", "Invalid"), + "question_answering_reasoning": eval_data.get("reasoning", "") + } + model_qa_records.append(qa_copy) + + if model_qa_records: + metrics = compute_qa_metrics(model_qa_records) + + print(f" Correct (all): {metrics['correct_qa_ratio(all)']:.4f}") + print(f" Hallucination (all): {metrics['hallucination_qa_ratio(all)']:.4f}") + print(f" Omission (all): {metrics['omission_qa_ratio(all)']:.4f}") + print(f" Correct (valid): {metrics['correct_qa_ratio(valid)']:.4f}") + print(f" Hallucination (valid): {metrics['hallucination_qa_ratio(valid)']:.4f}") + print(f" Omission (valid): {metrics['omission_qa_ratio(valid)']:.4f}") + print(f" Valid/Total: {metrics['qa_valid_num']}/{metrics['qa_num']}") + + # Save detailed results with all evaluations + report_file = data_path.parent / "reme_eval_stat_result_detailed.json" + + # Create summary for each model+prompt combination + evaluation_summary = {} + for model_name in models: + evaluation_summary[model_name] = {} + for prompt_version in prompt_versions: + prompt_name = "EVALUATION_PROMPT_FOR_QUESTION" if prompt_version == "v1" else "EVALUATION_PROMPT_FOR_QUESTION2" + + # Extract QA records for this combination + model_qa_records = [] + for qa in all_qa_records: + eval_data = qa.get("evaluations", {}).get(model_name, {}).get(prompt_version, {}) + if eval_data: + qa_copy = { + **qa, + "result_type": eval_data.get("result_type", "Invalid"), + "question_answering_reasoning": eval_data.get("reasoning", "") + } + model_qa_records.append(qa_copy) + + metrics = compute_qa_metrics(model_qa_records) + evaluation_summary[model_name][prompt_name] = { + "metrics": metrics, + "qa_records": model_qa_records + } + + final_results = { + "evaluation_summary": evaluation_summary, + "all_qa_records_with_evaluations": all_qa_records + } + + with open(report_file, "w", encoding="utf-8") as f: + json.dump(final_results, f, ensure_ascii=False, indent=2) + + print(f"\n💾 Detailed results saved: {report_file}") + print("\n" + "=" * 80) + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser( + description="Re-evaluate QA records from data directory using multiple models and prompts in parallel. " + "Request rate limiting is automatically handled by base_llm.py's request_interval mechanism." + ) + parser.add_argument( + "data_dir", + nargs='?', + default="./data", + type=str, + help="Path to data directory containing user session files (default: ./data)" + ) + parser.add_argument( + "--models", + type=str, + nargs='+', + default=["qwen3-max", "qwen-flash", "qwen-plus", "qwen3-30b-a3b-instruct-2507", "qwen3-235b-a22b-instruct-2507"], + help="LLM model names for evaluation (space-separated, default: qwen3-max)" + ) + # ["gpt-4o-2024-11-20"], ["qwen3-max", "qwen-flash", "qwen3-30b-a3b-instruct-2507", "qwen3-30b-a3b-instruct-2507"] + parser.add_argument( + "--prompts", + type=str, + nargs='+', + choices=["v1", "v2"], + default=["v1", "v2"], + help="Prompt versions to use: v1=EVALUATION_PROMPT_FOR_QUESTION, v2=EVALUATION_PROMPT_FOR_QUESTION2 (default: v1)" + ) + parser.add_argument( + "--serial", + action="store_true", + help="Use sequential execution instead of parallel (default: parallel)" + ) + + args = parser.parse_args() + + asyncio.run(main( + data_dir=args.data_dir, + models=args.models, + prompt_versions=args.prompts, + parallel=not args.serial + )) diff --git a/docs/todo.md b/docs/todo.md new file mode 100644 index 00000000..5d73e04d --- /dev/null +++ b/docs/todo.md @@ -0,0 +1,3 @@ +1. 如何更好的注册class +2. op的返回,使用return 还是 self.output +3. 如何把agent的东西放出来 \ No newline at end of file diff --git a/docs/work_memory/message_offload_ops.md b/docs/work_memory/message_offload_ops.md index d3965918..c05e22f7 100644 --- a/docs/work_memory/message_offload_ops.md +++ b/docs/work_memory/message_offload_ops.md @@ -96,7 +96,7 @@ When context grows too large, model performance degrades significantly—a pheno ### Usage Pattern For complete working examples of how to use MessageOffloadOp in practice, please refer to: -[test_message_offload_op.py](../../test_op/test_message_offload_op.py) +[test_message_offload_op.py](../../test/test_message_offload_op.py) This test file demonstrates: - **Compact mode**: How to configure and use compaction-only strategy diff --git a/docs/work_memory/message_reload_ops.md b/docs/work_memory/message_reload_ops.md index 6ee80136..c03f478d 100644 --- a/docs/work_memory/message_reload_ops.md +++ b/docs/work_memory/message_reload_ops.md @@ -114,7 +114,7 @@ Example: Reading `/workspace/context_store/tool_call_123.txt` with `offset=0` an ## Usage Pattern: Combining Grep and ReadFile For a complete working example of how to use these operations in practice, please refer to: -[test_agentic_retrieve_op.py](../../test_op/test_agentic_retrieve_op.py) +[test_agentic_retrieve_op.py](../../test/test_agentic_retrieve_op.py) This test file demonstrates: - How to configure the system prompt to guide AI in using Grep and ReadFile operations diff --git a/pyproject.toml b/pyproject.toml index ba7eb72d..47aa478d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -57,7 +57,7 @@ full = [ [tool.setuptools.packages.find] where = ["."] -include = ["reme_ai*"] +include = ["reme_ai*", "reme*"] exclude = ["test*", "cookbook*", "doc*", "library*", "dist*"] [tool.setuptools.package-data] diff --git a/reme/__init__.py b/reme/__init__.py new file mode 100644 index 00000000..32f34911 --- /dev/null +++ b/reme/__init__.py @@ -0,0 +1,19 @@ +"""ReMe""" + +from . import agent +from . import config +from . import core +from . import tool +from . import workflow +from .reme_app import ReMeApp + +__all__ = [ + "agent", + "config", + "core", + "tool", + "workflow", + "ReMeApp", +] + +__version__ = "0.3.0.0a1" diff --git a/reme/agent/__init__.py b/reme/agent/__init__.py new file mode 100644 index 00000000..45fed6e8 --- /dev/null +++ b/reme/agent/__init__.py @@ -0,0 +1,7 @@ +"""A simple chatbot.""" + +from . import chat + +__all__ = [ + "chat", +] diff --git a/reme_ai/mem_agent/chat/__init__.py b/reme/agent/chat/__init__.py similarity index 57% rename from reme_ai/mem_agent/chat/__init__.py rename to reme/agent/chat/__init__.py index be2fc055..ed3049ef 100644 --- a/reme_ai/mem_agent/chat/__init__.py +++ b/reme/agent/chat/__init__.py @@ -1,11 +1,13 @@ """chat agent""" -from .remy_agent import ReMyAgent from .simple_chat import SimpleChat from .stream_chat import StreamChat +from ...core import R __all__ = [ - "ReMyAgent", "StreamChat", "SimpleChat", ] + +R.op.register("simple_chat")(SimpleChat) +R.op.register("stream_chat")(StreamChat) diff --git a/reme_ai/mem_agent/chat/simple_chat.py b/reme/agent/chat/simple_chat.py similarity index 93% rename from reme_ai/mem_agent/chat/simple_chat.py rename to reme/agent/chat/simple_chat.py index 8a71c7a8..36181346 100644 --- a/reme_ai/mem_agent/chat/simple_chat.py +++ b/reme/agent/chat/simple_chat.py @@ -2,14 +2,12 @@ from loguru import logger -from ...core.context import C from ...core.enumeration import Role -from ...core.op import BaseOp +from ...core.op import BaseTool from ...core.schema import Message, ToolCall -@C.register_op() -class SimpleChat(BaseOp): +class SimpleChat(BaseTool): """Simple chat agent that handles non-streaming conversations.""" def _build_tool_call(self) -> ToolCall: @@ -59,4 +57,4 @@ class SimpleChat(BaseOp): logger.info(f"messages={messages}") assistant_message = await self.llm.chat(messages=messages) logger.info(f"assistant_message={assistant_message.simple_dump()}") - self.output = assistant_message.content + return assistant_message.content diff --git a/reme_ai/mem_agent/chat/stream_chat.py b/reme/agent/chat/stream_chat.py similarity index 95% rename from reme_ai/mem_agent/chat/stream_chat.py rename to reme/agent/chat/stream_chat.py index 470e4647..f5cd7a5a 100644 --- a/reme_ai/mem_agent/chat/stream_chat.py +++ b/reme/agent/chat/stream_chat.py @@ -2,14 +2,12 @@ from loguru import logger -from ...core.context import C from ...core.enumeration import Role, ChunkEnum -from ...core.op import BaseOp +from ...core.op import BaseTool from ...core.schema import Message, ToolCall -@C.register_op() -class StreamChat(BaseOp): +class StreamChat(BaseTool): """Streaming chat agent that handles real-time conversation streaming.""" def _build_tool_call(self) -> ToolCall: diff --git a/reme_ai/core/config/__init__.py b/reme/config/__init__.py similarity index 100% rename from reme_ai/core/config/__init__.py rename to reme/config/__init__.py diff --git a/reme_ai/core/config/default.yaml b/reme/config/default.yaml similarity index 87% rename from reme_ai/core/config/default.yaml rename to reme/config/default.yaml index 5891d180..836d866f 100644 --- a/reme_ai/core/config/default.yaml +++ b/reme/config/default.yaml @@ -16,13 +16,14 @@ llm: default: backend: openai model_name: qwen3-30b-a3b-instruct-2507 - max_concurrency: 20 +# model_name: qwen-flash + request_interval: 1 + temperature: 0.0001 qwen3_max_instruct: backend: openai model_name: qwen3-max -# temperature: 0.6 - max_concurrency: 20 + request_interval: 2 embedding_model: default: diff --git a/reme_ai/core/config/reme_config_parser.py b/reme/config/reme_config_parser.py similarity index 76% rename from reme_ai/core/config/reme_config_parser.py rename to reme/config/reme_config_parser.py index 798235b2..7a21f806 100644 --- a/reme_ai/core/config/reme_config_parser.py +++ b/reme/config/reme_config_parser.py @@ -1,6 +1,6 @@ """Configuration parser for ReMe framework.""" -from ..utils import PydanticConfigParser +from ..core.utils import PydanticConfigParser class ReMeConfigParser(PydanticConfigParser): diff --git a/reme_ai/core/__init__.py b/reme/core/__init__.py similarity index 52% rename from reme_ai/core/__init__.py rename to reme/core/__init__.py index 8eab5792..251c2433 100644 --- a/reme_ai/core/__init__.py +++ b/reme/core/__init__.py @@ -1,9 +1,5 @@ -"""Core module for ReMe AI framework.""" +"""Core""" -# pylint: disable=wrong-import-position -# flake8: noqa: F401 - -from . import config from . import context from . import embedding from . import enumeration @@ -15,3 +11,19 @@ from . import service from . import token_counter from . import utils from . import vector_store +from .context import R + +__all__ = [ + "context", + "embedding", + "enumeration", + "flow", + "llm", + "op", + "schema", + "service", + "token_counter", + "utils", + "vector_store", + "R", +] diff --git a/reme_ai/core/context/__init__.py b/reme/core/context/__init__.py similarity index 69% rename from reme_ai/core/context/__init__.py rename to reme/core/context/__init__.py index 7f26d600..7bbd5869 100644 --- a/reme_ai/core/context/__init__.py +++ b/reme/core/context/__init__.py @@ -2,15 +2,14 @@ from .base_context import BaseContext from .prompt_handler import PromptHandler -from .registry import Registry +from .registry_factory import R from .runtime_context import RuntimeContext -from .service_context import ServiceContext, C +from .service_context import ServiceContext __all__ = [ "BaseContext", "PromptHandler", - "Registry", + "R", "RuntimeContext", "ServiceContext", - "C", ] diff --git a/reme_ai/core/context/base_context.py b/reme/core/context/base_context.py similarity index 100% rename from reme_ai/core/context/base_context.py rename to reme/core/context/base_context.py diff --git a/reme/core/context/prompt_handler.py b/reme/core/context/prompt_handler.py new file mode 100644 index 00000000..e6b6d737 --- /dev/null +++ b/reme/core/context/prompt_handler.py @@ -0,0 +1,363 @@ +"""Module for managing and formatting prompt templates from files or dictionaries. + +This module provides a PromptHandler class that: +- Loads prompts from YAML/JSON files or dictionaries +- Supports multi-language prompts with automatic suffix handling +- Provides conditional line filtering using boolean flags +- Formats prompts with template variable substitution +- Validates format strings and provides helpful error messages +""" + +import json +from pathlib import Path +from string import Formatter +from typing import Any, Dict, Optional, Union + +import yaml +from loguru import logger + +from .base_context import BaseContext + + +class PromptNotFoundError(KeyError): + """Exception raised when a requested prompt template is not found.""" + + def __init__(self, prompt_name: str, available_prompts: list[str]): + self.prompt_name = prompt_name + self.available_prompts = available_prompts + super().__init__( + f"Prompt '{prompt_name}' not found. " + f"Available prompts: {', '.join(available_prompts[:10])}" + f"{'...' if len(available_prompts) > 10 else ''}", + ) + + +class PromptFormattingError(ValueError): + """Exception raised when prompt formatting fails.""" + + +class PromptHandler(BaseContext): + """A context-aware handler for loading, retrieving, and formatting prompt templates. + + This handler supports: + - Loading prompts from YAML/JSON files or dictionaries + - Multi-language prompt support with automatic language suffix + - Conditional line filtering using boolean flags (e.g., [debug], [verbose]) + - Template variable substitution with validation + - Method chaining for fluent API + + Examples: + >>> handler = PromptHandler(language="en") + >>> handler.load_prompt_dict({ + ... "greeting_en": "Hello, {name}!", + ... "farewell_en": "[debug]Debug mode\\nGoodbye, {name}!" + ... }) + >>> handler.prompt_format("greeting", name="Alice") + 'Hello, Alice!' + >>> handler.prompt_format("farewell", name="Bob", debug=False) + 'Goodbye, Bob!' + """ + + def __init__(self, language: str = "", **kwargs): + """Initialize the PromptHandler with optional language configuration. + + Args: + language: Language code to append as suffix (e.g., "en", "zh", "ja"). + If provided, get_prompt will automatically try to find + prompts with this suffix (e.g., "greeting" -> "greeting_en"). + **kwargs: Additional key-value pairs to initialize the context. + """ + super().__init__(**kwargs) + self.language: str = language.strip() + + def load_prompt_by_file( + self, + prompt_file_path: Optional[Union[Path, str]] = None, + overwrite: bool = True, + ) -> "PromptHandler": + """Load prompt configurations from a YAML or JSON file into the context. + + Supports both YAML (.yaml, .yml) and JSON (.json) file formats. + Non-existent files are silently skipped. + + Args: + prompt_file_path: Path to the prompt configuration file. + If None, returns self without changes. + overwrite: If True, allows overwriting existing prompts with warnings. + If False, skips existing prompts without overwriting. + + Returns: + Self for method chaining. + + Raises: + ValueError: If file format is not supported. + yaml.YAMLError: If YAML parsing fails. + json.JSONDecodeError: If JSON parsing fails. + """ + if prompt_file_path is None: + return self + + if isinstance(prompt_file_path, str): + prompt_file_path = Path(prompt_file_path) + + if not prompt_file_path.exists(): + logger.warning(f"Prompt file not found: {prompt_file_path}") + return self + + suffix = prompt_file_path.suffix.lower() + + try: + with prompt_file_path.open(encoding="utf-8") as f: + if suffix in [".yaml", ".yml"]: + prompt_dict = yaml.safe_load(f) + elif suffix == ".json": + prompt_dict = json.load(f) + else: + raise ValueError( + f"Unsupported file format: {suffix}. " f"Supported formats: .yaml, .yml, .json", + ) + + logger.info(f"Loaded {len(prompt_dict or {})} prompts from {prompt_file_path}") + self.load_prompt_dict(prompt_dict, overwrite=overwrite) + + except (yaml.YAMLError, json.JSONDecodeError) as e: + logger.error(f"Failed to parse prompt file {prompt_file_path}: {e}") + raise + + return self + + def load_prompt_dict( + self, + prompt_dict: Optional[Dict[str, Any]] = None, + overwrite: bool = True, + ) -> "PromptHandler": + """Merge a dictionary of prompt strings into the current context. + + Only string values are stored as prompts. Non-string values are skipped. + + Args: + prompt_dict: Dictionary mapping prompt names to prompt template strings. + overwrite: If True, allows overwriting existing prompts with warnings. + If False, skips existing prompts without overwriting. + + Returns: + Self for method chaining. + """ + if not prompt_dict: + return self + + for key, value in prompt_dict.items(): + if not isinstance(value, str): + logger.debug(f"Skipping non-string prompt: key={key}, type={type(value)}") + continue + + if key in self: + if overwrite: + logger.warning( + f"Overwriting prompt '{key}': " f"old length={len(self[key])}, new length={len(value)}", + ) + self[key] = value + else: + logger.debug(f"Skipping existing prompt: key={key}") + else: + logger.debug(f"Adding new prompt: key={key}, length={len(value)}") + self[key] = value + + return self + + def get_prompt(self, prompt_name: str, fallback_to_base: bool = True) -> str: + """Retrieve a prompt by name with automatic language suffix handling. + + If a language is configured, this method will: + 1. First try to find the prompt with language suffix (e.g., "greeting_en") + 2. If not found and fallback_to_base is True, try the base name (e.g., "greeting") + 3. Otherwise, raise PromptNotFoundError + + Args: + prompt_name: Name of the prompt to retrieve. + fallback_to_base: If True and language-specific prompt not found, + fallback to prompt without language suffix. + + Returns: + The prompt template string, stripped of leading/trailing whitespace. + + Raises: + PromptNotFoundError: If the prompt is not found. + """ + # Try with language suffix first + if self.language and not prompt_name.endswith(f"_{self.language}"): + key_with_lang = f"{prompt_name}_{self.language}" + if key_with_lang in self: + return self[key_with_lang].strip() + + # Try base name + if prompt_name in self: + return self[prompt_name].strip() + + # Try fallback if enabled + if fallback_to_base and self.language: + # Check if prompt_name already has language suffix, try without it + if prompt_name.endswith(f"_{self.language}"): + base_name = prompt_name[: -(len(self.language) + 1)] + if base_name in self: + return self[base_name].strip() + + # Not found, raise error with helpful message + available = list(self.keys()) + raise PromptNotFoundError(prompt_name, available) + + def has_prompt(self, prompt_name: str) -> bool: + """Check if a prompt exists (with or without language suffix). + + Args: + prompt_name: Name of the prompt to check. + + Returns: + True if the prompt exists, False otherwise. + """ + try: + self.get_prompt(prompt_name) + return True + except PromptNotFoundError: + return False + + def list_prompts(self, language_filter: Optional[str] = None) -> list[str]: + """List all available prompt names. + + Args: + language_filter: If provided, only return prompts for this language. + If None, return all prompts. + + Returns: + List of prompt names. + """ + if language_filter is None: + return list(self.keys()) + + suffix = f"_{language_filter.strip()}" + return [key for key in self.keys() if key.endswith(suffix)] + + @staticmethod + def _extract_format_fields(template: str) -> set[str]: + """Extract all format field names from a template string. + + Args: + template: Template string with {variable} placeholders. + + Returns: + Set of field names used in the template. + """ + return {field_name for _, field_name, _, _ in Formatter().parse(template) if field_name is not None} + + @staticmethod + def _filter_conditional_lines(prompt: str, flags: Dict[str, bool]) -> str: + """Filter lines based on boolean flags. + + Lines starting with [flag_name] are conditionally included based on + the value of flags[flag_name]. If True, the line is included (without + the flag marker). If False, the line is excluded. + + Args: + prompt: The prompt text with conditional markers. + flags: Dictionary of flag names to boolean values. + + Returns: + Filtered prompt text. + """ + filtered_lines = [] + + for line in prompt.split("\n"): + # Check each flag + matched_flag = None + for flag_name in flags: + marker = f"[{flag_name}]" + if line.startswith(marker): + matched_flag = flag_name + break + + if matched_flag is None: + # No flag marker, always include + filtered_lines.append(line) + elif flags[matched_flag]: + # Flag is True, include without marker + marker = f"[{matched_flag}]" + filtered_lines.append(line[len(marker) :]) + # else: Flag is False, skip this line + + return "\n".join(filtered_lines) + + def prompt_format( + self, + prompt_name: str, + validate: bool = True, + **kwargs, + ) -> str: + """Format a prompt with conditional line filtering and variable substitution. + + This method performs two-stage formatting: + 1. Conditional line filtering: Lines marked with [flag] are included only + if the corresponding boolean kwarg is True. + 2. Variable substitution: Template variables {var} are replaced with + provided values. + + Args: + prompt_name: Name of the prompt to format. + validate: If True, check that all required template variables are provided. + **kwargs: Keyword arguments for formatting. Boolean values are treated as + conditional flags, other values are used for template substitution. + + Returns: + Formatted prompt string. + + Raises: + PromptNotFoundError: If the prompt is not found. + PromptFormattingError: If validation fails or formatting errors occur. + + Examples: + >>> handler = PromptHandler() + >>> handler["test"] = "[debug]Debug: {info}\\nResult: {value}" + >>> handler.prompt_format("test", debug=False, info="test", value=42) + 'Result: 42' + >>> handler.prompt_format("test", debug=True, info="test", value=42) + 'Debug: test\\nResult: 42' + """ + # Get the prompt template + prompt = self.get_prompt(prompt_name) + + # Separate boolean flags from format variables + flag_kwargs = {k: v for k, v in kwargs.items() if isinstance(v, bool)} + format_kwargs = {k: v for k, v in kwargs.items() if not isinstance(v, bool)} + + # Step 1: Filter conditional lines + if flag_kwargs: + prompt = self._filter_conditional_lines(prompt, flag_kwargs) + + # Step 2: Validate required fields if requested + if validate: + required_fields = self._extract_format_fields(prompt) + missing_fields = required_fields - set(format_kwargs.keys()) + + if missing_fields: + raise PromptFormattingError( + f"Missing required format variables for prompt '{prompt_name}': " + f"{', '.join(sorted(missing_fields))}", + ) + + # Step 3: Format with variables + try: + if format_kwargs: + prompt = prompt.format(**format_kwargs) + except KeyError as e: + raise PromptFormattingError( + f"Format error in prompt '{prompt_name}': missing variable {e}", + ) from e + except (ValueError, IndexError) as e: + raise PromptFormattingError( + f"Format error in prompt '{prompt_name}': {e}", + ) from e + + return prompt.strip() + + def __repr__(self) -> str: + """Return a string representation of the PromptHandler.""" + return f"PromptHandler(language='{self.language}', " f"num_prompts={len(self)})" diff --git a/reme/core/context/registry_factory.py b/reme/core/context/registry_factory.py new file mode 100644 index 00000000..28cb1c82 --- /dev/null +++ b/reme/core/context/registry_factory.py @@ -0,0 +1,45 @@ +"""Module providing a registry class for managing class-to-name mappings via decorators.""" + +import inspect +from typing import Callable, TypeVar + +from .base_context import BaseContext +from ..utils import singleton + +T = TypeVar("T") + + +class Registry(BaseContext): + """A registry container that uses decorators to map and store class references.""" + + def register(self, name: str | type = "") -> Callable[[type[T]], type[T]] | type[T]: + """Return a decorator that registers a class under a specific name in the registry.""" + if inspect.isclass(name): + self[name.__name__] = name + return name + + else: + + def decorator(cls): + key: str = name if isinstance(name, str) and name else cls.__name__ + self[key] = cls + return cls + + return decorator + + +@singleton +class RegistryFactory: + """A factory class for creating registries.""" + + def __init__(self): + self.llm = Registry() + self.embedding_model = Registry() + self.vector_store = Registry() + self.op = Registry() + self.flow = Registry() + self.service = Registry() + self.token_counter = Registry() + + +R = RegistryFactory() diff --git a/reme_ai/core/context/runtime_context.py b/reme/core/context/runtime_context.py similarity index 87% rename from reme_ai/core/context/runtime_context.py rename to reme/core/context/runtime_context.py index d7112e1c..d4c02050 100644 --- a/reme_ai/core/context/runtime_context.py +++ b/reme/core/context/runtime_context.py @@ -3,6 +3,7 @@ import asyncio from .base_context import BaseContext +from .service_context import ServiceContext from ..enumeration import ChunkEnum from ..schema import Response, StreamChunk @@ -14,21 +15,24 @@ class RuntimeContext(BaseContext): self, response: Response | None = None, stream_queue: asyncio.Queue | None = None, + service_context: ServiceContext | None = None, **kwargs, ): """Initialize the context with optional response and queue.""" super().__init__(**kwargs) - self.response = response or Response() - self.stream_queue = stream_queue + self.response: Response | None = response or Response() + self.stream_queue: asyncio.Queue | None = stream_queue + self.service_context: ServiceContext | None = service_context @classmethod def from_context(cls, context: "RuntimeContext | None" = None, **kwargs) -> "RuntimeContext": """Create a new context from an existing instance or keywords.""" if context is None: return cls(**kwargs) - - context.update(kwargs) - return context + else: + if kwargs: + context.update(kwargs) + return context async def _enqueue(self, chunk: StreamChunk) -> None: """Internal helper to put a chunk into the queue if it exists.""" diff --git a/reme/core/context/service_context.py b/reme/core/context/service_context.py new file mode 100644 index 00000000..5be7b449 --- /dev/null +++ b/reme/core/context/service_context.py @@ -0,0 +1,230 @@ +"""Service context.""" + +import os +from concurrent.futures import ThreadPoolExecutor + +from loguru import logger + +from .base_context import BaseContext +from .registry_factory import R +from ..schema import ServiceConfig +from ..utils import MCPClient, print_logo, PydanticConfigParser, init_logger, load_env, run_coro_safely + + +class ServiceContext(BaseContext): + """Service context.""" + + def __init__( + self, + *args, + llm_api_key: str | None = None, + llm_api_base: str | None = None, + embedding_api_key: str | None = None, + embedding_api_base: str | None = None, + service_config: ServiceConfig | None = None, + parser: type[PydanticConfigParser] | None = None, + config_path: str | None = None, + enable_logo: bool = True, + llm: dict | None = None, + embedding_model: dict | None = None, + vector_store: dict | None = None, + token_counter: dict | None = None, + **kwargs, + ): + super().__init__() + # Set environment variables + load_env() + self._update_env("REME_LLM_API_KEY", llm_api_key) + self._update_env("REME_LLM_BASE_URL", llm_api_base) + self._update_env("REME_EMBEDDING_API_KEY", embedding_api_key) + self._update_env("REME_EMBEDDING_BASE_URL", embedding_api_base) + + # Use default parser if not provided + parser_class = parser if parser is not None else PydanticConfigParser + self.parser = parser_class(ServiceConfig) + + # Service configuration + if service_config is None: + input_args = [] + if config_path: + input_args.append(f"config={config_path}") + if args: + input_args.extend(args) + if kwargs: + input_args.extend([f"{k}={v}" for k, v in kwargs.items()]) + service_config = self.parser.parse_args(*input_args) + self.service_config: ServiceConfig = service_config + + # Initialize logger + if self.service_config.init_logger: + init_logger() + + # Update service config with provided arguments + if llm: + self.update_section_config("llm", **llm) + if embedding_model: + self.update_section_config("embedding_model", **embedding_model) + if token_counter: + self.update_section_config("token_counter", **token_counter) + if vector_store: + self.update_section_config("vector_store", **vector_store) + + # Print the ReMe logo if enabled in configuration. + self.service_config.enable_logo = enable_logo + if self.service_config.enable_logo: + print_logo(service_config=self.service_config) + + # Service configuration and runtime settings + self.language: str = self.service_config.language + self.thread_pool: ThreadPoolExecutor = ThreadPoolExecutor(max_workers=service_config.thread_pool_max_workers) + + # Initialize Ray for distributed computing if configured + if self.service_config.ray_max_workers > 1: + import ray + + ray.init(num_cpus=self.service_config.ray_max_workers) + + from ..llm import BaseLLM + from ..embedding import BaseEmbeddingModel + from ..vector_store import BaseVectorStore + from ..token_counter import BaseTokenCounter + from ..flow import BaseFlow, ExpressionFlow + from ..service import BaseService + + # Initialize LLM instances + self.llms: dict[str, BaseLLM] = {} + for name, config in self.service_config.llm.items(): + self.llms[name] = R.llm[config.backend](model_name=config.model_name, **config.model_extra) + + # Initialize Embedding model instances + self.embedding_models: dict[str, BaseEmbeddingModel] = {} + for name, config in self.service_config.embedding_model.items(): + self.embedding_models[name] = R.embedding_model[config.backend]( + model_name=config.model_name, + **config.model_extra, + ) + + # Initialize Token counter instances + self.token_counters: dict[str, BaseTokenCounter] = {} + for name, config in self.service_config.token_counter.items(): + self.token_counters[name] = R.token_counter[config.backend]( + model_name=config.model_name, + **config.model_extra, + ) + + # Initialize Vector store instances + self.vector_stores: dict[str, BaseVectorStore] = {} + for name, config in self.service_config.vector_store.items(): + self.vector_stores[name] = R.vector_store[config.backend]( + collection_name=config.collection_name, + embedding_model=self.embedding_models[config.embedding_model], + thread_pool=self.thread_pool, + **config.model_extra, + ) + + # Initialize flow instances + self.flows: dict[str, BaseFlow] = {} + for name, flow_cls in R.flow.items(): + if not self._filter_flows(name): + continue + flow: "BaseFlow" = flow_cls(name=name, service_context=self) + self.flows[flow.name] = flow + + # Initialize flow instances from service config + for name, flow_config in self.service_config.flow.items(): + if not self._filter_flows(name): + continue + flow_config.name = name + flow: BaseFlow = ExpressionFlow(flow_config=flow_config, service_context=self) + self.flows[flow.name] = flow + + # Initialize service instance + self.service: BaseService = R.service[self.service_config.backend](service_context=self) + + # MCP server mapping: maps server_name -> {tool_name: ToolCall} + if self.service_config.mcp_servers: + self.mcp_server_mapping: dict[str, dict] = run_coro_safely(self.prepare_mcp_servers()) + else: + self.mcp_server_mapping: dict[str, dict] = {} + + @staticmethod + def _update_env(key: str, value: str | None): + """Update environment variable if value is provided.""" + if value: + os.environ[key] = value + + def update_section_config(self, section_name: str, **kwargs): + """Update a specific section of the service config with new values.""" + section_dict: dict = getattr(self.service_config, section_name) + if "default" not in section_dict: + raise KeyError(f"Default `{section_name}` config not found") + + current_config = section_dict["default"] + section_dict["default"] = current_config.model_copy(update=kwargs, deep=True) + + def _filter_flows(self, name: str) -> bool: + """Filter flows based on enabled_flows and disabled_flows configuration.""" + if self.service_config.enabled_flows: + return name in self.service_config.enabled_flows + elif self.service_config.disabled_flows: + return name not in self.service_config.disabled_flows + else: + return True + + async def prepare_mcp_servers(self): + """Prepare and initialize MCP server connections.""" + mcp_client = MCPClient(config={"mcpServers": self.service_config.mcp_servers}) + for server_name in self.service_config.mcp_servers.keys(): + try: + # Retrieve all available tool calls from this MCP server + tool_calls = await mcp_client.list_tool_calls(server_name=server_name, return_dict=False) + + # Build mapping: tool_name -> ToolCall for quick lookup + self.mcp_server_mapping[server_name] = {tool_call.name: tool_call for tool_call in tool_calls} + + # Log discovered tools for debugging + for tool_call in tool_calls: + logger.info(f"list_tool_calls: {server_name}@{tool_call.name} {tool_call.simple_input_dump()}") + + except Exception as e: + logger.exception(f"list_tool_calls: {server_name} error: {e}") + + async def close(self): + """Close all service components asynchronously.""" + for _, vector_store in self.vector_stores.items(): + await vector_store.close() + + for _, llm in self.llms.items(): + await llm.close() + + for _, embedding_model in self.embedding_models.items(): + await embedding_model.close() + + self.shutdown_thread_pool() + self.shutdown_ray() + + def close_sync(self): + """Close all service components synchronously.""" + for _, vector_store in self.vector_stores.items(): + run_coro_safely(vector_store.close()) + + for _, llm in self.llms.items(): + llm.close_sync() + + for _, embedding_model in self.embedding_models.items(): + embedding_model.close_sync() + + self.shutdown_thread_pool() + self.shutdown_ray() + + def shutdown_thread_pool(self, wait: bool = True): + """Shutdown the thread pool executor.""" + if self.thread_pool: + self.thread_pool.shutdown(wait=wait) + + def shutdown_ray(self, wait: bool = True): + """Shutdown Ray cluster if it was initialized.""" + if self.service_config and self.service_config.ray_max_workers > 1: + import ray + + ray.shutdown(_exiting_interpreter=not wait) diff --git a/reme_ai/core/embedding/__init__.py b/reme/core/embedding/__init__.py similarity index 65% rename from reme_ai/core/embedding/__init__.py rename to reme/core/embedding/__init__.py index c1d92375..f694d065 100644 --- a/reme_ai/core/embedding/__init__.py +++ b/reme/core/embedding/__init__.py @@ -3,9 +3,13 @@ from .base_embedding_model import BaseEmbeddingModel from .openai_embedding_model import OpenAIEmbeddingModel from .openai_embedding_model_sync import OpenAIEmbeddingModelSync +from ..context import R __all__ = [ "BaseEmbeddingModel", "OpenAIEmbeddingModel", "OpenAIEmbeddingModelSync", ] + +R.embedding_model.register("openai")(OpenAIEmbeddingModel) +R.embedding_model.register("openai_sync")(OpenAIEmbeddingModelSync) diff --git a/reme_ai/core/embedding/base_embedding_model.py b/reme/core/embedding/base_embedding_model.py similarity index 100% rename from reme_ai/core/embedding/base_embedding_model.py rename to reme/core/embedding/base_embedding_model.py diff --git a/reme_ai/core/embedding/openai_embedding_model.py b/reme/core/embedding/openai_embedding_model.py similarity index 96% rename from reme_ai/core/embedding/openai_embedding_model.py rename to reme/core/embedding/openai_embedding_model.py index a7e9f0c8..435229b9 100644 --- a/reme_ai/core/embedding/openai_embedding_model.py +++ b/reme/core/embedding/openai_embedding_model.py @@ -6,10 +6,8 @@ from typing import Literal from openai import AsyncOpenAI from .base_embedding_model import BaseEmbeddingModel -from ..context import C -@C.register_embedding_model("openai") class OpenAIEmbeddingModel(BaseEmbeddingModel): """Asynchronous embedding model implementation compatible with OpenAI-style APIs.""" diff --git a/reme_ai/core/embedding/openai_embedding_model_sync.py b/reme/core/embedding/openai_embedding_model_sync.py similarity index 94% rename from reme_ai/core/embedding/openai_embedding_model_sync.py rename to reme/core/embedding/openai_embedding_model_sync.py index 760732cd..cf3aac14 100644 --- a/reme_ai/core/embedding/openai_embedding_model_sync.py +++ b/reme/core/embedding/openai_embedding_model_sync.py @@ -3,10 +3,8 @@ from openai import OpenAI from .openai_embedding_model import OpenAIEmbeddingModel -from ..context import C -@C.register_embedding_model("openai_sync") class OpenAIEmbeddingModelSync(OpenAIEmbeddingModel): """Synchronous embedding model implementation that extends the asynchronous OpenAI model.""" diff --git a/reme_ai/core/enumeration/__init__.py b/reme/core/enumeration/__init__.py similarity index 100% rename from reme_ai/core/enumeration/__init__.py rename to reme/core/enumeration/__init__.py diff --git a/reme_ai/core/enumeration/chunk_enum.py b/reme/core/enumeration/chunk_enum.py similarity index 100% rename from reme_ai/core/enumeration/chunk_enum.py rename to reme/core/enumeration/chunk_enum.py diff --git a/reme_ai/core/enumeration/http_enum.py b/reme/core/enumeration/http_enum.py similarity index 100% rename from reme_ai/core/enumeration/http_enum.py rename to reme/core/enumeration/http_enum.py diff --git a/reme/core/enumeration/json_schema_enum.py b/reme/core/enumeration/json_schema_enum.py new file mode 100644 index 00000000..d66882e2 --- /dev/null +++ b/reme/core/enumeration/json_schema_enum.py @@ -0,0 +1,38 @@ +"""Defines the standard data types supported by JSON Schema. + +This enum maps common JSON Schema primitive types to their corresponding +Python runtime types, and provides a convenient string representation +compatible with JSON Schema (`"string"`, `"number"`, etc.). +""" + +from enum import Enum + + +class JsonSchemaEnum(Enum): + """Enumeration of valid JSON Schema data types. + + The enum value is the corresponding Python type, while the string + representation (`str(...)`) is the canonical JSON Schema type name. + """ + + # Textual data + STRING = str + + # Numeric values, including integers and floats + NUMBER = float + + # Integer-only numeric values + INTEGER = int + + # JSON objects (key-value mappings) + OBJECT = dict + + # Ordered JSON lists/arrays + ARRAY = list + + # Boolean values: true / false + BOOLEAN = bool + + def __str__(self) -> str: + """Return the lowercase JSON Schema type name for this enum member.""" + return self.name.lower() diff --git a/reme/core/enumeration/memory_type.py b/reme/core/enumeration/memory_type.py new file mode 100644 index 00000000..b9f5ed29 --- /dev/null +++ b/reme/core/enumeration/memory_type.py @@ -0,0 +1,33 @@ +"""Defines the high-level categories of memory managed by ReMe. + +This enumeration is used across the system to tag, route, and store different +kinds of memories (identity, personal context, procedures, tools, etc.). +""" + +from enum import Enum + + +class MemoryType(str, Enum): + """Enumeration of memory categories used by the memory subsystem. + + These types describe *what* a piece of memory is about, which guides + storage, retrieval, and summarization strategies. + """ + + # Long‑term, relatively stable attributes about the user (name, roles, etc.) + IDENTITY = "identity" + + # User-specific preferences, habits, and evolving personal context + PERSONAL = "personal" + + # How‑to knowledge, workflows, and step‑by‑step instructions + PROCEDURAL = "procedural" + + # Information learned about tools, APIs, and their usage patterns + TOOL = "tool" + + # Condensed representation of larger memory collections + SUMMARY = "summary" + + # Raw chronological interaction history, typically before summarization + HISTORY = "history" diff --git a/reme_ai/core/enumeration/registry_enum.py b/reme/core/enumeration/registry_enum.py similarity index 100% rename from reme_ai/core/enumeration/registry_enum.py rename to reme/core/enumeration/registry_enum.py diff --git a/reme_ai/core/enumeration/role.py b/reme/core/enumeration/role.py similarity index 100% rename from reme_ai/core/enumeration/role.py rename to reme/core/enumeration/role.py diff --git a/reme_ai/core/flow/__init__.py b/reme/core/flow/__init__.py similarity index 77% rename from reme_ai/core/flow/__init__.py rename to reme/core/flow/__init__.py index e74a2b5c..6d5a053b 100644 --- a/reme_ai/core/flow/__init__.py +++ b/reme/core/flow/__init__.py @@ -3,11 +3,9 @@ from .base_flow import BaseFlow from .cmd_flow import CmdFlow from .expression_flow import ExpressionFlow -from .simple_flow import SimpleFlow __all__ = [ "BaseFlow", "CmdFlow", "ExpressionFlow", - "SimpleFlow", ] diff --git a/reme_ai/core/flow/base_flow.py b/reme/core/flow/base_flow.py similarity index 80% rename from reme_ai/core/flow/base_flow.py rename to reme/core/flow/base_flow.py index 7decce48..b79528f7 100644 --- a/reme_ai/core/flow/base_flow.py +++ b/reme/core/flow/base_flow.py @@ -7,30 +7,25 @@ from abc import ABC, abstractmethod from loguru import logger -from ..context import C, RuntimeContext -from ..enumeration import ChunkEnum, RegistryEnum +from ..context import RuntimeContext, ServiceContext, R +from ..enumeration import ChunkEnum from ..op import BaseOp, SequentialOp, ParallelOp -from ..schema import Response, ToolCall, ToolAttr +from ..schema import Response, ToolCall from ..utils import camel_to_snake, CacheHandler class BaseFlow(ABC): - """Abstract base class for flow execution with caching, streaming, and operation tree management. - - BaseFlow provides a framework for building complex workflows by composing operations - into executable trees. It supports both synchronous and asynchronous execution modes, - response caching, streaming outputs, and automatic tool call schema generation. - """ + """Abstract base class for flow execution with caching, streaming, and operation tree management.""" def __init__( self, name: str = "", - flow_op: BaseOp | None = None, stream: bool = False, raise_exception: bool = True, enable_cache: bool = False, cache_path: str = "cache/flow", cache_expire_hours: float = 0.1, + service_context: ServiceContext | None = None, **kwargs, ): """Initialize flow configuration and execution state.""" @@ -42,11 +37,12 @@ class BaseFlow(ABC): self.enable_cache: bool = enable_cache self.cache_path: str = cache_path self.cache_expire_hours: float = cache_expire_hours + self.service_context: ServiceContext | None = service_context self.flow_params: dict = kwargs - self._flow_op: BaseOp | None = flow_op self._cache: CacheHandler | None = None self._flow_printed: bool = False + self._flow_op: BaseOp | None = None self._tool_call: ToolCall | None = None def _build_tool_call(self) -> ToolCall | None: @@ -82,11 +78,7 @@ class BaseFlow(ABC): return if key := self._compute_cache_key(params): - self.cache.save( - key, - response.model_dump(exclude_none=True), - expire_hours=self.cache_expire_hours, - ) + self.cache.save(key, response.model_dump(exclude_none=True), expire_hours=self.cache_expire_hours) def _print_operation_tree(self, name: str, op: BaseOp, indent: int): """Recursively log the hierarchy of the flow's operation tree.""" @@ -100,19 +92,13 @@ class BaseFlow(ABC): @property def tool_call(self) -> ToolCall | None: """Lazily construct the ToolCall schema describing this flow.""" - if self.flow_op.tool_call: + if hasattr(self.flow_op, "tool_call"): return self.flow_op.tool_call if self._tool_call is None: self._tool_call = self._build_tool_call() if self._tool_call: self._tool_call.name = self._tool_call.name or self.name - self._tool_call.output = self._tool_call.output or { - f"{self.name}_result": ToolAttr( - type="string", - description=f"The execution result of the {self.name}", - ), - } return self._tool_call @property @@ -130,12 +116,6 @@ class BaseFlow(ABC): self._flow_op = self._build_flow() return self._flow_op - @flow_op.setter - def flow_op(self, op: BaseOp): - """Set the root operation of the flow.""" - self._flow_op = op - self._flow_printed = False - @property def async_mode(self) -> bool: """Check if the current flow operation tree is asynchronous.""" @@ -148,11 +128,10 @@ class BaseFlow(ABC): if not lines: raise ValueError("Expression is empty") - env: dict = C.registry_dict[RegistryEnum.OP] if len(lines) > 1: - exec("\n".join(lines[:-1]), {"__builtins__": {}}, env) + exec("\n".join(lines[:-1]), {"__builtins__": {}}, R.op) - result = eval(lines[-1], {"__builtins__": {}}, env) + result = eval(lines[-1], {"__builtins__": {}}, R.op) if not isinstance(result, BaseOp): raise TypeError(f"Expression evaluated to {type(result)}, expected BaseOp") return result @@ -172,30 +151,34 @@ class BaseFlow(ABC): if cached := self._maybe_load_cached(kwargs): return cached - context = RuntimeContext(**kwargs) + context = RuntimeContext(service_context=self.service_context, **kwargs) try: self.print_flow() flow_op: BaseOp = self._build_flow() assert self.flow_op.async_mode, "Async call requires an async flow operation." - await flow_op.call(context=context) - result = context.stream_queue if self.stream else context.response if self.stream: await context.add_stream_done() + return context.stream_queue + + else: + self._maybe_save_cache(kwargs, context.response) + return context.response - self._maybe_save_cache(kwargs, result) - return result except Exception as e: logger.exception(f"[{self.__class__.__name__}] {self.name} async call failed: {e}") if self.raise_exception: raise e + if self.stream: await context.add_stream_chunk_and_type(str(e), ChunkEnum.ERROR) await context.add_stream_done() return context.stream_queue - context.add_response_error(e) - return context.response + + else: + context.add_response_error(e) + return context.response def call_sync(self, **kwargs) -> Response: """Execute the flow synchronously with parameter caching.""" @@ -204,18 +187,20 @@ class BaseFlow(ABC): if cached := self._maybe_load_cached(kwargs): return cached - context = RuntimeContext(**kwargs) + context = RuntimeContext(service_context=self.service_context, **kwargs) try: self.print_flow() flow_op: BaseOp = self._build_flow() assert not self.flow_op.async_mode, "Sync call requires a sync flow operation." - flow_op.call_sync(context=context) + self._maybe_save_cache(kwargs, context.response) return context.response + except Exception as e: logger.exception(f"[{self.__class__.__name__}] {self.name} sync call failed: {e}") if self.raise_exception: raise e + context.add_response_error(e) return context.response diff --git a/reme_ai/core/flow/cmd_flow.py b/reme/core/flow/cmd_flow.py similarity index 100% rename from reme_ai/core/flow/cmd_flow.py rename to reme/core/flow/cmd_flow.py diff --git a/reme_ai/core/flow/expression_flow.py b/reme/core/flow/expression_flow.py similarity index 76% rename from reme_ai/core/flow/expression_flow.py rename to reme/core/flow/expression_flow.py index 5ac8f0d8..b32c9257 100644 --- a/reme_ai/core/flow/expression_flow.py +++ b/reme/core/flow/expression_flow.py @@ -1,6 +1,7 @@ """Expression-based flow implementation driven by configuration objects.""" from .base_flow import BaseFlow +from ..context import ServiceContext from ..op import BaseOp from ..schema import FlowConfig, ToolCall @@ -8,7 +9,7 @@ from ..schema import FlowConfig, ToolCall class ExpressionFlow(BaseFlow): """A flow implementation that constructs operations from a FlowConfig definition.""" - def __init__(self, flow_config: FlowConfig): + def __init__(self, flow_config: FlowConfig, service_context: ServiceContext): """Initialize the flow using settings and metadata from a FlowConfig instance.""" self.flow_config: FlowConfig = flow_config super().__init__( @@ -18,6 +19,7 @@ class ExpressionFlow(BaseFlow): enable_cache=self.flow_config.enable_cache, cache_path=self.flow_config.cache_path, cache_expire_hours=self.flow_config.cache_expire_hours, + service_context=service_context, **flow_config.model_extra, ) @@ -27,4 +29,9 @@ class ExpressionFlow(BaseFlow): def _build_tool_call(self) -> ToolCall: """Construct a tool call representation based on configuration parameters.""" - return ToolCall(**{"description": self.flow_config.description, "parameters": self.flow_config.parameters}) + return ToolCall( + **{ + "description": self.flow_config.description, + "parameters": self.flow_config.parameters, + }, + ) diff --git a/reme_ai/core/llm/__init__.py b/reme/core/llm/__init__.py similarity index 60% rename from reme_ai/core/llm/__init__.py rename to reme/core/llm/__init__.py index 57578f49..1b80641e 100644 --- a/reme_ai/core/llm/__init__.py +++ b/reme/core/llm/__init__.py @@ -5,6 +5,7 @@ from .lite_llm import LiteLLM from .lite_llm_sync import LiteLLMSync from .openai_llm import OpenAILLM from .openai_llm_sync import OpenAILLMSync +from ..context import R __all__ = [ "BaseLLM", @@ -13,3 +14,8 @@ __all__ = [ "OpenAILLM", "OpenAILLMSync", ] + +R.llm.register("litellm")(LiteLLM) +R.llm.register("litellm_sync")(LiteLLMSync) +R.llm.register("openai")(OpenAILLM) +R.llm.register("openai_sync")(OpenAILLMSync) diff --git a/reme_ai/core/llm/base_llm.py b/reme/core/llm/base_llm.py similarity index 64% rename from reme_ai/core/llm/base_llm.py rename to reme/core/llm/base_llm.py index f83a39a7..be7b5e0d 100644 --- a/reme_ai/core/llm/base_llm.py +++ b/reme/core/llm/base_llm.py @@ -1,4 +1,4 @@ -"""Abstract base interface for ReMe LLM implementations.""" +"""Base interface for LLM implementations.""" import asyncio import json @@ -15,37 +15,42 @@ from ..schema import ToolCall class BaseLLM(ABC): - """Abstract base class defining the standard interface for LLM interactions.""" + """Base class for LLM interactions.""" + + def __init__( + self, + model_name: str, + max_retries: int = 10, + raise_exception: bool = False, + request_interval: float = 0.0, + **kwargs, + ): + """Initialize LLM client. - def __init__(self, model_name: str, max_retries: int = 10, raise_exception: bool = False, max_concurrency: int | None = None, **kwargs): - """Initialize the LLM client with model configurations and retry policies. - Args: - model_name: The name of the model to use - max_retries: Maximum number of retry attempts on failure - raise_exception: Whether to raise exceptions or return default values - max_concurrency: Maximum concurrent requests for async operations. If None, no concurrency limit is applied. + model_name: Model name to use + max_retries: Maximum retry attempts on failure + raise_exception: Raise exceptions or return default values + request_interval: Minimum seconds between requests (default: 0.0) **kwargs: Additional model-specific parameters """ self.model_name: str = model_name self.max_retries: int = max_retries self.raise_exception: bool = raise_exception - self.max_concurrency: int | None = max_concurrency + self.request_interval: float = request_interval self.kwargs: dict = kwargs - - # Concurrency control for async operations - self._semaphore: asyncio.Semaphore | None = asyncio.Semaphore(max_concurrency) if max_concurrency else None + + self._last_request_time: float = 0.0 + self._request_lock: asyncio.Lock = asyncio.Lock() @staticmethod def _accumulate_tool_call_chunk(tool_call, ret_tools: list[ToolCall]): - """Assemble incremental tool call fragments into complete ToolCall objects.""" + """Assemble incremental tool call chunks into complete ToolCall objects.""" index = tool_call.index - # Ensure we have a ToolCall object at this index while len(ret_tools) <= index: ret_tools.append(ToolCall(index=index)) - # Accumulate tool call parts (id, name, arguments) if tool_call.id: ret_tools[index].id += tool_call.id @@ -57,7 +62,7 @@ class BaseLLM(ABC): @staticmethod def _validate_and_serialize_tools(ret_tool_calls: list[ToolCall], tools: list[ToolCall]) -> list[dict]: - """Validate tool call integrity and return serialized tool dictionaries.""" + """Validate and serialize tool calls.""" if not ret_tool_calls: return [] @@ -68,8 +73,9 @@ class BaseLLM(ABC): if tool.name not in tool_dict: continue - if not tool.check_argument(): - raise ValueError(f"Tool call {tool.name} has invalid JSON arguments: {tool.arguments}") + if not tool.sanitize_and_check_argument(): + logger.error(f"Invalid JSON arguments in {tool.name}: {tool.arguments}") + raise ValueError(f"Invalid JSON arguments in {tool.name}: {tool.arguments}") validated_tools.append(tool.simple_output_dump()) return validated_tools @@ -83,15 +89,7 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> dict: - """Construct provider-specific parameters for streaming API requests. - - Args: - messages: List of conversation messages - tools: Optional list of tool calls - log_params: Whether to log parameters - model_name: Optional model name to override self.model_name - **kwargs: Additional parameters - """ + """Build provider-specific streaming parameters.""" async def _stream_chat( self, @@ -99,7 +97,7 @@ class BaseLLM(ABC): tools: list[ToolCall] | None, stream_kwargs: dict, ) -> AsyncGenerator[StreamChunk, None]: - """Internal async generator for streaming raw response chunks.""" + """Async generator for streaming response chunks.""" raise NotImplementedError def _stream_chat_sync( @@ -108,7 +106,7 @@ class BaseLLM(ABC): tools: list[ToolCall] | None = None, stream_kwargs: dict | None = None, ) -> Generator[StreamChunk, None, None]: - """Internal synchronous generator for streaming raw response chunks.""" + """Sync generator for streaming response chunks.""" raise NotImplementedError async def stream_chat( @@ -118,23 +116,18 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> AsyncGenerator[StreamChunk, None]: - """Public async interface for streaming chat completions with retries. - - Args: - messages: List of conversation messages - tools: Optional list of tool calls - model_name: Optional model name to override self.model_name - **kwargs: Additional parameters - """ - # Apply concurrency control if configured - if self._semaphore: - async with self._semaphore: - async for chunk in self._stream_chat_impl(messages, tools, model_name, **kwargs): - yield chunk - else: - async for chunk in self._stream_chat_impl(messages, tools, model_name, **kwargs): - yield chunk - + """Stream chat completions with retries.""" + if self.request_interval > 0: + async with self._request_lock: + current_time = time.time() + elapsed = current_time - self._last_request_time + if elapsed < self.request_interval: + await asyncio.sleep(self.request_interval - elapsed) + self._last_request_time = time.time() + + async for chunk in self._stream_chat_impl(messages, tools, model_name, **kwargs): + yield chunk + async def _stream_chat_impl( self, messages: list[Message], @@ -142,7 +135,7 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> AsyncGenerator[StreamChunk, None]: - """Internal implementation of stream_chat with retry logic.""" + """Stream chat with retry logic.""" stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs) for i in range(self.max_retries): @@ -152,7 +145,7 @@ class BaseLLM(ABC): return except Exception as e: - logger.exception(f"stream chat with model={self.model_name} encounter error with e={e.args}") + logger.exception(f"Stream chat error (model={self.model_name}): {e.args}") if i == self.max_retries - 1: if self.raise_exception: @@ -170,14 +163,7 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> Generator[StreamChunk, None, None]: - """Public synchronous interface for streaming chat completions with retries. - - Args: - messages: List of conversation messages - tools: Optional list of tool calls - model_name: Optional model name to override self.model_name - **kwargs: Additional parameters - """ + """Stream chat completions synchronously with retries.""" stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs) for i in range(self.max_retries): @@ -186,7 +172,7 @@ class BaseLLM(ABC): return except Exception as e: - logger.exception(f"stream chat sync with model={self.model_name} encounter error with e={e.args}") + logger.exception(f"Stream chat sync error (model={self.model_name}): {e.args}") if i == self.max_retries - 1: if self.raise_exception: @@ -205,15 +191,7 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> Message: - """Internal async method to aggregate a full response by consuming the stream. - - Args: - messages: List of conversation messages - tools: Optional list of tool calls - enable_stream_print: Whether to print stream chunks - model_name: Optional model name to override self.model_name - **kwargs: Additional parameters - """ + """Aggregate full response by consuming the stream.""" state = { "enter_think": False, "enter_answer": False, @@ -224,7 +202,6 @@ class BaseLLM(ABC): stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs) async for stream_chunk in self._stream_chat(messages=messages, tools=tools, stream_kwargs=stream_kwargs): - # Process stream chunk if stream_chunk.chunk_type is ChunkEnum.USAGE: if enable_stream_print: print( @@ -273,15 +250,7 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> Message: - """Internal synchronous method to aggregate a full response by consuming the stream. - - Args: - messages: List of conversation messages - tools: Optional list of tool calls - enable_stream_print: Whether to print stream chunks - model_name: Optional model name to override self.model_name - **kwargs: Additional parameters - """ + """Aggregate full response synchronously by consuming the stream.""" state = { "enter_think": False, "enter_answer": False, @@ -292,7 +261,6 @@ class BaseLLM(ABC): stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs) for stream_chunk in self._stream_chat_sync(messages=messages, tools=tools, stream_kwargs=stream_kwargs): - # Process stream chunk if stream_chunk.chunk_type is ChunkEnum.USAGE: if enable_stream_print: print( @@ -343,24 +311,25 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> Message | Any: - """Perform an async chat completion with integrated retries and error handling. - - Args: - messages: List of conversation messages - tools: Optional list of tool calls - enable_stream_print: Whether to print stream chunks - callback_fn: Optional callback function to process the result - default_value: Default value to return on error - model_name: Optional model name to override self.model_name - **kwargs: Additional parameters - """ - # Apply concurrency control if configured - if self._semaphore: - async with self._semaphore: - return await self._chat_impl(messages, tools, enable_stream_print, callback_fn, default_value, model_name, **kwargs) - else: - return await self._chat_impl(messages, tools, enable_stream_print, callback_fn, default_value, model_name, **kwargs) - + """Chat completion with retries and error handling.""" + if self.request_interval > 0: + async with self._request_lock: + current_time = time.time() + elapsed = current_time - self._last_request_time + if elapsed < self.request_interval: + await asyncio.sleep(self.request_interval - elapsed) + self._last_request_time = time.time() + + return await self._chat_impl( + messages, + tools, + enable_stream_print, + callback_fn, + default_value, + model_name, + **kwargs, + ) + async def _chat_impl( self, messages: list[Message], @@ -371,10 +340,9 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> Message | Any: - """Internal implementation of chat with retry and error handling logic.""" - # Use the provided model_name or fall back to self.model_name + """Chat with retry and error handling logic.""" effective_model = model_name if model_name is not None else self.model_name - + for i in range(self.max_retries): try: result = await self._chat( @@ -387,15 +355,16 @@ class BaseLLM(ABC): return callback_fn(result) if callback_fn else result except Exception as e: - # Check if this is an inappropriate content error error_message = str(e.args[0]) if e.args else str(e) is_inappropriate_content = "inappropriate content" in error_message.lower() - is_rate_limit_error = "request rate increased too quickly" in error_message.lower() - + is_rate_limit_error = ( + "request rate increased too quickly" in error_message.lower() + or "exceeded your current quota" in error_message.lower() + or "insufficient_quota" in error_message.lower() + ) + if is_inappropriate_content: - logger.error(f"chat with model={effective_model} detected inappropriate content error") - logger.error("=" * 80) - logger.error("Full message content that triggered the error:") + logger.error(f"Inappropriate content detected (model={effective_model})") logger.error("=" * 80) for idx, msg in enumerate(messages): logger.error(f"Message {idx + 1} [role={msg.role}]:") @@ -406,15 +375,16 @@ class BaseLLM(ABC): logger.error(f"Tool calls: {msg.tool_calls}") logger.error("-" * 80) logger.error("=" * 80) - # Return empty Message immediately without retrying return Message(role=Role.ASSISTANT, content="") - + if is_rate_limit_error: - logger.warning(f"chat with model={effective_model} hit rate limit, sleeping for 60s before retry (attempt {i + 1}/{self.max_retries})") + logger.warning( + f"Rate limit hit (model={effective_model}), sleeping 60s (attempt {i + 1}/{self.max_retries})", + ) await asyncio.sleep(60) continue - - logger.exception(f"chat with model={effective_model} encounter error with e={e.args}") + + logger.exception(f"Chat error (model={effective_model}): {e.args}") if i == self.max_retries - 1: if self.raise_exception: @@ -434,20 +404,9 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> Message | Any: - """Perform a synchronous chat completion with integrated retries and error handling. - - Args: - messages: List of conversation messages - tools: Optional list of tool calls - enable_stream_print: Whether to print stream chunks - callback_fn: Optional callback function to process the result - default_value: Default value to return on error - model_name: Optional model name to override self.model_name - **kwargs: Additional parameters - """ - # Use the provided model_name or fall back to self.model_name + """Chat completion synchronously with retries and error handling.""" effective_model = model_name if model_name is not None else self.model_name - + for i in range(self.max_retries): try: result = self._chat_sync( @@ -460,15 +419,16 @@ class BaseLLM(ABC): return callback_fn(result) if callback_fn else result except Exception as e: - # Check if this is an inappropriate content error error_message = str(e.args[0]) if e.args else str(e) is_inappropriate_content = "inappropriate content" in error_message.lower() - is_rate_limit_error = "request rate increased too quickly" in error_message.lower() - + is_rate_limit_error = ( + "request rate increased too quickly" in error_message.lower() + or "exceeded your current quota" in error_message.lower() + or "insufficient_quota" in error_message.lower() + ) + if is_inappropriate_content: - logger.error(f"chat sync with model={effective_model} detected inappropriate content error") - logger.error("=" * 80) - logger.error("Full message content that triggered the error:") + logger.error(f"Inappropriate content detected (model={effective_model})") logger.error("=" * 80) for idx, msg in enumerate(messages): logger.error(f"Message {idx + 1} [role={msg.role}]:") @@ -479,15 +439,16 @@ class BaseLLM(ABC): logger.error(f"Tool calls: {msg.tool_calls}") logger.error("-" * 80) logger.error("=" * 80) - # Return empty Message immediately without retrying return Message(role=Role.ASSISTANT, content="") - + if is_rate_limit_error: - logger.warning(f"chat sync with model={effective_model} hit rate limit, sleeping for 60s before retry (attempt {i + 1}/{self.max_retries})") + logger.warning( + f"Rate limit hit (model={effective_model}), sleeping 60s (attempt {i + 1}/{self.max_retries})", + ) time.sleep(60) continue - - logger.exception(f"chat sync with model={effective_model} encounter error with e={e.args}") + + logger.exception(f"Chat sync error (model={effective_model}): {e.args}") if i == self.max_retries - 1: if self.raise_exception: @@ -498,7 +459,7 @@ class BaseLLM(ABC): return default_value async def close(self): - """Release any asynchronous resources or connections held by the client.""" + """Release async resources.""" def close_sync(self): - """Release any synchronous resources or connections held by the client.""" + """Release sync resources.""" diff --git a/reme_ai/core/llm/lite_llm.py b/reme/core/llm/lite_llm.py similarity index 97% rename from reme_ai/core/llm/lite_llm.py rename to reme/core/llm/lite_llm.py index 88177184..5663702b 100644 --- a/reme_ai/core/llm/lite_llm.py +++ b/reme/core/llm/lite_llm.py @@ -7,14 +7,12 @@ import litellm from loguru import logger from .base_llm import BaseLLM -from ..context import C from ..enumeration import ChunkEnum from ..schema import Message from ..schema import StreamChunk from ..schema import ToolCall -@C.register_llm("litellm") class LiteLLM(BaseLLM): """Async LLM implementation using LiteLLM to support multiple providers.""" @@ -40,7 +38,7 @@ class LiteLLM(BaseLLM): **kwargs, ) -> dict: """Construct and log the parameters dictionary for LiteLLM API calls. - + Args: messages: List of conversation messages tools: Optional list of tool calls @@ -50,7 +48,7 @@ class LiteLLM(BaseLLM): """ # Use the provided model_name or fall back to self.model_name effective_model = model_name if model_name is not None else self.model_name - + # Construct the API parameters by merging multiple sources llm_kwargs = { "model": effective_model, @@ -98,7 +96,7 @@ class LiteLLM(BaseLLM): if not chunk.choices: if hasattr(chunk, "usage") and chunk.usage: yield StreamChunk(chunk_type=ChunkEnum.USAGE, chunk=chunk.usage.model_dump()) - continue + continue delta = chunk.choices[0].delta diff --git a/reme_ai/core/llm/lite_llm_sync.py b/reme/core/llm/lite_llm_sync.py similarity index 95% rename from reme_ai/core/llm/lite_llm_sync.py rename to reme/core/llm/lite_llm_sync.py index c3a925d2..778eaed0 100644 --- a/reme_ai/core/llm/lite_llm_sync.py +++ b/reme/core/llm/lite_llm_sync.py @@ -5,14 +5,12 @@ from typing import Generator import litellm from .lite_llm import LiteLLM -from ..context import C from ..enumeration import ChunkEnum from ..schema import Message from ..schema import StreamChunk from ..schema import ToolCall -@C.register_llm("litellm_sync") class LiteLLMSync(LiteLLM): """Synchronous LiteLLM client for executing chat completions and streaming responses.""" @@ -31,7 +29,7 @@ class LiteLLMSync(LiteLLM): if not chunk.choices: if hasattr(chunk, "usage") and chunk.usage: yield StreamChunk(chunk_type=ChunkEnum.USAGE, chunk=chunk.usage.model_dump()) - continue + continue delta = chunk.choices[0].delta diff --git a/reme_ai/core/llm/openai_llm.py b/reme/core/llm/openai_llm.py similarity index 97% rename from reme_ai/core/llm/openai_llm.py rename to reme/core/llm/openai_llm.py index 9e9ffc4a..ea647e07 100644 --- a/reme_ai/core/llm/openai_llm.py +++ b/reme/core/llm/openai_llm.py @@ -7,14 +7,12 @@ from loguru import logger from openai import AsyncOpenAI from .base_llm import BaseLLM -from ..context import C from ..enumeration import ChunkEnum from ..schema import Message from ..schema import StreamChunk from ..schema import ToolCall -@C.register_llm("openai") class OpenAILLM(BaseLLM): """Asynchronous LLM client for OpenAI-compatible APIs supporting streaming completions and tool execution.""" @@ -45,7 +43,7 @@ class OpenAILLM(BaseLLM): **kwargs, ) -> dict: """Construct the parameter dictionary for the OpenAI Chat Completions API call. - + Args: messages: List of conversation messages tools: Optional list of tool calls @@ -55,7 +53,7 @@ class OpenAILLM(BaseLLM): """ # Use the provided model_name or fall back to self.model_name effective_model = model_name if model_name is not None else self.model_name - + # Construct the API parameters by merging multiple sources llm_kwargs = { "model": effective_model, @@ -93,7 +91,7 @@ class OpenAILLM(BaseLLM): if not chunk.choices: if hasattr(chunk, "usage") and chunk.usage: yield StreamChunk(chunk_type=ChunkEnum.USAGE, chunk=chunk.usage.model_dump()) - continue + continue delta = chunk.choices[0].delta diff --git a/reme_ai/core/llm/openai_llm_sync.py b/reme/core/llm/openai_llm_sync.py similarity index 96% rename from reme_ai/core/llm/openai_llm_sync.py rename to reme/core/llm/openai_llm_sync.py index 51da29f4..86dd119d 100644 --- a/reme_ai/core/llm/openai_llm_sync.py +++ b/reme/core/llm/openai_llm_sync.py @@ -5,14 +5,12 @@ from typing import Generator from openai import OpenAI from .openai_llm import OpenAILLM -from ..context import C from ..enumeration import ChunkEnum from ..schema import Message from ..schema import StreamChunk from ..schema import ToolCall -@C.register_llm("openai_sync") class OpenAILLMSync(OpenAILLM): """Synchronous LLM client for OpenAI-compatible APIs, inheriting from OpenAILLM.""" @@ -35,7 +33,7 @@ class OpenAILLMSync(OpenAILLM): if not chunk.choices: if hasattr(chunk, "usage") and chunk.usage: yield StreamChunk(chunk_type=ChunkEnum.USAGE, chunk=chunk.usage.model_dump()) - continue + continue delta = chunk.choices[0].delta diff --git a/reme_ai/core/op/__init__.py b/reme/core/op/__init__.py similarity index 72% rename from reme_ai/core/op/__init__.py rename to reme/core/op/__init__.py index a48bf276..53a6f379 100644 --- a/reme_ai/core/op/__init__.py +++ b/reme/core/op/__init__.py @@ -2,14 +2,19 @@ from .base_op import BaseOp from .base_ray_op import BaseRayOp +from .base_tool import BaseTool from .mcp_tool import MCPTool from .parallel_op import ParallelOp from .sequential_op import SequentialOp +from ..context import R __all__ = [ "BaseOp", "BaseRayOp", + "BaseTool", "MCPTool", "ParallelOp", "SequentialOp", ] + +R.op.register("mcp_tool")(MCPTool) diff --git a/reme_ai/core/op/base_op.py b/reme/core/op/base_op.py similarity index 61% rename from reme_ai/core/op/base_op.py rename to reme/core/op/base_op.py index 357db44e..ef469d0f 100644 --- a/reme_ai/core/op/base_op.py +++ b/reme/core/op/base_op.py @@ -3,22 +3,23 @@ import asyncio import copy import inspect +from abc import ABCMeta from pathlib import Path -from typing import Callable, Optional +from typing import Callable, Optional, Any from loguru import logger from tqdm import tqdm -from ..context import RuntimeContext, PromptHandler, C +from ..context import RuntimeContext, PromptHandler, ServiceContext from ..embedding import BaseEmbeddingModel from ..llm import BaseLLM -from ..schema import ToolCall, ToolAttr, Response +from ..schema import Response from ..token_counter import BaseTokenCounter from ..utils import camel_to_snake, CacheHandler, timer from ..vector_store import BaseVectorStore -class BaseOp: +class BaseOp(metaclass=ABCMeta): """Base operator class for LLM workflow execution and composition.""" def __new__(cls, *args, **kwargs): @@ -34,6 +35,7 @@ class BaseOp: async_mode: bool = True, language: str = "", prompt_name: str = "", + prompt_path: str = "", llm: str | BaseLLM = "default", embedding_model: str | BaseEmbeddingModel = "default", vector_store: str | BaseVectorStore = "default", @@ -44,7 +46,6 @@ class BaseOp: sub_ops: dict[str, "BaseOp"] | list["BaseOp"] | Optional["BaseOp"] = None, input_mapping: dict[str, str] | None = None, output_mapping: dict[str, str] | None = None, - save_response_result: bool = False, enable_sync_thread_pool: bool = True, max_retries: int = 1, raise_exception: bool = False, @@ -53,8 +54,8 @@ class BaseOp: """Initialize operator configurations and internal state.""" self.name = name or camel_to_snake(self.__class__.__name__) self.async_mode = async_mode - self.language = language or C.language - self.prompt = self._get_prompt_handler(prompt_name) + self.language = language + self.prompt = self._get_prompt_handler(prompt_name, prompt_path) self._llm = llm self._embedding_model = embedding_model @@ -64,12 +65,12 @@ class BaseOp: self.enable_cache = enable_cache self.cache_path = cache_path self.cache_expire_hours = cache_expire_hours - self.sub_ops: list[BaseOp] = [] + + self.sub_ops: list["BaseOp"] = [] self.add_sub_ops(sub_ops) self.input_mapping = input_mapping self.output_mapping = output_mapping - self.save_response_result = save_response_result self.enable_sync_thread_pool = enable_sync_thread_pool self.max_retries = max(1, max_retries) self.raise_exception = raise_exception @@ -78,86 +79,29 @@ class BaseOp: self._pending_tasks: list = [] self.context: RuntimeContext | None = None self._cache: CacheHandler | None = None - self._tool_call: ToolCall | None = None - def _get_prompt_handler(self, prompt_name: str) -> PromptHandler: + def _get_prompt_handler(self, prompt_name: str, prompt_path: str) -> PromptHandler: """Load prompt configuration from the associated YAML file.""" - path = Path(inspect.getfile(self.__class__)) - path = path.with_stem(prompt_name) if prompt_name else path + if prompt_path: + path = Path(prompt_path) + else: + path = Path(inspect.getfile(self.__class__)) + if prompt_name: + path = path.with_stem(prompt_name) return PromptHandler(language=self.language).load_prompt_by_file(path.with_suffix(".yaml")) - def _build_tool_call(self) -> ToolCall | None: - """Build and return the tool call schema; override in subclasses.""" - - def _validate_inputs(self): - """Ensure all required tool inputs are present in context.""" - if self.tool_call is not None: - parameters = self.tool_call.parameters - if parameters.type == "object" and parameters.properties: - required_list = parameters.required or [] - required_keys = {k: (k in required_list) for k in parameters.properties.keys()} - self.context.validate_required_keys(required_keys, self.name) - - def _handle_failure(self, e: Exception, attempt: int): + def _handle_failure(self, e: Exception, attempt: int) -> str | None: """Log failures and handle final retry logic.""" message = f"[{self.__class__.__name__}] {self.name} failed (attempt {attempt + 1}): {e}" if attempt == self.max_retries - 1: logger.exception(message) if self.raise_exception: raise e - - if self.tool_call is not None: - self.output = f"{self.name} failed: {e}" + return f"{self.name} failed: {e}" else: logger.warning(message) - - @property - def tool_call(self) -> ToolCall | None: - """Lazily construct and return the tool call metadata.""" - if self._tool_call is None: - self._tool_call = self._build_tool_call() - if self._tool_call is None: - return None - - self._tool_call.name = self._tool_call.name or self.name - if not self._tool_call.output.properties: - self._tool_call.output.properties = { - f"{self.name}_result": ToolAttr(type="string", description=f"Execution result of {self.name}"), - } - return self._tool_call - - @property - def input_dict(self) -> dict: - """Extract required and optional inputs from context based on schema.""" - parameters = self.tool_call.parameters - if parameters.type != "object" or not parameters.properties: - return {} - required_keys = set(parameters.required or []) - return {k: self.context[k] for k in parameters.properties.keys() if (k in required_keys or k in self.context)} - - @property - def output(self): - """Get the single output value from context.""" - output_properties = self.tool_call.output.properties - if not output_properties: return None - keys = list(output_properties.keys()) - if len(keys) >= 1 and keys[0] in self.context: - return self.context[keys[0]] - else: - return None - - @output.setter - def output(self, value): - """Set the single output value into context.""" - output_properties = self.tool_call.output.properties - if not output_properties: - return - - keys = list(output_properties.keys()) - self.context[keys[0]] = value - @property def cache(self) -> CacheHandler: """Access the operator-specific cache handler.""" @@ -166,119 +110,120 @@ class BaseOp: self._cache = CacheHandler(f"{self.cache_path}/{self.name}") return self._cache + @property + def service_context(self) -> ServiceContext: + """Access the service context.""" + return self.context.service_context + @property def llm(self) -> BaseLLM: """Get the LLM instance from ServiceContext.""" if isinstance(self._llm, str): - self._llm = C.get_llm(self._llm) + self._llm = self.service_context.llms[self._llm] return self._llm @property def embedding_model(self) -> BaseEmbeddingModel: """Get the embedding model instance from ServiceContext.""" if isinstance(self._embedding_model, str): - self._embedding_model = C.get_embedding_model(self._embedding_model) + self._embedding_model = self.service_context.embedding_models[self._embedding_model] return self._embedding_model @property def vector_store(self) -> BaseVectorStore: """Lazily initialize and return the vector store instance.""" if isinstance(self._vector_store, str): - self._vector_store = C.get_vector_store(self._vector_store) + self._vector_store = self.service_context.vector_stores[self._vector_store] return self._vector_store @property def token_counter(self) -> BaseTokenCounter: """Get the token counter instance from ServiceContext.""" if isinstance(self._token_counter, str): - self._token_counter = C.get_token_counter(self._token_counter) + self._token_counter = self.service_context.token_counters[self._token_counter] return self._token_counter @property def service_metadata(self) -> dict: """Get service configuration metadata.""" - return C.service_config.model_extra + return self.service_context.service_config.model_extra @property def response(self) -> Response: - """Get the response object.""" + """Access the response object.""" return self.context.response - def set_tool_call(self, tool_call: ToolCall | dict): - """Set the tool call.""" - if isinstance(tool_call, dict): - self._tool_call = ToolCall(**tool_call) - elif isinstance(tool_call, ToolCall): - self._tool_call = tool_call - else: - raise ValueError(f"Invalid tool call: {tool_call}") - - self._tool_call.name = self._tool_call.name or self.name - if not self._tool_call.output.properties: - self._tool_call.output.properties = { - f"{self.name}_result": ToolAttr(type="string", description=f"Execution result of {self.name}"), - } - - def set_language(self, language: str): - """Set the language.""" - self.language = language - return self - def before_execute_sync(self): """Prepare context and validate before sync execution.""" self.context.apply_mapping(self.input_mapping) - self._validate_inputs() + + async def before_execute(self): + """Prepare context and validate before async execution.""" + self.context.apply_mapping(self.input_mapping) def execute_sync(self): """Define core sync logic in subclasses.""" - def after_execute_sync(self): - """Finalize context and mappings after sync execution.""" - self.context.apply_mapping(self.output_mapping) - if self.tool_call is not None and self.save_response_result: - self.context.response.answer = self.output - - async def before_execute(self): - """Prepare context and validate before async execution.""" - self.before_execute_sync() - async def execute(self): """Define core async logic in subclasses.""" - async def after_execute(self): + def after_execute_sync(self, response: Any): + """Finalize context and mappings after sync execution.""" + self.context.apply_mapping(self.output_mapping) + if response is not None: + if isinstance(response, dict): + for k, v in response.items(): + if k == "answer": + self.response.answer = v + elif k == "success": + self.response.success = v.lower() == "true" + else: + self.response.metadata[k] = v + else: + self.response.answer = response + return response + + async def after_execute(self, output: Any): """Finalize context and mappings after async execution.""" - self.after_execute_sync() + return self.after_execute_sync(output) @timer def call_sync(self, context: RuntimeContext = None, **kwargs): """Execute the operator synchronously with retry logic.""" self.context = RuntimeContext.from_context(context, **kwargs) + response = None for i in range(self.max_retries): try: self.before_execute_sync() - self.execute_sync() - self.after_execute_sync() + response = self.execute_sync() + response = self.after_execute_sync(response) break except Exception as e: - self._handle_failure(e, i) - return self.output if self.tool_call is not None else None + response = self._handle_failure(e, i) + return response + + @timer async def call(self, context: RuntimeContext = None, **kwargs): """Execute the operator asynchronously with retry logic.""" self.context = RuntimeContext.from_context(context, **kwargs) + response = None for i in range(self.max_retries): try: await self.before_execute() - await self.execute() - await self.after_execute() + response = await self.execute() + response = await self.after_execute(response) break except Exception as e: - self._handle_failure(e, i) - return self.output if self.tool_call is not None else None + response = self._handle_failure(e, i) + return response def submit_sync_task(self, fn: Callable, *args, **kwargs) -> "BaseOp": """Submit a task to the thread pool or local queue.""" - task = C.thread_pool.submit(fn, *args, **kwargs) if self.enable_sync_thread_pool else (fn, args, kwargs) + if self.enable_sync_thread_pool: + task = self.service_context.thread_pool.submit(fn, *args, **kwargs) + else: + task = (fn, args, kwargs) self._pending_tasks.append(task) return self @@ -292,26 +237,32 @@ class BaseOp: """Wait for all pending sync tasks and return flattened results.""" results = [] for task in tqdm(self._pending_tasks, desc=task_desc or self.name): - res = task.result() if self.enable_sync_thread_pool else task[0](*task[1], **task[2]) - if res: - results.extend(res if isinstance(res, list) else [res]) + if self.enable_sync_thread_pool: + result = task.result() + else: + result = task[0](*task[1], **task[2]) + if result: + if isinstance(result, list): + results.extend(result) + else: + results.append(result) self._pending_tasks.clear() return results async def join_async_tasks(self, return_exceptions: bool = True) -> list: """Wait for all pending async tasks and aggregate results.""" - try: - raw_results = await asyncio.gather(*self._pending_tasks, return_exceptions=return_exceptions) - results = [] - for res in raw_results: - if isinstance(res, Exception): - logger.error(f"[{self.__class__.__name__}] Async task failed: {res}") - continue - if res: - results.extend(res if isinstance(res, list) else [res]) - return results - finally: - self._pending_tasks.clear() + raw_results = await asyncio.gather(*self._pending_tasks, return_exceptions=return_exceptions) + results = [] + for result in raw_results: + if isinstance(result, Exception): + logger.error(f"[{self.__class__.__name__}] Async task failed: {result}") + elif result: + if isinstance(result, list): + results.extend(result) + else: + result.append(result) + self._pending_tasks.clear() + return results def add_sub_ops(self, sub_ops: dict[str, "BaseOp"] | list["BaseOp"] | Optional["BaseOp"]): """Add child operators to this operator's sub_ops.""" diff --git a/reme_ai/core/op/base_ray_op.py b/reme/core/op/base_ray_op.py similarity index 94% rename from reme_ai/core/op/base_ray_op.py rename to reme/core/op/base_ray_op.py index 7a8fe113..84c0ed25 100644 --- a/reme_ai/core/op/base_ray_op.py +++ b/reme/core/op/base_ray_op.py @@ -8,14 +8,14 @@ from loguru import logger from tqdm import tqdm from .base_op import BaseOp -from ..context import BaseContext, C +from ..context import BaseContext _RAY_IMPORT_ERROR = None try: import ray -except ImportError as e: - _RAY_IMPORT_ERROR = e +except ImportError as _e: + _RAY_IMPORT_ERROR = _e ray = None @@ -35,7 +35,7 @@ class BaseRayOp(BaseOp, metaclass=ABCMeta): def submit_and_join_ray_task(self, fn: Callable, parallel_key: str = "", task_desc: str = "", **kwargs) -> list: """Divide data into chunks and execute them across Ray workers.""" - max_workers = C.service_config.ray_max_workers + max_workers = self.service_context.ray_max_workers self._ray_task_list.clear() # Automatically detect the key containing the list to parallelize @@ -94,7 +94,7 @@ class BaseRayOp(BaseOp, metaclass=ABCMeta): def submit_ray_task(self, fn, *args, **kwargs): """Submit a single Ray task to the task list for later execution.""" if not ray.is_initialized(): - ray.init(num_cpus=C.service_config.ray_max_workers, ignore_reinit_error=True) + ray.init(num_cpus=self.service_context.ray_max_workers, ignore_reinit_error=True) remote_fn = ray.remote(fn) task = remote_fn.remote(*args, **kwargs) diff --git a/reme/core/op/base_tool.py b/reme/core/op/base_tool.py new file mode 100644 index 00000000..5b41b1fb --- /dev/null +++ b/reme/core/op/base_tool.py @@ -0,0 +1,58 @@ +"""Base class for tools""" + +from abc import ABCMeta + +from . import BaseOp +from ..schema import ToolCall + + +class BaseTool(BaseOp, metaclass=ABCMeta): + """Base class for tools""" + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self._tool_call: ToolCall | None = None + + def _build_tool_call(self) -> ToolCall: + """Build and return the tool call schema; override in subclasses.""" + + def _validate_inputs(self): + """Validate the inputs.""" + parameters = self.tool_call.parameters + if parameters.type == "object" and parameters.properties: + required_list = parameters.required or [] + required_keys = {k: (k in required_list) for k in parameters.properties.keys()} + self.context.validate_required_keys(required_keys, self.name) + + @property + def tool_call(self) -> ToolCall: + """Get the tool call schema.""" + if self._tool_call is None: + self._tool_call = self._build_tool_call() + self._tool_call.name = self._tool_call.name or self.name + return self._tool_call + + def set_tool_call(self, tool_call: ToolCall | dict): + """Set the tool call schema.""" + if isinstance(tool_call, dict): + self._tool_call = ToolCall(**tool_call) + elif isinstance(tool_call, ToolCall): + self._tool_call = tool_call + else: + raise ValueError(f"Invalid tool call: {tool_call}") + + self._tool_call.name = self._tool_call.name or self.name + + @property + def input_dict(self) -> dict: + """Get the input dict.""" + parameters = self.tool_call.parameters + if parameters.type != "object" or not parameters.properties: + return {} + required_keys = set(parameters.required or []) + return {k: self.context[k] for k in parameters.properties.keys() if (k in required_keys or k in self.context)} + + def before_execute_sync(self): + """Hook before execute""" + super().before_execute_sync() + self._validate_inputs() diff --git a/reme_ai/core/op/mcp_tool.py b/reme/core/op/mcp_tool.py similarity index 75% rename from reme_ai/core/op/mcp_tool.py rename to reme/core/op/mcp_tool.py index 57cbcbd4..c59e105c 100644 --- a/reme_ai/core/op/mcp_tool.py +++ b/reme/core/op/mcp_tool.py @@ -2,25 +2,20 @@ from typing import List -from .base_op import BaseOp -from ..context import C +from mcp.types import CallToolResult, TextContent + +from .base_tool import BaseTool from ..schema import ToolCall from ..utils import MCPClient -@C.register_op() -class MCPTool(BaseOp): - """Operator for calling remote MCP (Model Context Protocol) tools. - - This class enables integration with external MCP servers to execute tools - and retrieve their results. It supports parameter customization and retry logic. - """ +class MCPTool(BaseTool): + """Operator for calling remote MCP (Model Context Protocol) tools.""" def __init__( self, mcp_server: str = "", tool_name: str = "", - save_response_result: bool = True, parameter_required: List[str] | None = None, parameter_optional: List[str] | None = None, parameter_deleted: List[str] | None = None, @@ -29,13 +24,7 @@ class MCPTool(BaseOp): raise_exception: bool = False, **kwargs, ): - - super().__init__( - save_response_result=save_response_result, - max_retries=max_retries, - raise_exception=raise_exception, - **kwargs, - ) + super().__init__(max_retries=max_retries, raise_exception=raise_exception, **kwargs) self.mcp_server: str = mcp_server self.tool_name: str = tool_name @@ -43,12 +32,12 @@ class MCPTool(BaseOp): self.parameter_optional: List[str] | None = parameter_optional self.parameter_deleted: List[str] | None = parameter_deleted self.timeout: float | None = timeout - # Example MCP marketplace: https://bailian.console.aliyun.com/?tab=mcp#/mcp-market - self._client = MCPClient(C.service_config.mcp_servers) + # Example MCP marketplace: https://bailian.console.aliyun.com/?tab=mcp#/mcp-market + self._client = MCPClient(self.service_context.service_config.mcp_servers) def _build_tool_call(self) -> ToolCall: - tool_call_dict = C.mcp_server_mapping[self.mcp_server] + tool_call_dict = self.service_context.mcp_server_mapping[self.mcp_server] tool_call: ToolCall = tool_call_dict[self.tool_name].model_copy(deep=True) # Initialize required list if not exists @@ -74,9 +63,16 @@ class MCPTool(BaseOp): return tool_call async def execute(self): - self.output = await self._client.call_tool( + tool_result: CallToolResult = await self._client.call_tool( server_name=self.mcp_server, tool_name=self.tool_name, arguments=self.input_dict, - parse_text_result=True, ) + self.context.tool_result = tool_result + + text_result = [] + for block in tool_result.content: + if isinstance(block, TextContent): + text_result.append(block.text) + output: str = "\n".join(text_result) + return output diff --git a/reme_ai/core/op/parallel_op.py b/reme/core/op/parallel_op.py similarity index 93% rename from reme_ai/core/op/parallel_op.py rename to reme/core/op/parallel_op.py index 18b84d0c..8bca5790 100644 --- a/reme_ai/core/op/parallel_op.py +++ b/reme/core/op/parallel_op.py @@ -11,14 +11,14 @@ class ParallelOp(BaseOp): for op in self.sub_ops: assert op.async_mode self.submit_async_task(op.call, context=self.context) - await self.join_async_tasks() + return await self.join_async_tasks() def execute_sync(self): """Executes all sub-operations concurrently using synchronous task management.""" for op in self.sub_ops: assert not op.async_mode self.submit_sync_task(op.call_sync, context=self.context) - self.join_sync_tasks() + return self.join_sync_tasks() def __lshift__(self, op: dict[str, BaseOp] | list[BaseOp] | BaseOp): """Raises RuntimeError as the shift operator is not supported for parallel operations.""" diff --git a/reme_ai/core/op/sequential_op.py b/reme/core/op/sequential_op.py similarity index 84% rename from reme_ai/core/op/sequential_op.py rename to reme/core/op/sequential_op.py index 3dabb0c9..fed44dc3 100644 --- a/reme_ai/core/op/sequential_op.py +++ b/reme/core/op/sequential_op.py @@ -8,15 +8,19 @@ class SequentialOp(BaseOp): async def execute(self): """Executes sub-operations sequentially using asynchronous awaits.""" + result = None for op in self.sub_ops: assert op.async_mode - await op.call(context=self.context) + result = await op.call(context=self.context) + return result def execute_sync(self): """Executes sub-operations sequentially in a synchronous blocking manner.""" + result = None for op in self.sub_ops: assert not op.async_mode - op.call_sync(context=self.context) + result = op.call_sync(context=self.context) + return result def __lshift__(self, op: dict[str, BaseOp] | list[BaseOp] | BaseOp): """Raises RuntimeError as the left shift operator is not supported.""" diff --git a/reme_ai/core/schema/__init__.py b/reme/core/schema/__init__.py similarity index 100% rename from reme_ai/core/schema/__init__.py rename to reme/core/schema/__init__.py diff --git a/reme_ai/core/schema/memory_node.py b/reme/core/schema/memory_node.py similarity index 91% rename from reme_ai/core/schema/memory_node.py rename to reme/core/schema/memory_node.py index 304dfe5c..67ed7c43 100644 --- a/reme_ai/core/schema/memory_node.py +++ b/reme/core/schema/memory_node.py @@ -6,7 +6,6 @@ memories in the ReMe system. import datetime import hashlib -import json from typing import Any from pydantic import BaseModel, Field, model_validator @@ -146,30 +145,6 @@ class MemoryNode(BaseModel): metadata=metadata, ) - def format_memory(self) -> str: - """Format memory as human-readable string. - - Returns: - str: Formatted string with when_to_use, content, and ref_memory_id. - """ - parts: list[str] = [ - f"memory_id={self.memory_id}", - ] - - if self.when_to_use: - parts.append(self.when_to_use) - - if self.content: - parts.append(self.content) - - if self.metadata: - parts.append(f"metadata={json.dumps(self.metadata, ensure_ascii=False)}") - - if self.ref_memory_id: - parts.append(f"ref_memory_id={self.ref_memory_id}") - - return " ".join(parts) - @classmethod def from_vector_node(cls, node: VectorNode) -> "MemoryNode": """Reconstruct MemoryNode from VectorNode. diff --git a/reme_ai/core/schema/message.py b/reme/core/schema/message.py similarity index 96% rename from reme_ai/core/schema/message.py rename to reme/core/schema/message.py index 6c3299e7..321dd8c8 100644 --- a/reme_ai/core/schema/message.py +++ b/reme/core/schema/message.py @@ -132,7 +132,7 @@ class Message(BaseModel): def strip_md_func(line): if strip_markdown_headers: - line = re.sub(r'\n##+ +', '\n', line) + line = re.sub(r"\n##+ +", "\n", line) return line if add_reasoning and self.reasoning_content: @@ -143,8 +143,9 @@ class Message(BaseModel): elif isinstance(self.content, list): for block in self.content: - text = block.content if isinstance(block.content, str) else \ - json.dumps(block.content, ensure_ascii=False) + text = ( + block.content if isinstance(block.content, str) else json.dumps(block.content, ensure_ascii=False) + ) text = str(text) lines.append(strip_md_func(text)) diff --git a/reme_ai/core/schema/request.py b/reme/core/schema/request.py similarity index 100% rename from reme_ai/core/schema/request.py rename to reme/core/schema/request.py diff --git a/reme_ai/core/schema/response.py b/reme/core/schema/response.py similarity index 83% rename from reme_ai/core/schema/response.py rename to reme/core/schema/response.py index 3104bc6e..fe753232 100644 --- a/reme_ai/core/schema/response.py +++ b/reme/core/schema/response.py @@ -1,11 +1,13 @@ """Defines the standardized data structure for model output responses.""" +from typing import Any + from pydantic import Field, BaseModel class Response(BaseModel): """Represents a structured response containing the execution result, status, and metadata.""" - answer: str | dict | list = Field(default="") + answer: str | Any = Field(default="") success: bool = Field(default=True) metadata: dict = Field(default_factory=dict) diff --git a/reme_ai/core/schema/service_config.py b/reme/core/schema/service_config.py similarity index 84% rename from reme_ai/core/schema/service_config.py rename to reme/core/schema/service_config.py index 4c6eb543..95461196 100644 --- a/reme_ai/core/schema/service_config.py +++ b/reme/core/schema/service_config.py @@ -1,7 +1,6 @@ """Configuration schemas for service components using Pydantic models.""" import os -from typing import Dict, List from pydantic import BaseModel, Field, ConfigDict @@ -99,15 +98,15 @@ class ServiceConfig(BaseModel): thread_pool_max_workers: int = Field(default=16) ray_max_workers: int = Field(default=-1) init_logger: bool = Field(default=True) - disabled_flows: List[str] = Field(default_factory=list) - enabled_flows: List[str] = Field(default_factory=list) - mcp_servers: Dict[str, dict] = Field(default_factory=dict, description="External MCP Server configuration") + disabled_flows: list[str] = Field(default_factory=list) + enabled_flows: list[str] = Field(default_factory=list) + mcp_servers: dict[str, dict] = Field(default_factory=dict) mcp: MCPConfig = Field(default_factory=MCPConfig) http: HttpConfig = Field(default_factory=HttpConfig) cmd: CmdConfig = Field(default_factory=CmdConfig) - flow: Dict[str, FlowConfig] = Field(default_factory=dict) - llm: Dict[str, LLMConfig] = Field(default_factory=dict) - embedding_model: Dict[str, EmbeddingModelConfig] = Field(default_factory=dict) - vector_store: Dict[str, VectorStoreConfig] = Field(default_factory=dict) - token_counter: Dict[str, TokenCounterConfig] = Field(default_factory=dict) + flow: dict[str, FlowConfig] = Field(default_factory=dict) + llm: dict[str, LLMConfig] = Field(default_factory=dict) + embedding_model: dict[str, EmbeddingModelConfig] = Field(default_factory=dict) + vector_store: dict[str, VectorStoreConfig] = Field(default_factory=dict) + token_counter: dict[str, TokenCounterConfig] = Field(default_factory=dict) diff --git a/reme_ai/core/schema/stream_chunk.py b/reme/core/schema/stream_chunk.py similarity index 100% rename from reme_ai/core/schema/stream_chunk.py rename to reme/core/schema/stream_chunk.py diff --git a/reme_ai/core/schema/tool_call.py b/reme/core/schema/tool_call.py similarity index 81% rename from reme_ai/core/schema/tool_call.py rename to reme/core/schema/tool_call.py index 0a96f8c4..e7ba9b78 100644 --- a/reme_ai/core/schema/tool_call.py +++ b/reme/core/schema/tool_call.py @@ -1,6 +1,4 @@ -""" -MCP Tool Schema definitions for recursive JSON Schema representation. -""" +"""MCP Tool Schema definitions for recursive JSON Schema representation.""" import json from typing import Any, Dict, List, Optional, Union @@ -41,14 +39,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 @@ -105,11 +103,6 @@ class ToolCall(BaseModel): description="Specification for input parameters", ) - output: ToolAttr = Field( - default_factory=lambda: ToolAttr(type="object", properties={}), - description="Specification for the execution result (Schema)", - ) - @model_validator(mode="before") @classmethod def init_tool_call(cls, data: dict) -> dict: @@ -147,6 +140,68 @@ class ToolCall(BaseModel): }, } + def simple_output_dump(self) -> dict: + """Convert ToolCall to output format dictionary for API responses.""" + return { + "index": self.index, + "id": self.id, + self.type: { + "arguments": self.arguments, + "name": self.name, + }, + "type": self.type, + } + + @property + def argument_dict(self) -> dict: + """Parse and return arguments as a dictionary.""" + return json.loads(self.arguments) + + def check_argument(self) -> bool: + """Check if arguments can be parsed as valid JSON.""" + try: + _ = self.argument_dict + return True + except Exception: + return False + + def sanitize_and_check_argument(self) -> bool: + """ + Attempt to sanitize and validate arguments JSON. + Common issues from LLM streaming: + - Extra closing brackets: }]}] -> }] + - Missing closing brackets + - Trailing commas + """ + if not self.arguments or not self.arguments.strip(): + return False + + try: + # First try parsing as-is + _ = json.loads(self.arguments) + return True + except json.JSONDecodeError: + pass + + # Try to fix common issues + sanitized = self.arguments.strip() + + # Remove trailing extra brackets/braces + # Pattern: if it ends with multiple closing chars, try removing extras + while len(sanitized) > 1: + try: + json.loads(sanitized) + self.arguments = sanitized # Update with sanitized version + return True + except json.JSONDecodeError: + # Try removing last character + if sanitized[-1] in "]}": + sanitized = sanitized[:-1].rstrip() + else: + break + + return False + @classmethod def from_mcp_tool(cls, tool: Tool) -> "ToolCall": """Creates a ToolCall instance from an MCP Tool object.""" @@ -164,28 +219,3 @@ class ToolCall(BaseModel): description=self.description, inputSchema=self.parameters.simple_input_dump(), ) - - @property - def argument_dict(self) -> dict: - """Parse and return arguments as a dictionary.""" - return json.loads(self.arguments) - - def check_argument(self) -> bool: - """Check if arguments can be parsed as valid JSON.""" - try: - _ = self.argument_dict - return True - except Exception: - return False - - def simple_output_dump(self) -> dict: - """Convert ToolCall to output format dictionary for API responses.""" - return { - "index": self.index, - "id": self.id, - self.type: { - "arguments": self.arguments, - "name": self.name, - }, - "type": self.type, - } diff --git a/reme_ai/core/schema/vector_node.py b/reme/core/schema/vector_node.py similarity index 100% rename from reme_ai/core/schema/vector_node.py rename to reme/core/schema/vector_node.py diff --git a/reme_ai/core/service/__init__.py b/reme/core/service/__init__.py similarity index 64% rename from reme_ai/core/service/__init__.py rename to reme/core/service/__init__.py index e9f00a65..1cd4ae58 100644 --- a/reme_ai/core/service/__init__.py +++ b/reme/core/service/__init__.py @@ -4,6 +4,7 @@ from .base_service import BaseService from .cmd_service import CmdService from .http_service import HttpService from .mcp_service import MCPService +from ..context import R __all__ = [ "BaseService", @@ -11,3 +12,7 @@ __all__ = [ "HttpService", "MCPService", ] + +R.service.register("cmd")(CmdService) +R.service.register("http")(HttpService) +R.service.register("mcp")(MCPService) diff --git a/reme_ai/core/service/base_service.py b/reme/core/service/base_service.py similarity index 76% rename from reme_ai/core/service/base_service.py rename to reme/core/service/base_service.py index 28c82fb3..9496391c 100644 --- a/reme_ai/core/service/base_service.py +++ b/reme/core/service/base_service.py @@ -5,7 +5,7 @@ from abc import ABC, abstractmethod from loguru import logger from pydantic import BaseModel -from ..context import C +from ..context import ServiceContext from ..flow import BaseFlow from ..schema import ToolCall from ..utils import create_pydantic_model @@ -14,8 +14,10 @@ from ..utils import create_pydantic_model class BaseService(ABC): """Abstract base class for services that integrate and execute flows.""" - def __init__(self, **kwargs): + def __init__(self, service_context: ServiceContext, **kwargs): """Initialize the base service.""" + self.service_context: ServiceContext = service_context + self.service_config = self.service_context.service_config self.kwargs = kwargs @abstractmethod @@ -32,10 +34,10 @@ class BaseService(ABC): def run(self): """Initialize and integrate all flows registered in the global context.""" flow_names: list[str] = [] - for _, flow in C.flow_dict.items(): + for flow in self.service_context.flows.values(): flow_name = self.integrate_flow(flow) if flow_name: flow_names.append(flow_name) if flow_names: - logger.info(f"integrate {','.join(flow_names)}") + logger.info(f"Integrated {','.join(flow_names)}") diff --git a/reme_ai/core/service/cmd_service.py b/reme/core/service/cmd_service.py similarity index 75% rename from reme_ai/core/service/cmd_service.py rename to reme/core/service/cmd_service.py index c02efa20..18d147a6 100644 --- a/reme_ai/core/service/cmd_service.py +++ b/reme/core/service/cmd_service.py @@ -3,12 +3,10 @@ from loguru import logger from .base_service import BaseService -from ..context import C from ..flow import CmdFlow, BaseFlow from ..utils.common_utils import run_coro_safely -@C.register_service("cmd") class CmdService(BaseService): """Service implementation for handling command flow execution logic.""" @@ -19,16 +17,16 @@ class CmdService(BaseService): def integrate_flow(self, flow: BaseFlow) -> str | None: """Integrate the workflow configuration into the command service.""" - self._cmd_flow = CmdFlow(flow=C.service_config.flow) + self._cmd_flow = CmdFlow(flow=self.service_config.cmd.flow) def run(self): """Execute the command flow in either asynchronous or synchronous mode.""" super().run() - + kwargs = self.service_config.cmd.model_extra if self._cmd_flow.async_mode: - response = run_coro_safely(self._cmd_flow.call(**C.service_config.cmd.model_extra)) + response = run_coro_safely(self._cmd_flow.call(**kwargs)) else: - response = self._cmd_flow.call_sync(**C.service_config.cmd.model_extra) + response = self._cmd_flow.call_sync(**kwargs) if response.answer: logger.info(f"response.answer={response.answer}") diff --git a/reme_ai/core/service/http_service.py b/reme/core/service/http_service.py similarity index 95% rename from reme_ai/core/service/http_service.py rename to reme/core/service/http_service.py index 8ef5d4ce..694dd3bd 100644 --- a/reme_ai/core/service/http_service.py +++ b/reme/core/service/http_service.py @@ -9,20 +9,18 @@ from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import StreamingResponse from .base_service import BaseService -from ..context import C from ..flow import BaseFlow from ..schema import Response from ..utils.common_utils import execute_stream_task -@C.register_service("http") class HttpService(BaseService): """Expose flows via HTTP REST and SSE endpoints.""" def __init__(self, **kwargs): """Initialize FastAPI app with CORS and health checks.""" super().__init__(**kwargs) - self.app = FastAPI(title=C.service_config.app_name) + self.app = FastAPI(title=self.service_config.app_name) self.app.add_middleware( CORSMiddleware, allow_origins=["*"], @@ -75,7 +73,7 @@ class HttpService(BaseService): def run(self): """Start the Uvicorn server.""" super().run() - cfg = C.service_config.http + cfg = self.service_config.http uvicorn.run( self.app, host=cfg.host, diff --git a/reme_ai/core/service/mcp_service.py b/reme/core/service/mcp_service.py similarity index 85% rename from reme_ai/core/service/mcp_service.py rename to reme/core/service/mcp_service.py index 7ac62a26..65183230 100644 --- a/reme_ai/core/service/mcp_service.py +++ b/reme/core/service/mcp_service.py @@ -1,23 +1,19 @@ """Model Context Protocol (MCP) service implementation.""" -from typing import Any - from fastmcp import FastMCP from fastmcp.tools import FunctionTool from .base_service import BaseService -from ..context import C from ..flow import BaseFlow -@C.register_service("mcp") class MCPService(BaseService): """Expose flows as Model Context Protocol (MCP) tools.""" - def __init__(self, **kwargs: Any): + def __init__(self, **kwargs): """Initialize FastMCP instance with service settings.""" super().__init__(**kwargs) - self.mcp = FastMCP(name=C.service_config.app_name) + self.mcp = FastMCP(name=self.service_config.app_name) def integrate_flow(self, flow: BaseFlow) -> str | None: """Register a non-streaming flow as an MCP tool.""" @@ -45,12 +41,8 @@ class MCPService(BaseService): def run(self): """Run the MCP server with specified transport protocol.""" super().run() - cfg = C.service_config.mcp - + cfg = self.service_config.mcp run_args: dict = {"transport": cfg.transport, "show_banner": False, **cfg.model_extra} - - # Add network settings for non-stdio transports if cfg.transport != "stdio": run_args.update({"host": cfg.host, "port": cfg.port}) - self.mcp.run(**run_args) diff --git a/reme_ai/core/token_counter/__init__.py b/reme/core/token_counter/__init__.py similarity index 58% rename from reme_ai/core/token_counter/__init__.py rename to reme/core/token_counter/__init__.py index a9b50826..f0cb5e28 100644 --- a/reme_ai/core/token_counter/__init__.py +++ b/reme/core/token_counter/__init__.py @@ -3,9 +3,14 @@ from .base_token_counter import BaseTokenCounter from .hf_token_counter import HFTokenCounter from .openai_token_counter import OpenAITokenCounter +from ..context import R __all__ = [ "BaseTokenCounter", "HFTokenCounter", "OpenAITokenCounter", ] + +R.token_counter.register("base")(BaseTokenCounter) +R.token_counter.register("hf")(HFTokenCounter) +R.token_counter.register("openai")(OpenAITokenCounter) diff --git a/reme_ai/core/token_counter/base_token_counter.py b/reme/core/token_counter/base_token_counter.py similarity index 96% rename from reme_ai/core/token_counter/base_token_counter.py rename to reme/core/token_counter/base_token_counter.py index 76ccdb78..6f98cd02 100644 --- a/reme_ai/core/token_counter/base_token_counter.py +++ b/reme/core/token_counter/base_token_counter.py @@ -2,13 +2,12 @@ import math import re + from loguru import logger -from ..context import C from ..schema import Message, ToolCall -@C.register_token_counter("base") class BaseTokenCounter: """A rule-based token counter for Chinese and non-Chinese text.""" diff --git a/reme_ai/core/token_counter/hf_token_counter.py b/reme/core/token_counter/hf_token_counter.py similarity index 97% rename from reme_ai/core/token_counter/hf_token_counter.py rename to reme/core/token_counter/hf_token_counter.py index 0bad4c78..4dd072a9 100644 --- a/reme_ai/core/token_counter/hf_token_counter.py +++ b/reme/core/token_counter/hf_token_counter.py @@ -5,11 +5,9 @@ import os from loguru import logger from .base_token_counter import BaseTokenCounter -from ..context import C from ..schema import Message, ToolCall -@C.register_token_counter("hf") class HFTokenCounter(BaseTokenCounter): """Token counter using transformers.AutoTokenizer.apply_chat_template.""" diff --git a/reme_ai/core/token_counter/openai_token_counter.py b/reme/core/token_counter/openai_token_counter.py similarity index 97% rename from reme_ai/core/token_counter/openai_token_counter.py rename to reme/core/token_counter/openai_token_counter.py index 5793e88a..672f8a45 100644 --- a/reme_ai/core/token_counter/openai_token_counter.py +++ b/reme/core/token_counter/openai_token_counter.py @@ -1,13 +1,13 @@ """Token counting implementation for OpenAI-compatible models.""" import json + from loguru import logger + from .base_token_counter import BaseTokenCounter -from ..context import C from ..schema import Message, ToolCall -@C.register_token_counter("openai") class OpenAITokenCounter(BaseTokenCounter): """Token counter for OpenAI models using tiktoken.""" diff --git a/reme_ai/core/utils/__init__.py b/reme/core/utils/__init__.py similarity index 94% rename from reme_ai/core/utils/__init__.py rename to reme/core/utils/__init__.py index 242f7be9..3d54e4e4 100644 --- a/reme_ai/core/utils/__init__.py +++ b/reme/core/utils/__init__.py @@ -4,7 +4,7 @@ from .cache_handler import CacheHandler from .case_converter import snake_to_camel, camel_to_snake from .common_utils import run_coro_safely, execute_stream_task from .env_utils import load_env -from .execute_tuils import exec_code, run_shell_command +from .execute_utils import exec_code, run_shell_command from .http_client import HttpClient from .llm_utils import extract_content, format_messages, deduplicate_memories from .logger_utils import init_logger diff --git a/reme_ai/core/utils/cache_handler.py b/reme/core/utils/cache_handler.py similarity index 91% rename from reme_ai/core/utils/cache_handler.py rename to reme/core/utils/cache_handler.py index 70c8585a..f3b0072f 100644 --- a/reme_ai/core/utils/cache_handler.py +++ b/reme/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/utils/case_converter.py b/reme/core/utils/case_converter.py similarity index 100% rename from reme_ai/core/utils/case_converter.py rename to reme/core/utils/case_converter.py diff --git a/reme_ai/core/utils/common_utils.py b/reme/core/utils/common_utils.py similarity index 100% rename from reme_ai/core/utils/common_utils.py rename to reme/core/utils/common_utils.py diff --git a/reme_ai/core/utils/env_utils.py b/reme/core/utils/env_utils.py similarity index 100% rename from reme_ai/core/utils/env_utils.py rename to reme/core/utils/env_utils.py diff --git a/reme_ai/core/utils/execute_tuils.py b/reme/core/utils/execute_utils.py similarity index 100% rename from reme_ai/core/utils/execute_tuils.py rename to reme/core/utils/execute_utils.py diff --git a/reme_ai/core/utils/http_client.py b/reme/core/utils/http_client.py similarity index 100% rename from reme_ai/core/utils/http_client.py rename to reme/core/utils/http_client.py diff --git a/reme_ai/core/utils/llm_utils.py b/reme/core/utils/llm_utils.py similarity index 100% rename from reme_ai/core/utils/llm_utils.py rename to reme/core/utils/llm_utils.py diff --git a/reme_ai/core/utils/logger_utils.py b/reme/core/utils/logger_utils.py similarity index 100% rename from reme_ai/core/utils/logger_utils.py rename to reme/core/utils/logger_utils.py diff --git a/reme_ai/core/utils/logo_utils.py b/reme/core/utils/logo_utils.py similarity index 100% rename from reme_ai/core/utils/logo_utils.py rename to reme/core/utils/logo_utils.py diff --git a/reme_ai/core/utils/mcp_client.py b/reme/core/utils/mcp_client.py similarity index 100% rename from reme_ai/core/utils/mcp_client.py rename to reme/core/utils/mcp_client.py diff --git a/reme_ai/core/utils/pydantic_config_parser.py b/reme/core/utils/pydantic_config_parser.py similarity index 100% rename from reme_ai/core/utils/pydantic_config_parser.py rename to reme/core/utils/pydantic_config_parser.py diff --git a/reme_ai/core/utils/pydantic_utils.py b/reme/core/utils/pydantic_utils.py similarity index 100% rename from reme_ai/core/utils/pydantic_utils.py rename to reme/core/utils/pydantic_utils.py diff --git a/reme_ai/core/utils/singleton.py b/reme/core/utils/singleton.py similarity index 100% rename from reme_ai/core/utils/singleton.py rename to reme/core/utils/singleton.py diff --git a/reme_ai/core/utils/time.py b/reme/core/utils/time.py similarity index 100% rename from reme_ai/core/utils/time.py rename to reme/core/utils/time.py diff --git a/reme_ai/core/vector_store/__init__.py b/reme/core/vector_store/__init__.py similarity index 62% rename from reme_ai/core/vector_store/__init__.py rename to reme/core/vector_store/__init__.py index 79500294..bd8a1869 100644 --- a/reme_ai/core/vector_store/__init__.py +++ b/reme/core/vector_store/__init__.py @@ -6,6 +6,7 @@ from .es_vector_store import ESVectorStore from .local_vector_store import LocalVectorStore from .pgvector_store import PGVectorStore from .qdrant_vector_store import QdrantVectorStore +from ..context import R __all__ = [ "BaseVectorStore", @@ -15,3 +16,9 @@ __all__ = [ "PGVectorStore", "QdrantVectorStore", ] + +R.vector_store.register("chroma")(ChromaVectorStore) +R.vector_store.register("es")(ESVectorStore) +R.vector_store.register("local")(LocalVectorStore) +R.vector_store.register("pgvector")(PGVectorStore) +R.vector_store.register("qdrant")(QdrantVectorStore) diff --git a/reme_ai/core/vector_store/base_vector_store.py b/reme/core/vector_store/base_vector_store.py similarity index 91% rename from reme_ai/core/vector_store/base_vector_store.py rename to reme/core/vector_store/base_vector_store.py index a4a8ca8e..40a84e99 100644 --- a/reme_ai/core/vector_store/base_vector_store.py +++ b/reme/core/vector_store/base_vector_store.py @@ -3,11 +3,11 @@ import asyncio from abc import ABC, abstractmethod from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor from functools import partial -from reme_ai.core.context import C -from reme_ai.core.embedding import BaseEmbeddingModel -from reme_ai.core.schema import VectorNode +from ..embedding import BaseEmbeddingModel +from ..schema import VectorNode class BaseVectorStore(ABC): @@ -17,20 +17,19 @@ class BaseVectorStore(ABC): self, collection_name: str, embedding_model: BaseEmbeddingModel, + thread_pool: ThreadPoolExecutor, **kwargs, ): """Initialize the vector store with a collection name and an embedding model.""" - if embedding_model is None: - raise ValueError("embedding_model is required") self.collection_name: str = collection_name self.embedding_model: BaseEmbeddingModel = embedding_model + self.thread_pool: ThreadPoolExecutor = thread_pool self.kwargs: dict = kwargs - @staticmethod - async def _run_sync_in_executor(sync_func: Callable, *args, **kwargs): + async def _run_sync_in_executor(self, sync_func: Callable, *args, **kwargs): """Run a synchronous function in the context-defined thread pool executor.""" loop = asyncio.get_running_loop() - return await loop.run_in_executor(C.thread_pool, partial(sync_func, *args, **kwargs)) + return await loop.run_in_executor(self.thread_pool, partial(sync_func, *args, **kwargs)) # noqa async def get_node_embedding(self, node: VectorNode) -> VectorNode: """Generate and assign embedding for a single vector node.""" diff --git a/reme_ai/core/vector_store/chroma_vector_store.py b/reme/core/vector_store/chroma_vector_store.py similarity index 89% rename from reme_ai/core/vector_store/chroma_vector_store.py rename to reme/core/vector_store/chroma_vector_store.py index 7baa5eab..710ce73e 100644 --- a/reme_ai/core/vector_store/chroma_vector_store.py +++ b/reme/core/vector_store/chroma_vector_store.py @@ -1,11 +1,11 @@ """ChromaDB vector store implementation for the ReMe framework.""" +from concurrent.futures import ThreadPoolExecutor from typing import Any from loguru import logger from .base_vector_store import BaseVectorStore -from ..context import C from ..embedding import BaseEmbeddingModel from ..schema import VectorNode @@ -20,7 +20,6 @@ except ImportError as e: Settings = None -@C.register_vector_store("chroma") class ChromaVectorStore(BaseVectorStore): """ChromaDB-based vector store implementation for local or remote storage.""" @@ -28,6 +27,7 @@ class ChromaVectorStore(BaseVectorStore): self, collection_name: str, embedding_model: BaseEmbeddingModel, + thread_pool: ThreadPoolExecutor, client: chromadb.ClientAPI | None = None, host: str | None = None, port: int | None = None, @@ -46,6 +46,7 @@ class ChromaVectorStore(BaseVectorStore): super().__init__( collection_name=collection_name, embedding_model=embedding_model, + thread_pool=thread_pool, **kwargs, ) @@ -117,14 +118,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 +161,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 +174,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 +191,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 @@ -209,8 +240,8 @@ class ChromaVectorStore(BaseVectorStore): try: self.client.delete_collection(name=collection_name) return True - except Exception as e: - logger.warning(f"Failed to delete collection {collection_name}: {e}") + except Exception as _e: + logger.warning(f"Failed to delete collection {collection_name}: {_e}") return False deleted = await self._run_sync_in_executor(_delete) diff --git a/reme_ai/core/vector_store/es_vector_store.py b/reme/core/vector_store/es_vector_store.py similarity index 91% rename from reme_ai/core/vector_store/es_vector_store.py rename to reme/core/vector_store/es_vector_store.py index 16226749..9a1019b8 100644 --- a/reme_ai/core/vector_store/es_vector_store.py +++ b/reme/core/vector_store/es_vector_store.py @@ -4,12 +4,12 @@ This module provides an Elasticsearch-based vector store that implements the Bas interface for high-performance dense vector storage and retrieval. """ +from concurrent.futures import ThreadPoolExecutor from typing import Any from loguru import logger from .base_vector_store import BaseVectorStore -from ..context import C from ..embedding import BaseEmbeddingModel from ..schema import VectorNode @@ -24,7 +24,6 @@ except ImportError as e: async_bulk = None -@C.register_vector_store("es") class ESVectorStore(BaseVectorStore): """Elasticsearch-based vector store for dense vector storage and kNN search.""" @@ -32,6 +31,7 @@ class ESVectorStore(BaseVectorStore): self, collection_name: str, embedding_model: BaseEmbeddingModel, + thread_pool: ThreadPoolExecutor, hosts: str | list[str] | None = None, basic_auth: tuple[str, str] | None = None, cloud_id: str | None = None, @@ -45,6 +45,7 @@ class ESVectorStore(BaseVectorStore): Args: collection_name: Name of the Elasticsearch index (converted to lowercase). embedding_model: Model instance used to generate vector embeddings. + thread_pool: ThreadPoolExecutor for running synchronous operations. hosts: Connection host(s) for the Elasticsearch cluster. basic_auth: Credentials for basic authentication. cloud_id: Deployment ID for Elastic Cloud. @@ -61,7 +62,12 @@ class ESVectorStore(BaseVectorStore): # Elasticsearch requires lowercase index names collection_name = collection_name.lower() - super().__init__(collection_name=collection_name, embedding_model=embedding_model, **kwargs) + super().__init__( + collection_name=collection_name, + embedding_model=embedding_model, + thread_pool=thread_pool, + **kwargs, + ) # Initialize AsyncElasticsearch client self.client = AsyncElasticsearch( @@ -262,9 +268,21 @@ 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 +466,21 @@ 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/core/vector_store/local_vector_store.py similarity index 92% rename from reme_ai/core/vector_store/local_vector_store.py rename to reme/core/vector_store/local_vector_store.py index cce3cae2..a86e226a 100644 --- a/reme_ai/core/vector_store/local_vector_store.py +++ b/reme/core/vector_store/local_vector_store.py @@ -1,17 +1,16 @@ """Local file system vector store implementation for ReMe.""" import json +from concurrent.futures import ThreadPoolExecutor from pathlib import Path from loguru import logger from .base_vector_store import BaseVectorStore -from ..context import C from ..embedding import BaseEmbeddingModel from ..schema import VectorNode -@C.register_vector_store("local") class LocalVectorStore(BaseVectorStore): """Local file system-based vector store using JSON files and manual cosine similarity.""" @@ -19,11 +18,17 @@ class LocalVectorStore(BaseVectorStore): self, collection_name: str, embedding_model: BaseEmbeddingModel, + thread_pool: ThreadPoolExecutor, root_path: str = "./local_vector_store", **kwargs, ): """Initialize the local vector store with a root path and collection name.""" - super().__init__(collection_name=collection_name, embedding_model=embedding_model, **kwargs) + super().__init__( + collection_name=collection_name, + embedding_model=embedding_model, + thread_pool=thread_pool, + **kwargs, + ) self.root_path = Path(root_path) self.collection_path = self.root_path / collection_name self.root_path.mkdir(parents=True, exist_ok=True) @@ -91,17 +96,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/core/vector_store/pgvector_store.py similarity index 85% rename from reme_ai/core/vector_store/pgvector_store.py rename to reme/core/vector_store/pgvector_store.py index a23c84b9..279975d0 100644 --- a/reme_ai/core/vector_store/pgvector_store.py +++ b/reme/core/vector_store/pgvector_store.py @@ -1,12 +1,13 @@ """PostgreSQL pgvector implementation for vector storage and retrieval.""" import json +import re +from concurrent.futures import ThreadPoolExecutor from typing import Any from loguru import logger from .base_vector_store import BaseVectorStore -from ..context import C from ..embedding import BaseEmbeddingModel from ..schema import VectorNode @@ -21,14 +22,33 @@ except ImportError as e: Pool = None -@C.register_vector_store("pgvector") 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, embedding_model: BaseEmbeddingModel, + thread_pool: ThreadPoolExecutor, host: str = "localhost", port: int = 5432, database: str = "postgres", @@ -47,7 +67,15 @@ class PGVectorStore(BaseVectorStore): "PGVector requires extra dependencies. Install with `pip install asyncpg pgvector`", ) from _ASYNCPG_IMPORT_ERROR - super().__init__(collection_name=collection_name, embedding_model=embedding_model, **kwargs) + # Validate collection name to prevent SQL injection + self._validate_table_name(collection_name) + + super().__init__( + collection_name=collection_name, + embedding_model=embedding_model, + thread_pool=thread_pool, + **kwargs, + ) self.dsn = dsn self.host = host @@ -106,6 +134,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 +179,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 +187,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 +283,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 +299,29 @@ 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 " + f"(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 +345,11 @@ 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) + for i in range(len(filter_params), 0, -1): + new_placeholder = f"${i + 1}" + filter_clause = re.sub(rf"\${i}\b", new_placeholder, filter_clause) async with pool.acquire() as conn: sql = f""" @@ -365,7 +420,7 @@ class PGVectorStore(BaseVectorStore): async with pool.acquire() as conn: result = await conn.execute(f"DELETE FROM {self.collection_name}") - logger.info(f"Deleted all documents from {self.collection_name}") + logger.info(f"Deleted all documents from {self.collection_name} result={result}") async def update(self, nodes: VectorNode | list[VectorNode], **kwargs): """Update existing vector nodes with new content, embeddings, or metadata.""" diff --git a/reme_ai/core/vector_store/qdrant_vector_store.py b/reme/core/vector_store/qdrant_vector_store.py similarity index 84% rename from reme_ai/core/vector_store/qdrant_vector_store.py rename to reme/core/vector_store/qdrant_vector_store.py index 1ac4db64..93ccee70 100644 --- a/reme_ai/core/vector_store/qdrant_vector_store.py +++ b/reme/core/vector_store/qdrant_vector_store.py @@ -1,11 +1,11 @@ """Qdrant vector store implementation for the ReMe project.""" +from concurrent.futures import ThreadPoolExecutor from typing import Any from loguru import logger from .base_vector_store import BaseVectorStore -from ..context import C from ..embedding import BaseEmbeddingModel from ..schema import VectorNode @@ -36,7 +36,6 @@ except ImportError as e: VectorParams = None -@C.register_vector_store("qdrant") class QdrantVectorStore(BaseVectorStore): """Vector store implementation using Qdrant for dense vector search.""" @@ -44,6 +43,7 @@ class QdrantVectorStore(BaseVectorStore): self, collection_name: str, embedding_model: BaseEmbeddingModel, + thread_pool: ThreadPoolExecutor, host: str | None = None, port: int = 6333, path: str | None = None, @@ -61,6 +61,7 @@ class QdrantVectorStore(BaseVectorStore): Args: collection_name: Name of the collection. embedding_model: Model used for generating vector embeddings. + thread_pool: ThreadPoolExecutor for running synchronous operations. host: Server host address. port: HTTP port for the server. path: Local storage path for on-disk/in-memory mode. @@ -78,7 +79,14 @@ class QdrantVectorStore(BaseVectorStore): "Qdrant requires extra dependencies. Install with `pip install qdrant-client`", ) from _QDRANT_IMPORT_ERROR - super().__init__(collection_name=collection_name, embedding_model=embedding_model, **kwargs) + super().__init__( + collection_name=collection_name, + embedding_model=embedding_model, + thread_pool=thread_pool, + **kwargs, + ) + + client_kwargs = {k: v for k, v in kwargs.items() if k != "thread_pool"} self.client = AsyncQdrantClient( host=host, @@ -89,7 +97,7 @@ class QdrantVectorStore(BaseVectorStore): https=https, grpc_port=grpc_port, prefer_grpc=prefer_grpc, - **kwargs, + **client_kwargs, ) self.is_local = path is not None @@ -246,29 +254,67 @@ 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, " + f"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, " + f"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/reme_app.py b/reme/reme_app.py new file mode 100644 index 00000000..e41aa8a6 --- /dev/null +++ b/reme/reme_app.py @@ -0,0 +1,90 @@ +"""ReMe application classes for simplified configuration and execution.""" + +import asyncio +import sys + +from .config import ReMeConfigParser +from .core.context import ServiceContext +from .core.flow import BaseFlow +from .core.schema import Response +from .core.utils import execute_stream_task + + +class ReMeApp: + """ReMe application with config file support and flow execution methods.""" + + def __init__( + self, + *args, + llm_api_key: str | None = None, + llm_api_base: str | None = None, + embedding_api_key: str | None = None, + embedding_api_base: str | None = None, + enable_logo: bool = True, + **kwargs, + ): + self.service_context = ServiceContext( + *args, + llm_api_key=llm_api_key, + llm_api_base=llm_api_base, + embedding_api_key=embedding_api_key, + embedding_api_base=embedding_api_base, + service_config=None, + parser=ReMeConfigParser, + config_path=None, + enable_logo=enable_logo, + **kwargs, + ) + + async def __aenter__(self): + """Async context manager entry.""" + return self + + def __enter__(self): + """Context manager entry.""" + return self + + async def __aexit__(self, exc_type=None, exc_val=None, exc_tb=None): + """Async context manager exit.""" + await self.service_context.close() + return False + + def __exit__(self, exc_type=None, exc_val=None, exc_tb=None): + """Context manager exit.""" + self.service_context.close_sync() + return False + + async def execute_flow(self, name: str, **kwargs) -> Response: + """Execute a flow with the given name and parameters.""" + assert name in self.service_context.flows, f"Flow {name} not found" + flow: BaseFlow = self.service_context.flows[name] + return await flow.call(**kwargs) + + async def execute_stream_flow(self, name: str, **kwargs): + """Execute a stream flow with the given name and parameters.""" + assert name in self.service_context.flows, f"Flow {name} not found" + flow: BaseFlow = self.service_context.flows[name] + assert flow.stream is True, "non-stream flow is not supported in execute_stream_flow!" + stream_queue = asyncio.Queue() + task = asyncio.create_task(flow.call(stream_queue=stream_queue, **kwargs)) + async for chunk in execute_stream_task( + stream_queue=stream_queue, + task=task, + task_name=name, + as_bytes=False, + ): + yield chunk + + def run_service(self): + """Run the configured service (HTTP, MCP, or CMD).""" + self.service_context.service.run() + + +def main(): + """Main entry point for running ReMe application from command line.""" + with ReMeApp(*sys.argv[1:]) as app: + app.run_service() + + +if __name__ == "__main__": + main() diff --git a/reme/tool/__init__.py b/reme/tool/__init__.py new file mode 100644 index 00000000..5c0d019b --- /dev/null +++ b/reme/tool/__init__.py @@ -0,0 +1,11 @@ +"""Tool""" + +from . import gallery +from . import memory +from . import search + +__all__ = [ + "gallery", + "memory", + "search", +] diff --git a/reme/tool/gallery/__init__.py b/reme/tool/gallery/__init__.py new file mode 100644 index 00000000..30e1d931 --- /dev/null +++ b/reme/tool/gallery/__init__.py @@ -0,0 +1,16 @@ +"""execute tool""" + +from .execute_code import ExecuteCode +from .execute_shell import ExecuteShell +from .think_tool import ThinkTool +from ...core import R + +__all__ = [ + "ExecuteCode", + "ExecuteShell", + "ThinkTool", +] + +for name in __all__: + tool_class = globals()[name] + R.op.register()(tool_class) diff --git a/reme_ai/tool/execute/execute_code.py b/reme/tool/gallery/execute_code.py similarity index 87% rename from reme_ai/tool/execute/execute_code.py rename to reme/tool/gallery/execute_code.py index ea259487..6cb77c85 100644 --- a/reme_ai/tool/execute/execute_code.py +++ b/reme/tool/gallery/execute_code.py @@ -4,15 +4,13 @@ This module provides an operation that can execute Python code strings and return the output or error messages. """ -from ...core.context import C -from ...core.op import BaseOp +from ...core.op import BaseTool from ...core.schema import ToolCall from ...core.utils import exec_code -@C.register_op() -class ExecuteCode(BaseOp): +class ExecuteCode(BaseTool): """Operation for executing Python code dynamically. This operation takes Python code as input, executes it in a safe context, @@ -40,4 +38,4 @@ class ExecuteCode(BaseOp): self.execute_sync() def execute_sync(self): - self.output = exec_code(self.context.code) + return exec_code(self.context.code) diff --git a/reme_ai/tool/execute/execute_code.yaml b/reme/tool/gallery/execute_code.yaml similarity index 100% rename from reme_ai/tool/execute/execute_code.yaml rename to reme/tool/gallery/execute_code.yaml diff --git a/reme_ai/tool/execute/execute_shell.py b/reme/tool/gallery/execute_shell.py similarity index 90% rename from reme_ai/tool/execute/execute_shell.py rename to reme/tool/gallery/execute_shell.py index 6e244ddb..de862602 100644 --- a/reme_ai/tool/execute/execute_shell.py +++ b/reme/tool/gallery/execute_shell.py @@ -4,15 +4,13 @@ This module provides an operation that can execute shell commands asynchronously and return the output, error, and exit code. """ -from ...core.context import C -from ...core.op import BaseOp +from ...core.op import BaseTool from ...core.schema import ToolCall from ...core.utils import run_shell_command -@C.register_op() -class ExecuteShell(BaseOp): +class ExecuteShell(BaseTool): """Operation for executing shell commands asynchronously. This operation takes a shell command as input, executes it asynchronously, @@ -46,4 +44,4 @@ class ExecuteShell(BaseOp): f"Exit Code: {return_code if return_code is not None else '(none)'}", ] - self.output = "\n".join(result_parts) + return "\n".join(result_parts) diff --git a/reme_ai/tool/execute/execute_shell.yaml b/reme/tool/gallery/execute_shell.yaml similarity index 100% rename from reme_ai/tool/execute/execute_shell.yaml rename to reme/tool/gallery/execute_shell.yaml diff --git a/reme_ai/mem_tool/think_tool.py b/reme/tool/gallery/think_tool.py similarity index 60% rename from reme_ai/mem_tool/think_tool.py rename to reme/tool/gallery/think_tool.py index 1d26446a..f01646db 100644 --- a/reme_ai/mem_tool/think_tool.py +++ b/reme/tool/gallery/think_tool.py @@ -4,29 +4,15 @@ This module provides a tool that prompts the model for explicit reflection before taking actions, helping agents reason about their next steps. """ -from .base_memory_tool import BaseMemoryTool -from ..core.context import C -from ..core.schema import ToolCall +from ...core.op import BaseTool +from ...core.schema import ToolCall -@C.register_op() -class ThinkTool(BaseMemoryTool): - """Utility that prompts the model for explicit reflection text. - - This tool provides a thinking mechanism for agents to reflect on: - 1. Whether current context is sufficient to answer - 2. What information is missing - 3. Which tool and parameters to use next - """ +class ThinkTool(BaseTool): + """Utility that prompts the model for explicit reflection text.""" def __init__(self, add_output_reflection: bool = False, **kwargs): - """Initialize the think tool. - - Args: - add_output_reflection: If True, outputs the reflection content; - if False, outputs a confirmation message - **kwargs: Additional arguments passed to BaseOp - """ + """Initialize the think tool.""" super().__init__(**kwargs) self.add_output_reflection: bool = add_output_reflection @@ -51,6 +37,6 @@ class ThinkTool(BaseMemoryTool): async def execute(self): """Execute the think tool by processing reflection input.""" if self.add_output_reflection: - self.output = self.context["reflection"] + return self.context["reflection"] else: - self.output = self.get_prompt("reflection_output") + return self.get_prompt("reflection_output") diff --git a/reme_ai/mem_tool/think_tool.yaml b/reme/tool/gallery/think_tool.yaml similarity index 100% rename from reme_ai/mem_tool/think_tool.yaml rename to reme/tool/gallery/think_tool.yaml diff --git a/reme/tool/memory/__init__.py b/reme/tool/memory/__init__.py new file mode 100644 index 00000000..9113f35a --- /dev/null +++ b/reme/tool/memory/__init__.py @@ -0,0 +1,28 @@ +"""memory tools""" + +from .base_memory_tool import BaseMemoryTool +from .history.add_history import AddHistory +from .history.read_history import ReadHistory +from .identity.add_identity import AddIdentity +from .identity.read_identity import ReadIdentity +from .meta.add_meta_memory import AddMetaMemory +from .meta.read_meta_memory import ReadMetaMemory +from .user_profile.read_user_profile import ReadUserProfile +from .user_profile.update_user_profile import UpdateUserProfile +from ...core import R + +__all__ = [ + "BaseMemoryTool", + "AddHistory", + "ReadHistory", + "AddIdentity", + "ReadIdentity", + "AddMetaMemory", + "ReadMetaMemory", + "ReadUserProfile", + "UpdateUserProfile", +] + +for name in __all__: + tool_class = globals()[name] + R.op.register()(tool_class) diff --git a/reme/tool/memory/base_memory_tool.py b/reme/tool/memory/base_memory_tool.py new file mode 100644 index 00000000..adc2f6b2 --- /dev/null +++ b/reme/tool/memory/base_memory_tool.py @@ -0,0 +1,96 @@ +"""Base class for memory tool""" + +from abc import ABCMeta +from pathlib import Path + +from ...core.enumeration import MemoryType +from ...core.op import BaseTool +from ...core.schema import ToolCall, MemoryNode, ToolAttr +from ...core.utils import CacheHandler + + +class BaseMemoryTool(BaseTool, metaclass=ABCMeta): + """Base class for memory tool""" + + def __init__( + self, + enable_multiple: bool = True, + enable_thinking_params: bool = False, + local_memory_path: str = "./reme_local_memory", + **kwargs, + ): + super().__init__(**kwargs) + self.enable_multiple: bool = enable_multiple + self.enable_thinking_params: bool = enable_thinking_params + self.local_memory_path: str = local_memory_path + self.memory_nodes: list[MemoryNode | str] = [] + + def _build_tool_call(self) -> ToolCall: + """Build and return the tool call schema""" + + def _build_multiple_tool_call(self) -> ToolCall: + """Build and return the multiple tool call schema""" + + @property + def tool_call(self) -> ToolCall | None: + """Get the tool call schema.""" + if self._tool_call is None: + if self.enable_multiple: + self._tool_call = self._build_multiple_tool_call() + else: + self._tool_call = self._build_tool_call() + self._tool_call.name = self._tool_call.name or self.name + + # Add thinking parameter if enabled + if self.enable_thinking_params: + parameters = self._tool_call.parameters + if parameters and parameters.properties is not None: + if "thinking" not in parameters.properties: + parameters.properties = { + "thinking": ToolAttr( + type="string", + description="Your complete and detailed thinking process " + "about how to fill in each parameter", + ), + **parameters.properties, + } + if parameters.required is not None: + parameters.required = ["thinking", *parameters.required] + else: + parameters.required = ["thinking"] + return self._tool_call + + @property + def local_memory(self) -> CacheHandler: + """Create the meta memory cache handler.""" + return CacheHandler(Path(self.local_memory_path) / self.vector_store.collection_name) + + @property + def memory_type(self) -> MemoryType: + """Get the memory type from context.""" + return MemoryType(self.context.get("memory_type")) + + @property + def memory_target(self) -> str: + """Get the memory target from context.""" + return self.context.get("memory_target", "") + + @property + def memory_cache_key(self) -> str: + """Get the memory cache key from context.""" + return f"{self.memory_type.value}_{self.memory_target}".replace(" ", "_").lower() + + @property + def history_node(self) -> MemoryNode: + """Get the history node from context.""" + return self.context.get("history_node") + + @property + def retrieved_nodes(self) -> list[MemoryNode]: + """Get the retrieved nodes from context.""" + return self.context.get("retrieved_nodes") + + @property + def author(self) -> str: + """Get the author from context.""" + return self.context.get("author", "") diff --git a/reme_ai/mem_tool/meta/__init__.py b/reme/tool/memory/history/__init__.py similarity index 100% rename from reme_ai/mem_tool/meta/__init__.py rename to reme/tool/memory/history/__init__.py diff --git a/reme/tool/memory/history/add_history.py b/reme/tool/memory/history/add_history.py new file mode 100644 index 00000000..b46259f3 --- /dev/null +++ b/reme/tool/memory/history/add_history.py @@ -0,0 +1,47 @@ +"""Add history tool""" + +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ....core.enumeration import MemoryType +from ....core.schema import ToolCall, MemoryNode, Message +from ....core.utils import format_messages + + +class AddHistory(BaseMemoryTool): + """Tool to add historical dialogue to vector store""" + + def __init__(self, **kwargs): + kwargs["enable_multiple"] = False + super().__init__(**kwargs) + + def _build_tool_call(self) -> ToolCall: + """Build and return the tool call schema""" + return ToolCall( + **{ + "description": "Add original history dialogue.", + "parameters": { + "type": "object", + "properties": {}, + "required": [], + }, + }, + ) + + async def execute(self): + """Execute the add history operation""" + self.context.messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages] + history_content: str = (self.context.description + "\n" + format_messages(self.context.messages)).strip() + history_node = MemoryNode( + memory_type=MemoryType.HISTORY, + when_to_use=history_content[:100], + content=history_content, + author=self.author, + ) + logger.info(f"Adding history node: {history_node.model_dump_json(indent=2, exclude={'content'})}") + + vector_node = history_node.to_vector_node() + await self.vector_store.delete(vector_node.memory_id) + await self.vector_store.insert([vector_node]) + + return f"Successfully added history: {history_node.memory_id}" diff --git a/reme/tool/memory/history/read_history.py b/reme/tool/memory/history/read_history.py new file mode 100644 index 00000000..089a3309 --- /dev/null +++ b/reme/tool/memory/history/read_history.py @@ -0,0 +1,46 @@ +"""Read history memory tool""" + +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ....core.schema import MemoryNode, ToolCall + + +class ReadHistory(BaseMemoryTool): + """Read history memory tool""" + + def __init__(self, **kwargs): + kwargs["enable_multiple"] = False + super().__init__(**kwargs) + + def _build_tool_call(self) -> ToolCall: + """Build and return the tool call schema""" + return ToolCall( + **{ + "description": "Read original history dialogue.", + "parameters": { + "type": "object", + "properties": { + "history_id": { + "type": "string", + "description": "history_id", + }, + }, + "required": ["history_id"], + }, + }, + ) + + async def execute(self): + history_id = self.context.history_id + nodes = await self.vector_store.get(vector_ids=[history_id]) + + if not nodes: + output = f"No history: {history_id}" + logger.warning(output) + return output + + memory = MemoryNode.from_vector_node(nodes[0]) + output = f"Historical Dialogue[{history_id}]\n{memory.content}" + logger.info(f"Successfully read history memory: {history_id}") + return output diff --git a/reme/tool/memory/identity/__init__.py b/reme/tool/memory/identity/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/reme/tool/memory/identity/add_identity.py b/reme/tool/memory/identity/add_identity.py new file mode 100644 index 00000000..5f6ab984 --- /dev/null +++ b/reme/tool/memory/identity/add_identity.py @@ -0,0 +1,42 @@ +"""Add identity memory tool""" + +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ....core.schema import ToolCall + + +class AddIdentity(BaseMemoryTool): + """Tool to add or update agent identity memory""" + + def __init__(self, **kwargs): + kwargs["enable_multiple"] = False + super().__init__(**kwargs) + + def _build_tool_call(self) -> ToolCall: + return ToolCall( + **{ + "description": "add or update agent identity memory.", + "parameters": { + "type": "object", + "properties": { + "identity_memory": { + "type": "string", + "description": "Agent identity content, such as role, personality, or current state.", + }, + }, + "required": ["identity_memory"], + }, + }, + ) + + async def execute(self): + identity_memory = self.context.get("identity_memory", "") + + if not identity_memory: + logger.warning("No valid identity memory provided") + return "No valid identity memory provided for update." + + self.local_memory.save("identity_memory", identity_memory) + logger.info(f"Successfully updated identity memory: {identity_memory}") + return "Successfully updated identity memory." diff --git a/reme/tool/memory/identity/read_identity.py b/reme/tool/memory/identity/read_identity.py new file mode 100644 index 00000000..856aea55 --- /dev/null +++ b/reme/tool/memory/identity/read_identity.py @@ -0,0 +1,36 @@ +"""Read identity memory tool""" + +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ....core.schema import ToolCall + + +class ReadIdentity(BaseMemoryTool): + """Tool to read agent identity memory""" + + def __init__(self, **kwargs): + kwargs["enable_multiple"] = False + super().__init__(**kwargs) + + def _build_tool_call(self) -> ToolCall: + return ToolCall( + **{ + "description": "read agent identity memory.", + "parameters": { + "type": "object", + "properties": {}, + "required": [], + }, + }, + ) + + async def execute(self): + identity_memory = self.local_memory.load("identity_memory") + + if not identity_memory: + logger.info("No identity memory found") + return "No identity memory found." + + logger.info(f"Read identity memory: {identity_memory}") + return f"Identity\n{identity_memory}" diff --git a/reme/tool/memory/meta/__init__.py b/reme/tool/memory/meta/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/reme/tool/memory/meta/add_meta_memory.py b/reme/tool/memory/meta/add_meta_memory.py new file mode 100644 index 00000000..43e9f7db --- /dev/null +++ b/reme/tool/memory/meta/add_meta_memory.py @@ -0,0 +1,89 @@ +"""Add meta memory tool""" + +import json + +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ....core.enumeration import MemoryType +from ....core.schema import ToolCall + + +class AddMetaMemory(BaseMemoryTool): + """Tool to add memory metadata entries to meta storage""" + + def __init__(self, **kwargs): + kwargs["enable_multiple"] = True + super().__init__(**kwargs) + + def _build_multiple_tool_call(self) -> ToolCall: + """Build and return the multiple tool call schema""" + return ToolCall( + **{ + "description": "add memory metadata entries to register memory types and targets. " + "Before using, verify Main Agent's Meta Memory doesn't already contain the " + "same memory_type(memory_target) combinations.", + "parameters": { + "type": "object", + "properties": { + "meta_memories": { + "type": "array", + "description": "List of memory metadata entries to add", + "items": { + "type": "object", + "properties": { + "memory_type": { + "type": "string", + "description": "Type of memory: 'personal' for person-specific preferences, " + "'procedural' for how-to knowledge", + "enum": [MemoryType.PERSONAL.value, MemoryType.PROCEDURAL.value], + }, + "memory_target": { + "type": "string", + "description": "Target identifier, " + "e.g., person's name ('John') or domain ('deployment')", + }, + }, + "required": ["memory_type", "memory_target"], + }, + }, + }, + "required": ["meta_memories"], + }, + }, + ) + + async def execute(self): + existing_memories: list[dict] = self.local_memory.load("meta_memories") or [] + existing_set = {(m["memory_type"], m["memory_target"]) for m in existing_memories} + + # Filter and build new memories to add + new_memories: list[dict] = [] + meta_memories: list[dict] = self.context.get("meta_memories", []) + + for mem in meta_memories: + memory_type = mem.get("memory_type", "") + memory_target = mem.get("memory_target", "") + + # Check if valid and not duplicate + if ( + memory_type in [MemoryType.PERSONAL.value, MemoryType.PROCEDURAL.value] + and memory_target + and (memory_type, memory_target) not in existing_set + ): + new_memories.append({"memory_type": memory_type, "memory_target": memory_target}) + existing_set.add((memory_type, memory_target)) + + if not new_memories: + output = "No new meta memories to add (all entries already exist or invalid)." + logger.info(output) + return output + + # Merge, sort and save + all_memories = sorted(existing_memories + new_memories, key=lambda m: (m["memory_type"], m["memory_target"])) + self.local_memory.save("meta_memories", all_memories) + + # Format output + output = f"Successfully update meta memory entries: {json.dumps(new_memories, ensure_ascii=False)}" + logger.info(output) + return output diff --git a/reme/tool/memory/meta/read_meta_memory.py b/reme/tool/memory/meta/read_meta_memory.py new file mode 100644 index 00000000..5fb819bc --- /dev/null +++ b/reme/tool/memory/meta/read_meta_memory.py @@ -0,0 +1,66 @@ +"""Read meta memory tool""" + +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ....core.enumeration import MemoryType +from ....core.schema import ToolCall + + +class ReadMetaMemory(BaseMemoryTool): + """Tool to read memory metadata from meta storage""" + + TYPE_DESC_DICT = { + MemoryType.IDENTITY.value: "self-cognition memory storing agent's identity and state", + MemoryType.PERSONAL.value: "person-specific memory storing preferences and context", + MemoryType.PROCEDURAL.value: "procedural memory storing how-to knowledge and processes", + } + + def __init__(self, enable_identity_memory: bool = False, **kwargs): + kwargs["enable_multiple"] = False + super().__init__(**kwargs) + self.enable_identity_memory = enable_identity_memory + + def _build_tool_call(self) -> ToolCall: + return ToolCall( + **{ + "description": "read memory metadata registry to see what types of memories are being tracked.", + "parameters": { + "type": "object", + "properties": {}, + "required": [], + }, + }, + ) + + async def execute(self): + # Load and filter meta memories + result = self.local_memory.load("meta_memories") + all_memories = result if result is not None else [] + + memories = [ + m for m in all_memories if m.get("memory_type") in [MemoryType.PERSONAL.value, MemoryType.PROCEDURAL.value] + ] + + if self.enable_identity_memory: + memories.append( + { + "memory_type": MemoryType.IDENTITY.value, + "memory_target": "self", + }, + ) + + # Format output + if memories: + lines = [ + f"- {m['memory_type']}({m['memory_target']}): {self.TYPE_DESC_DICT.get(m['memory_type'], '')}" + for m in memories + ] + + output = "\n".join(lines) + logger.info(f"Retrieved {len(memories)} meta memory entries") + else: + output = "No memory metadata found." + logger.info(output) + + return output diff --git a/reme/tool/memory/user_profile/__init__.py b/reme/tool/memory/user_profile/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/reme/tool/memory/user_profile/read_user_profile.py b/reme/tool/memory/user_profile/read_user_profile.py new file mode 100644 index 00000000..dfbccd63 --- /dev/null +++ b/reme/tool/memory/user_profile/read_user_profile.py @@ -0,0 +1,60 @@ +"""Read user profile tool""" + +from typing import Literal + +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ....core.schema import ToolCall +from ....core.schema.memory_node import MemoryNode + + +class ReadUserProfile(BaseMemoryTool): + """Tool to read user profile from local memory""" + + def __init__(self, show_id: Literal["profile", "history"] = "profile", **kwargs): + kwargs["enable_multiple"] = False + super().__init__(**kwargs) + self.show_id = show_id + + def _build_tool_call(self) -> ToolCall: + return ToolCall( + **{ + "description": "read user profile.", + "parameters": { + "type": "object", + "properties": {}, + "required": [], + }, + }, + ) + + async def execute(self): + cached_data = self.local_memory.load(self.memory_cache_key, auto_clean=False) + + if not cached_data: + logger.info(f"No cached data found for {self.memory_cache_key}") + return "" + + nodes = [MemoryNode(**data) for data in cached_data] + nodes.sort(key=lambda n: n.metadata.get("conversation_time", "")) + + formatted_profiles = [] + for node in nodes: + parts = [] + if self.show_id == "profile": + parts.append(f"profile_id={node.memory_id}") + + if conv_time := node.metadata.get("conversation_time"): + parts.append(f"conversation_time={conv_time}") + + parts.append(f"{node.when_to_use}: {node.content}") + + if self.show_id == "history": + parts.append(f"history_id={node.ref_memory_id}") + + formatted_profiles.append(" ".join(parts)) + + logger.info(f"Read {len(formatted_profiles)} profiles from cache key: {self.memory_cache_key}") + + return "### User Profile\n" + "\n".join(formatted_profiles).strip() diff --git a/reme/tool/memory/user_profile/update_user_profile.py b/reme/tool/memory/user_profile/update_user_profile.py new file mode 100644 index 00000000..93f3a5dc --- /dev/null +++ b/reme/tool/memory/user_profile/update_user_profile.py @@ -0,0 +1,109 @@ +"""Update user profile tool""" + +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ....core.schema import ToolCall +from ....core.schema.memory_node import MemoryNode +from ....core.utils import deduplicate_memories + + +class UpdateUserProfile(BaseMemoryTool): + """Tool to update user profile by adding or removing profile entries""" + + def __init__(self, **kwargs): + kwargs["enable_multiple"] = True + super().__init__(**kwargs) + + def _build_multiple_tool_call(self) -> ToolCall: + """Build and return the multiple tool call schema""" + return ToolCall( + **{ + "description": "update user profile by adding or removing profile entries.", + "parameters": { + "type": "object", + "properties": { + "profile_ids_to_delete": { + "type": "array", + "description": "List of profile IDs to delete", + "items": {"type": "string"}, + }, + "profiles_to_add": { + "type": "array", + "description": "List of profiles to add", + "items": { + "type": "object", + "properties": { + "conversation_time": { + "type": "string", + "description": "Conversation time, e.g. '2020-01-01 00:00:00'", + }, + "profile_key": { + "type": "string", + "description": "Profile key or category, e.g. 'name'", + }, + "profile_value": { + "type": "string", + "description": "Profile value or content, e.g. 'John Smith'", + }, + }, + "required": ["conversation_time", "profile_key", "profile_value"], + }, + }, + }, + "required": ["profile_ids_to_delete", "profiles_to_add"], + }, + }, + ) + + async def execute(self): + # Get and deduplicate profile IDs to delete + profile_ids_to_delete = self.context.get("profile_ids_to_delete", []) + profile_ids_to_delete = list(dict.fromkeys([pid for pid in profile_ids_to_delete if pid])) + profiles_to_add = self.context.get("profiles_to_add", []) + + if not profile_ids_to_delete and not profiles_to_add: + return "No profiles to remove or add. Operation completed." + + # Load existing profiles from local memory + cached_data = self.local_memory.load(self.memory_cache_key, auto_clean=False) + existing_nodes = [MemoryNode(**data) for data in cached_data] if cached_data else [] + + # Remove profiles + removed_count = 0 + if profile_ids_to_delete: + original_count = len(existing_nodes) + existing_nodes = [n for n in existing_nodes if n.memory_id not in profile_ids_to_delete] + removed_count = original_count - len(existing_nodes) + logger.info(f"Removed {removed_count} profiles.") + + # Add new profiles + new_nodes = [] + if profiles_to_add: + for profile in profiles_to_add: + node = MemoryNode( + memory_type=self.memory_type, + memory_target=self.memory_target, + when_to_use=profile.get("profile_key", ""), + content=profile.get("profile_value", ""), + ref_memory_id=self.history_node.memory_id, + author=self.author, + metadata={"conversation_time": profile.get("conversation_time", "")}, + ) + new_nodes.append(node) + logger.info(f"Added {len(new_nodes)} new profiles.") + + # Deduplicate and save updated profiles + updated_nodes = deduplicate_memories(existing_nodes + new_nodes) + nodes_data = [node.model_dump(exclude_none=True) for node in updated_nodes] + self.local_memory.save(self.memory_cache_key, nodes_data) + + # Build output message + operations = [] + if removed_count > 0: + operations.append(f"removed {removed_count} old profiles.") + if len(new_nodes) > 0: + operations.append(f"added {len(new_nodes)} new profiles.") + operations.append("Operation completed.") + logger.info("\n".join(operations)) + return operations diff --git a/reme_ai/tool/search/__init__.py b/reme/tool/search/__init__.py similarity index 66% rename from reme_ai/tool/search/__init__.py rename to reme/tool/search/__init__.py index 5a73dc2e..6230c7a2 100644 --- a/reme_ai/tool/search/__init__.py +++ b/reme/tool/search/__init__.py @@ -3,9 +3,14 @@ from .dashscope_search import DashscopeSearch from .mock_search import MockSearch from .tavily_search import TavilySearch +from ...core import R __all__ = [ "DashscopeSearch", "MockSearch", "TavilySearch", ] + +for name in __all__: + tool_class = globals()[name] + R.op.register()(tool_class) diff --git a/reme_ai/tool/search/dashscope_search.py b/reme/tool/search/dashscope_search.py similarity index 92% rename from reme_ai/tool/search/dashscope_search.py rename to reme/tool/search/dashscope_search.py index 19bd8104..f0e0ae9d 100644 --- a/reme_ai/tool/search/dashscope_search.py +++ b/reme/tool/search/dashscope_search.py @@ -9,13 +9,11 @@ from typing import Literal from loguru import logger -from ...core.context import C -from ...core.op import BaseOp +from ...core.op import BaseTool from ...core.schema import ToolCall -@C.register_op() -class DashscopeSearch(BaseOp): +class DashscopeSearch(BaseTool): """Operation for performing web searches using Dashscope API. This operation uses Alibaba Cloud's Dashscope service to search the web @@ -61,8 +59,7 @@ class DashscopeSearch(BaseOp): if self.enable_cache: cached_result = self.cache.load(query) if cached_result: - self.output = cached_result["response_content"] - return + return cached_result["response_content"] if self.enable_role_prompt: user_query = self.prompt_format("role_prompt", query=query) @@ -108,4 +105,4 @@ class DashscopeSearch(BaseOp): if self.enable_cache: self.cache.save(query, final_result, expire_hours=self.cache_expire_hours) - self.output = final_result["response_content"] + return final_result["response_content"] diff --git a/reme_ai/tool/search/dashscope_search.yaml b/reme/tool/search/dashscope_search.yaml similarity index 100% rename from reme_ai/tool/search/dashscope_search.yaml rename to reme/tool/search/dashscope_search.yaml diff --git a/reme_ai/tool/search/mock_search.py b/reme/tool/search/mock_search.py similarity index 91% rename from reme_ai/tool/search/mock_search.py rename to reme/tool/search/mock_search.py index 187463dc..ed695b0c 100644 --- a/reme_ai/tool/search/mock_search.py +++ b/reme/tool/search/mock_search.py @@ -9,15 +9,13 @@ import random from loguru import logger -from ...core.context import C from ...core.enumeration import Role -from ...core.op import BaseOp +from ...core.op import BaseTool from ...core.schema import ToolCall, Message from ...core.utils import extract_content -@C.register_op() -class MockSearch(BaseOp): +class MockSearch(BaseTool): """Operation for generating mock search results. This operation generates simulated search results using an LLM, @@ -61,4 +59,4 @@ class MockSearch(BaseOp): return extract_content(message.content, "json") search_results: str = await self.llm.chat(messages=messages, callback_fn=callback_fn) - self.output = json.dumps(search_results, ensure_ascii=False, indent=2) + return json.dumps(search_results, ensure_ascii=False, indent=2) diff --git a/reme_ai/tool/search/mock_search.yaml b/reme/tool/search/mock_search.yaml similarity index 100% rename from reme_ai/tool/search/mock_search.yaml rename to reme/tool/search/mock_search.yaml diff --git a/reme_ai/tool/search/tavily_search.py b/reme/tool/search/tavily_search.py similarity index 90% rename from reme_ai/tool/search/tavily_search.py rename to reme/tool/search/tavily_search.py index 5c194bdc..bfe29879 100644 --- a/reme_ai/tool/search/tavily_search.py +++ b/reme/tool/search/tavily_search.py @@ -9,13 +9,11 @@ import os from loguru import logger -from ...core.context import C -from ...core.op import BaseOp +from ...core.op import BaseTool from ...core.schema import ToolCall -@C.register_op() -class TavilySearch(BaseOp): +class TavilySearch(BaseTool): """Operation for performing web searches using Tavily API. This operation uses the Tavily search service to find web content @@ -73,8 +71,7 @@ class TavilySearch(BaseOp): if self.enable_cache: cached_result = self.cache.load(query) if cached_result: - self.output = json.dumps(cached_result, ensure_ascii=False, indent=2) - return + return json.dumps(cached_result, ensure_ascii=False, indent=2) response = await self.client.search(query=query) logger.info(f"tavily_search response={response}") @@ -88,8 +85,7 @@ class TavilySearch(BaseOp): if self.enable_cache and final_result: self.cache.save(query, final_result, expire_hours=self.cache_expire_hours) - self.output = json.dumps(final_result, ensure_ascii=False, indent=2) - return + return json.dumps(final_result, ensure_ascii=False, indent=2) url_info_dict = {item["url"]: item for item in response["results"]} response_extract = await self.client.extract(urls=[item["url"] for item in response["results"]]) @@ -116,4 +112,4 @@ class TavilySearch(BaseOp): if self.enable_cache and final_result: self.cache.save(query, final_result, expire_hours=self.cache_expire_hours) - self.output = json.dumps(final_result, ensure_ascii=False, indent=2) + return json.dumps(final_result, ensure_ascii=False, indent=2) diff --git a/reme_ai/tool/search/tavily_search.yaml b/reme/tool/search/tavily_search.yaml similarity index 100% rename from reme_ai/tool/search/tavily_search.yaml rename to reme/tool/search/tavily_search.yaml diff --git a/reme/workflow/__init__.py b/reme/workflow/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/reme/workflow/procedural_memory/__init__.py b/reme/workflow/procedural_memory/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/reme/workflow/tool_memory/__init__.py b/reme/workflow/tool_memory/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/reme_ai/core/application.py b/reme_ai/core/application.py deleted file mode 100644 index 29eb846e..00000000 --- a/reme_ai/core/application.py +++ /dev/null @@ -1,221 +0,0 @@ -"""Main application module for managing ReMe AI service lifecycle and flow execution.""" - -import asyncio -import os - -from .context import C -from .flow import BaseFlow -from .schema import ServiceConfig, Response -from .utils import PydanticConfigParser, init_logger, execute_stream_task, run_coro_safely, load_env - - -class Application: - """ - Main application class for managing the lifecycle of ReMe AI services. - - Handles initialization, configuration, service management, and flow execution - for both synchronous and asynchronous contexts. - """ - - def __init__( - self, - *args, - llm_api_key: str | None = None, - llm_api_base: str | None = None, - embedding_api_key: str | None = None, - embedding_api_base: str | None = None, - service_config: ServiceConfig | None = None, - parser: type[PydanticConfigParser] | None = None, - config_path: str | None = None, - enable_logo: bool = True, - llm: dict | None = None, - embedding_model: dict | None = None, - vector_store: dict | None = None, - token_counter: dict | None = None, - **kwargs, - ): - """ - Initialize the Application with configuration settings. - - Args: - *args: Additional arguments passed to parser. Examples: - - "llm.default.model_name=qwen3-30b-a3b-thinking-2507" - - "llm.default.backend=openai_compatible" - - "llm.default.temperature=0.6" - - "embedding_model.default.model_name=text-embedding-v4" - - "embedding_model.default.backend=openai_compatible" - - "embedding_model.default.dimensions=1024" - - "vector_store.default.backend=memory" - - "vector_store.default.embedding_model=default" - llm_api_key: API key for LLM service - llm_api_base: Base URL for LLM service - embedding_api_key: API key for embedding service - embedding_api_base: Base URL for embedding service - service_config: Pre-built service configuration object - parser: Custom parser class for configuration (defaults to PydanticConfigParser) - config_path: Path to configuration file - enable_logo: Whether to display the ReMe logo on startup - llm: LLM configuration dictionary - embedding_model: Embedding model configuration dictionary - vector_store: Vector store configuration dictionary - token_counter: Token counter configuration dictionary - **kwargs: Additional keyword arguments passed to parser. Same format as args but as kwargs. Examples: - - **{"llm.default.model_name": "qwen3-30b-a3b-thinking-2507"} - """ - - load_env() - self._update_env("REME_LLM_API_KEY", llm_api_key) - self._update_env("REME_LLM_BASE_URL", llm_api_base) - self._update_env("REME_EMBEDDING_API_KEY", embedding_api_key) - self._update_env("REME_EMBEDDING_BASE_URL", embedding_api_base) - - # Use default parser if not provided - parser_class = parser if parser is not None else PydanticConfigParser - self.parser = parser_class(ServiceConfig) - - if service_config is None: - input_args = [] - if config_path: - input_args.append(f"config={config_path}") - if args: - input_args.extend(args) - if kwargs: - input_args.extend([f"{k}={v}" for k, v in kwargs.items()]) - service_config = self.parser.parse_args(*input_args) - - C.service_config = service_config - - if C.service_config.init_logger: - init_logger() - - if llm: - C.update_section_config("llm", **llm) - if embedding_model: - C.update_section_config("embedding_model", **embedding_model) - if vector_store: - C.update_section_config("vector_store", **vector_store) - if token_counter: - C.update_section_config("token_counter", **token_counter) - C.service_config.enable_logo = enable_logo - C.print_logo() - - @staticmethod - def _update_env(key: str, value: str | None): - """Update environment variable if value is provided.""" - if value: - os.environ[key] = value - - @staticmethod - async def start(): - """Initialize the service context and prepare external MCP servers.""" - C.initialize_service_context() - await C.prepare_mcp_servers() - - @staticmethod - def start_sync(): - """Synchronous version of start().""" - C.initialize_service_context() - run_coro_safely(C.prepare_mcp_servers()) - - @staticmethod - async def stop(wait_thread_pool: bool = True, wait_ray: bool = True): - """ - Stop the application and cleanup resources. - - Args: - wait_thread_pool: Whether to wait for thread pool shutdown - wait_ray: Whether to wait for Ray shutdown - """ - await C.close() - C.shutdown_thread_pool(wait=wait_thread_pool) - C.shutdown_ray(wait=wait_ray) - - @staticmethod - def stop_sync(wait_thread_pool: bool = True, wait_ray: bool = True): - """Synchronous version of stop().""" - C.close_sync() - C.shutdown_thread_pool(wait=wait_thread_pool) - C.shutdown_ray(wait=wait_ray) - - async def __aenter__(self): - """Async context manager entry.""" - await self.start() - return self - - def __enter__(self): - """Context manager entry.""" - self.start_sync() - return self - - async def __aexit__(self, exc_type=None, exc_val=None, exc_tb=None): - """Async context manager exit.""" - await self.stop() - return False - - def __exit__(self, exc_type=None, exc_val=None, exc_tb=None): - """Context manager exit.""" - self.stop_sync() - return False - - @staticmethod - async def execute_flow(name: str, **kwargs) -> Response: - """ - Execute a flow asynchronously. - - Args: - name: Name of the flow to execute - **kwargs: Arguments to pass to the flow - - Returns: - Response object from the flow execution - """ - flow: BaseFlow = C.get_flow(name) - return await flow.call(**kwargs) - - @staticmethod - def execute_flow_sync(name: str, **kwargs) -> Response: - """ - Execute a flow synchronously. - - Args: - name: Name of the flow to execute - **kwargs: Arguments to pass to the flow - - Returns: - Response object from the flow execution - """ - flow: BaseFlow = C.get_flow(name) - return flow.call_sync(**kwargs) - - @staticmethod - async def execute_stream_flow(name: str, **kwargs): - """ - Execute a streaming flow asynchronously. - - Args: - name: Name of the streaming flow to execute - **kwargs: Arguments to pass to the flow - - Yields: - Stream chunks from the flow execution - - Raises: - AssertionError: If the flow is not configured for streaming - """ - flow: BaseFlow = C.get_flow(name) - assert flow.stream is True, "non-stream flow is not supported in execute_stream_flow!" - stream_queue = asyncio.Queue() - task = asyncio.create_task(flow.call(stream_queue=stream_queue, **kwargs)) - - async for chunk in execute_stream_task( - stream_queue=stream_queue, - task=task, - task_name=name, - as_bytes=False, - ): - yield chunk - - @staticmethod - def run_service(): - """Run the configured service (HTTP, MCP, or CMD).""" - C.get_service().run() diff --git a/reme_ai/core/context/prompt_handler.py b/reme_ai/core/context/prompt_handler.py deleted file mode 100644 index e48f98ac..00000000 --- a/reme_ai/core/context/prompt_handler.py +++ /dev/null @@ -1,95 +0,0 @@ -"""Module for managing and formatting prompt templates from files or dictionaries.""" - -from pathlib import Path - -import yaml -from loguru import logger - -from .base_context import BaseContext -from .service_context import C - - -class PromptHandler(BaseContext): - """A context-aware handler for loading, retrieving, and formatting prompt templates.""" - - def __init__(self, language: str = "", **kwargs): - """Initialize the handler with a specific language and optional context data.""" - super().__init__(**kwargs) - self.language: str = language or C.language - - def load_prompt_by_file(self, prompt_file_path: Path | str = None): - """Load prompt configurations from a YAML file into the context.""" - if prompt_file_path is None: - return self - - if isinstance(prompt_file_path, str): - prompt_file_path = Path(prompt_file_path) - - if not prompt_file_path.exists(): - return self - - with prompt_file_path.open(encoding="utf-8") as f: - # Load YAML content using the full loader - prompt_dict = yaml.load(f, yaml.FullLoader) - self.load_prompt_dict(prompt_dict) - return self - - def load_prompt_dict(self, prompt_dict: dict = None): - """Merge a dictionary of prompt strings into the current context.""" - if not prompt_dict: - return self - - for key, value in prompt_dict.items(): - if isinstance(value, str): - if key in self: - logger.warning(f"Overwriting prompt key={key}, old_value={self[key]}, new_value={value}") - else: - logger.debug(f"Adding new prompt key={key}, value={value}") - self[key] = value - return self - - def get_prompt(self, prompt_name: str): - """Retrieve a prompt by name, automatically appending the language suffix if needed.""" - key: str = prompt_name - if self.language and not key.endswith(self.language.strip()): - key += "_" + self.language.strip() - - assert key in self, f"prompt_name={key} not found." - return self[key] - - def prompt_format(self, prompt_name: str, **kwargs) -> str: - """Format a prompt by filtering flagged lines and filling template variables.""" - prompt = self.get_prompt(prompt_name) - - # Separate boolean flags from string formatting arguments - flag_kwargs = {k: v for k, v in kwargs.items() if isinstance(v, bool)} - other_kwargs = {k: v for k, v in kwargs.items() if not isinstance(v, bool)} - - if flag_kwargs: - split_prompt = [] - for line in prompt.strip().split("\n"): - hit = False - hit_flag = True - for key, flag in flag_kwargs.items(): - if not line.startswith(f"[{key}]"): - continue - - hit = True - hit_flag = flag - # Remove the flag prefix from the line - line = line.strip(f"[{key}]") - break - - # Include line if no flag is present or if the flag evaluates to True - if not hit: - split_prompt.append(line) - elif hit_flag: - split_prompt.append(line) - - prompt = "\n".join(split_prompt) - - if other_kwargs: - # Apply standard Python string formatting - prompt = prompt.format(**other_kwargs) - - return prompt diff --git a/reme_ai/core/context/registry.py b/reme_ai/core/context/registry.py deleted file mode 100644 index f403037d..00000000 --- a/reme_ai/core/context/registry.py +++ /dev/null @@ -1,19 +0,0 @@ -"""Module providing a registry class for managing class-to-name mappings via decorators.""" - -from .base_context import BaseContext - - -class Registry(BaseContext): - """A registry container that uses decorators to map and store class references.""" - - def register(self, name: str = "", add_cls: bool = True): - """Return a decorator that registers a class under a specific name in the registry.""" - - def decorator(cls): - if add_cls: - # Use provided name or default to the class name as the key - key = name or cls.__name__ - self[key] = cls - return cls - - return decorator diff --git a/reme_ai/core/context/service_context.py b/reme_ai/core/context/service_context.py deleted file mode 100644 index e22147e9..00000000 --- a/reme_ai/core/context/service_context.py +++ /dev/null @@ -1,537 +0,0 @@ -"""Module for managing global service configurations and component registries via a singleton context.""" - -from concurrent.futures import ThreadPoolExecutor -from typing import TYPE_CHECKING - -from loguru import logger - -from .base_context import BaseContext -from .registry import Registry -from ..enumeration import RegistryEnum -from ..schema import ServiceConfig -from ..utils import singleton, print_logo - -if TYPE_CHECKING: - from ..llm import BaseLLM - from ..embedding import BaseEmbeddingModel - from ..vector_store import BaseVectorStore - from ..token_counter import BaseTokenCounter - from ..flow import BaseFlow - from ..service import BaseService - - -@singleton -class ServiceContext(BaseContext): - """A singleton container for global application state, thread pools, and component registries. - - This class serves as the central management hub for the entire ReMe application, providing: - - Service configuration management - - Component registration and instantiation (LLMs, embeddings, vector stores, etc.) - - Thread pool and Ray distributed computing management - - MCP (Model Context Protocol) server integration - - The singleton pattern ensures only one instance exists throughout the application lifecycle, - accessible via the global `C` variable exported at the bottom of this module. - """ - - def __init__(self, **kwargs): - """Initialize the global context with configuration objects and specialized registries. - - Sets up: - - Empty service configuration placeholder - - Thread pool for concurrent operations - - Registry dictionaries for class registration (templates) - - Instance dictionaries for instantiated objects (actual instances) - - MCP server mapping for external tool integration - """ - super().__init__(**kwargs) - - # Service configuration and runtime settings - self.service_config: ServiceConfig | None = None - self.language: str = "" - self.thread_pool: ThreadPoolExecutor | None = None - - # Registry system: stores class definitions for different component types - self.registry_dict: dict[RegistryEnum, Registry] = {v: Registry() for v in RegistryEnum.__members__.values()} - - # Instance system: stores instantiated objects created from registered classes - self.instance_dict: dict[RegistryEnum, dict] = {v: {} for v in RegistryEnum.__members__.values()} - - # MCP server mapping: maps server_name -> {tool_name: ToolCall} - self.mcp_server_mapping: dict[str, dict] = {} - - # Initialization flag: ensures initialize_service_context is called only once - self._initialized: bool = False - - def register(self, name: str, register_type: RegistryEnum): - """Return a decorator to register a component within a specific registry category. - - Args: - name: The registration name for the component (used for lookup) - register_type: The type of registry (LLM, EMBEDDING_MODEL, VECTOR_STORE, etc.) - - Returns: - A decorator function that registers the decorated class - - Example: - @C.register("my_llm", RegistryEnum.LLM) - class MyLLM(BaseLLM): - pass - """ - return self.registry_dict[register_type].register(name=name) - - def register_llm(self, name: str = ""): - """Register a Large Language Model class.""" - return self.register(name=name, register_type=RegistryEnum.LLM) - - def register_embedding_model(self, name: str = ""): - """Register an embedding model class.""" - return self.register(name=name, register_type=RegistryEnum.EMBEDDING_MODEL) - - def register_vector_store(self, name: str = ""): - """Register a vector store implementation class.""" - return self.register(name=name, register_type=RegistryEnum.VECTOR_STORE) - - def register_op(self, name: str = ""): - """Register an operation (Op) class.""" - return self.register(name=name, register_type=RegistryEnum.OP) - - def register_flow(self, name: str = ""): - """Register a workflow or logic flow class.""" - return self.register(name=name, register_type=RegistryEnum.FLOW) - - def register_service(self, name: str = ""): - """Register a backend service class.""" - return self.register(name=name, register_type=RegistryEnum.SERVICE) - - def register_token_counter(self, name: str = ""): - """Register a token counting utility class.""" - return self.register(name=name, register_type=RegistryEnum.TOKEN_COUNTER) - - def get_model_class(self, name: str, register_type: RegistryEnum): - """Retrieve a registered class by name from a specific registry category. - - Args: - name: The registration name of the class - register_type: The type of registry to search in - - Returns: - The registered class (not an instance, but the class itself) - - Raises: - AssertionError: If the class is not found in the registry - """ - assert name in self.registry_dict[register_type], f"{name} not in registry_dict[{register_type}]" - return self.registry_dict[register_type][name] - - def get_llm_class(self, name: str): - """Get the LLM class registered under the given name.""" - return self.get_model_class(name, RegistryEnum.LLM) - - def get_embedding_model_class(self, name: str): - """Get the embedding model class registered under the given name.""" - return self.get_model_class(name, RegistryEnum.EMBEDDING_MODEL) - - def get_vector_store_class(self, name: str): - """Get the vector store class registered under the given name.""" - return self.get_model_class(name, RegistryEnum.VECTOR_STORE) - - def get_op_class(self, name: str): - """Get the operation class registered under the given name.""" - return self.get_model_class(name, RegistryEnum.OP) - - def get_flow_class(self, name: str): - """Get the flow class registered under the given name.""" - return self.get_model_class(name, RegistryEnum.FLOW) - - def get_service_class(self, name: str): - """Get the service class registered under the given name.""" - return self.get_model_class(name, RegistryEnum.SERVICE) - - def get_token_counter_class(self, name: str): - """Get the token counter class registered under the given name.""" - return self.get_model_class(name, RegistryEnum.TOKEN_COUNTER) - - def get_llm(self, name: str) -> "BaseLLM": - """Retrieve a specific LLM instance by name. - - Args: - name: The name of the LLM instance (typically 'default' or custom name) - - Returns: - The instantiated LLM object - - Raises: - KeyError: If no LLM with the given name exists - """ - return self.instance_dict[RegistryEnum.LLM][name] - - def get_embedding_model(self, name: str) -> "BaseEmbeddingModel": - """Retrieve a specific embedding model instance by name. - - Args: - name: The name of the embedding model instance (typically 'default') - - Returns: - The instantiated embedding model object - - Raises: - KeyError: If no embedding model with the given name exists - """ - return self.instance_dict[RegistryEnum.EMBEDDING_MODEL][name] - - def get_vector_store(self, name: str) -> "BaseVectorStore": - """Retrieve a specific vector store instance by name. - - Args: - name: The name of the vector store instance (typically 'default') - - Returns: - The instantiated vector store object - - Raises: - KeyError: If no vector store with the given name exists - """ - return self.instance_dict[RegistryEnum.VECTOR_STORE][name] - - def get_token_counter(self, name: str) -> "BaseTokenCounter": - """Retrieve a specific token counter instance by name. - - Args: - name: The name of the token counter instance (typically 'default') - - Returns: - The instantiated token counter object - - Raises: - KeyError: If no token counter with the given name exists - """ - return self.instance_dict[RegistryEnum.TOKEN_COUNTER][name] - - def get_flow(self, name: str) -> "BaseFlow": - """Retrieve a specific flow instance by name. - - Args: - name: The name of the flow instance - - Returns: - The instantiated flow object - - Raises: - KeyError: If no flow with the given name exists - """ - return self.instance_dict[RegistryEnum.FLOW][name] - - def get_service(self) -> "BaseService": - """Retrieve the default service instance. - - Returns: - The instantiated service backend (HTTP, MCP, or CMD service) - - Raises: - KeyError: If the default service was not initialized - """ - return self.instance_dict[RegistryEnum.SERVICE]["default"] - - def update_section_config(self, section_name: str, **kwargs): - """Update a specific section of the service config with new values. - - Args: - section_name: Name of the config section (e.g., 'llm', 'embedding_model') - **kwargs: Key-value pairs to update in the default configuration - - Raises: - KeyError: If the default config for the section doesn't exist - - Example: - update_section_config('llm', temperature=0.8, max_tokens=1000) - """ - if not hasattr(self.service_config, section_name) or not kwargs: - return - - section_dict: dict = getattr(self.service_config, section_name) - if "default" not in section_dict: - raise KeyError(f"Default `{section_name}` config not found") - - current_config = section_dict["default"] - section_dict["default"] = current_config.model_copy(update=kwargs, deep=True) - - def initialize_service_context(self): - """Initialize the service context with the configuration. - - This is the main initialization method that sets up all system components in order: - 1. Language settings - 2. Thread pool for concurrent operations - 3. Ray cluster (if configured for distributed computing) - 4. LLM instances - 5. Embedding model instances - 6. Token counter instances - 7. Vector store instances (with their embedding models) - 8. Flow instances (both registered and configured) - 9. Service backend instance - - Note: This method should be called after service_config is set. - This method can only be called once. Subsequent calls will be ignored. - """ - if self._initialized: - logger.warning("initialize_service_context has already been called. Skipping re-initialization.") - return - - self.language = self.service_config.language - self.thread_pool = ThreadPoolExecutor(max_workers=self.service_config.thread_pool_max_workers) - - # Initialize Ray for distributed computing if configured - if self.service_config.ray_max_workers > 1: - import ray - - ray.init(num_cpus=self.service_config.ray_max_workers) - - # Initialize components in dependency order - self._initialize_llm() - self._initialize_embedding_model() - self._initialize_token_counter() - self._initialize_vector_store() # Depends on embedding models - self._initialize_flow() - self._initialize_service() - - # Mark as initialized - self._initialized = True - - def _initialize_llm(self): - """Initialize all configured LLM instances. - - For each LLM configuration: - - Retrieves the corresponding registered LLM class by backend name - - Instantiates it with model_name and additional configuration - - Stores the instance in instance_dict for later retrieval - """ - for name, config in self.service_config.llm.items(): - llm_cls = self.get_llm_class(config.backend) - self.instance_dict[RegistryEnum.LLM][name] = llm_cls(model_name=config.model_name, **config.model_extra) - - def _initialize_embedding_model(self): - """Initialize all configured embedding model instances. - - For each embedding model configuration: - - Retrieves the corresponding registered embedding model class by backend name - - Instantiates it with model_name and additional configuration - - Stores the instance in instance_dict for later retrieval - """ - for name, config in self.service_config.embedding_model.items(): - embedding_model_cls = self.get_embedding_model_class(config.backend) - self.instance_dict[RegistryEnum.EMBEDDING_MODEL][name] = embedding_model_cls( - model_name=config.model_name, - **config.model_extra, - ) - - def _initialize_token_counter(self): - """Initialize all configured token counter instances. - - For each token counter configuration: - - Retrieves the corresponding registered token counter class by backend name - - Instantiates it with model_name and additional configuration - - Stores the instance in instance_dict for later retrieval - """ - for name, config in self.service_config.token_counter.items(): - token_counter_cls = self.get_token_counter_class(config.backend) - self.instance_dict[RegistryEnum.TOKEN_COUNTER][name] = token_counter_cls( - model_name=config.model_name, - **config.model_extra, - ) - - def _initialize_vector_store(self): - """Initialize all configured vector stores with their embedding models. - - For each vector store configuration: - - Retrieves the corresponding registered vector store class by backend name - - Retrieves the associated embedding model instance by name - - Instantiates the vector store with collection name, embedding model, and extra config - - Stores the instance in instance_dict for later retrieval - - Note: This must be called after _initialize_embedding_model() since vector stores - depend on embedding model instances. - """ - for name, config in self.service_config.vector_store.items(): - vector_store_cls = self.get_vector_store_class(config.backend) - self.instance_dict[RegistryEnum.VECTOR_STORE][name] = vector_store_cls( - collection_name=config.collection_name, - embedding_model=self.instance_dict[RegistryEnum.EMBEDDING_MODEL][config.embedding_model], - **config.model_extra, - ) - - def _filter_flows(self, name: str) -> bool: - """Filter flows based on enabled_flows and disabled_flows configuration. - - The filtering logic follows this priority: - 1. If enabled_flows is set: only flows in the list are loaded - 2. Else if disabled_flows is set: all flows except those in the list are loaded - 3. Otherwise: all flows are loaded - - Args: - name: The flow name to check - - Returns: - True if the flow should be loaded, False otherwise - """ - if self.service_config.enabled_flows: - return name in self.service_config.enabled_flows - elif self.service_config.disabled_flows: - return name not in self.service_config.disabled_flows - else: - return True - - def _initialize_flow(self): - """Initialize all flows from both registry and configuration. - - Flows can be defined in two ways: - 1. Registered flows: Python classes decorated with @register_flow - 2. Configuration flows: Defined in config as ExpressionFlow instances - - Process: - 1. First, instantiate all registered flow classes (from decorators) - - Filter based on enabled_flows/disabled_flows - - Create instance with the flow name - - 2. Then, instantiate all configured flows (from config file) - - Filter based on enabled_flows/disabled_flows - - Create ExpressionFlow instances with flow configuration - - Note: Configuration flows can override registered flows with the same name. - """ - - # Initialize flows from registry (decorator-based registration) - for name, flow_cls in self.registry_dict[RegistryEnum.FLOW].items(): - if not self._filter_flows(name): - continue - flow: "BaseFlow" = flow_cls(name=name) - self.instance_dict[RegistryEnum.FLOW][flow.name] = flow - - # Initialize flows from configuration (config-based definition) - from ..flow import ExpressionFlow - - for name, flow_config in self.service_config.flow.items(): - if not self._filter_flows(name): - continue - flow_config.name = name - flow: BaseFlow = ExpressionFlow(flow_config=flow_config) - self.instance_dict[RegistryEnum.FLOW][name] = flow - - def _initialize_service(self): - """Initialize the service backend instance. - - Creates an instance of the configured service backend (e.g., HTTP, MCP, or CMD service) - and stores it in the instance dictionary under the 'default' key. - """ - service_cls = self.get_service_class(self.service_config.backend) - self.instance_dict[RegistryEnum.SERVICE]["default"] = service_cls() - - async def prepare_mcp_servers(self): - """Prepare and initialize MCP (Model Context Protocol) server connections. - - This method: - 1. Checks if MCP servers are configured - 2. Creates an MCP client instance - 3. For each configured server: - - Lists available tool calls from the server - - Builds a mapping of tool_name -> ToolCall object - - Logs available tools for debugging - - The mcp_server_mapping is structured as: - { - "server_name": { - "tool_name": ToolCall(...), - ... - }, - ... - } - - This allows the application to discover and use external tools provided by MCP servers. - """ - if not self.service_config.mcp_servers: - return - - from ..utils import MCPClient - - mcp_client = MCPClient(config={"mcpServers": self.service_config.mcp_servers}) - for server_name in self.service_config.mcp_servers.keys(): - try: - # Retrieve all available tool calls from this MCP server - tool_calls = await mcp_client.list_tool_calls(server_name=server_name, return_dict=False) - - # Build mapping: tool_name -> ToolCall for quick lookup - self.mcp_server_mapping[server_name] = {tool_call.name: tool_call for tool_call in tool_calls} - - # Log discovered tools for debugging - for tool_call in tool_calls: - logger.info(f"list_tool_calls: {server_name}@{tool_call.name} {tool_call.simple_input_dump()}") - - except Exception as e: - logger.exception(f"list_tool_calls: {server_name} error: {e}") - - def print_logo(self): - """Print the ReMe logo if enabled in configuration.""" - if self.service_config.enable_logo: - print_logo(service_config=self.service_config) - - async def close(self): - """Close all service components asynchronously. - - Gracefully closes all instantiated components in order: - 1. Vector stores (closes database connections) - 2. LLMs (closes API clients and connections) - 3. Embedding models (closes API clients and connections) - - This method should be called when shutting down the application - to ensure all resources are properly released. - """ - for _, vector_store in self.instance_dict[RegistryEnum.VECTOR_STORE].items(): - await vector_store.close() - - for _, llm in self.instance_dict[RegistryEnum.LLM].items(): - await llm.close() - - for _, embedding_model in self.instance_dict[RegistryEnum.EMBEDDING_MODEL].items(): - await embedding_model.close() - - def close_sync(self): - """Close all service components synchronously. - - Synchronous version of close() for non-async contexts. - Closes LLMs and embedding models without using async/await. - - Note: Vector stores are not closed here as they typically require async operations. - """ - for _, llm in self.instance_dict[RegistryEnum.LLM].items(): - llm.close_sync() - - for _, embedding_model in self.instance_dict[RegistryEnum.EMBEDDING_MODEL].items(): - embedding_model.close_sync() - - def shutdown_thread_pool(self, wait: bool = True): - """Shutdown the thread pool executor. - - Args: - wait: If True, blocks until all pending futures are executed. - If False, returns immediately and pending futures may be cancelled. - """ - if self.thread_pool: - self.thread_pool.shutdown(wait=wait) - - def shutdown_ray(self, wait: bool = True): - """Shutdown Ray cluster if it was initialized. - - Args: - wait: If True, waits for Ray to fully shutdown. - If False, returns immediately without waiting. - - Note: Only shuts down Ray if it was configured with ray_max_workers > 1. - """ - if self.service_config and self.service_config.ray_max_workers > 1: - import ray - - ray.shutdown(_exiting_interpreter=not wait) - - -# Export a global singleton instance for easy access across the application -# This is the primary way to access the service context throughout the codebase -C = ServiceContext() diff --git a/reme_ai/core/enumeration/json_schema_enum.py b/reme_ai/core/enumeration/json_schema_enum.py deleted file mode 100644 index 507645f4..00000000 --- a/reme_ai/core/enumeration/json_schema_enum.py +++ /dev/null @@ -1,18 +0,0 @@ -"""Defines the standard data types supported by JSON Schema.""" - -from enum import Enum - - -class JsonSchemaEnum(Enum): - """Enumeration of valid JSON Schema data types.""" - - STRING = str - NUMBER = float - INTEGER = int - OBJECT = dict - ARRAY = list - BOOLEAN = bool - - def __str__(self) -> str: - """Returns the string representation of the enum value.""" - return self.name.lower() diff --git a/reme_ai/core/enumeration/memory_type.py b/reme_ai/core/enumeration/memory_type.py deleted file mode 100644 index 22d35481..00000000 --- a/reme_ai/core/enumeration/memory_type.py +++ /dev/null @@ -1,25 +0,0 @@ -"""Memory type enumeration for the three-layer memory architecture.""" - -from enum import Enum - - -class MemoryType(str, Enum): - """ - Three-layer memory architecture for agent memory management. - - Layer 1 - High-level Abstraction Memory: - - IDENTITY: Self-cognition (identity, personality, current state) - - PERSONAL: Person-specific memory (preferences and context about specific individuals) - - PROCEDURAL: Procedural memory (how-to knowledge, e.g., 4 steps to write financial reports) - - TOOL: Tool memory (tool usage patterns, success rates, token consumption, latency) - - Layer 2 - Summary Memory (Compressed): Summarized digest of raw message history - Layer 3 - History Memory (Raw): Raw message history - """ - - IDENTITY = "identity" - PERSONAL = "personal" - PROCEDURAL = "procedural" - TOOL = "tool" - SUMMARY = "summary" - HISTORY = "history" diff --git a/reme_ai/core/flow/simple_flow.py b/reme_ai/core/flow/simple_flow.py deleted file mode 100644 index 51a1381d..00000000 --- a/reme_ai/core/flow/simple_flow.py +++ /dev/null @@ -1,17 +0,0 @@ -"""Simple flow implementation that directly uses a predefined flow operation.""" - -from .base_flow import BaseFlow -from ..op import BaseOp -from ..schema import ToolCall - - -class SimpleFlow(BaseFlow): - """Simple flow that directly uses a predefined flow operation.""" - - def _build_flow(self) -> BaseOp: - assert self._flow_op is not None - return self._flow_op.copy() - - def _build_tool_call(self) -> ToolCall: - assert self._flow_op is not None - return self._flow_op.tool_call diff --git a/reme_ai/core/main.py b/reme_ai/core/main.py deleted file mode 100644 index 81811c1f..00000000 --- a/reme_ai/core/main.py +++ /dev/null @@ -1,48 +0,0 @@ -"""ReMe application classes for simplified configuration and execution.""" - -import sys - -from .application import Application -from .config import ReMeConfigParser - - -class ReMeApp(Application): - """ReMe application with config file support and flow execution methods.""" - - def __init__( - self, - *args, - llm_api_key: str | None = None, - llm_api_base: str | None = None, - embedding_api_key: str | None = None, - embedding_api_base: str | None = None, - config_path: str | None = None, - enable_logo: bool = True, - **kwargs, - ): - super().__init__( - *args, - llm_api_key=llm_api_key, - llm_api_base=llm_api_base, - embedding_api_key=embedding_api_key, - embedding_api_base=embedding_api_base, - service_config=None, - parser=ReMeConfigParser, - config_path=config_path, - enable_logo=enable_logo, - **kwargs, - ) - - async def async_execute(self, name: str, **kwargs) -> dict: - """Execute a flow asynchronously and return the result as a dictionary.""" - return (await self.execute_flow(name=name, **kwargs)).model_dump() - - -def main(): - """Main entry point for running ReMe application from command line.""" - with ReMeApp(*sys.argv[1:]) as app: - app.run_service() - - -if __name__ == "__main__": - main() diff --git a/reme_ai/mem_agent/base_memory_agent.py b/reme_ai/mem_agent/base_memory_agent.py index fb52dc0d..7c21bb19 100644 --- a/reme_ai/mem_agent/base_memory_agent.py +++ b/reme_ai/mem_agent/base_memory_agent.py @@ -22,7 +22,7 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): tools: list[BaseMemoryTool], add_think_tool: bool = False, # only for instruct model tool_call_interval: float = 0, - max_steps: int = 20, + max_steps: int = 8, **kwargs, ): tools = tools or [] @@ -35,10 +35,11 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): self.max_steps: int = max_steps self.messages: list[Message] = [] + self.tool_messages: list[Message] = [] self.success: bool = True - self.retrieved_nodes: list[MemoryNode] = [] self.memory_nodes: list[MemoryNode | str] = [] + self.meta_info: str = "" def _build_tool_call(self) -> ToolCall: return ToolCall( @@ -97,35 +98,37 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): """Builds and returns the initial messages for the agent.""" return self.get_messages() - async def _reasoning_step(self, messages: list[Message], step: int, **kwargs) -> tuple[Message, bool]: + async def _reasoning_step(self, messages: list[Message], step: int, stage: str = "", **kwargs) -> tuple[Message, bool]: assistant_message: Message = await self.llm.chat( messages=messages, tools=[t.tool_call for t in self.tools], **kwargs, ) messages.append(assistant_message) + stage_prefix = f"-{stage}" if stage else "" logger.info( - f"[{self.__class__.__name__}] " + f"[{self.__class__.__name__}{stage_prefix}] " f"step{step + 1}.assistant={assistant_message.simple_dump(enable_json_dump=True)}", ) should_act = bool(assistant_message.tool_calls) return assistant_message, should_act - async def _acting_step(self, assistant_message: Message, step: int, **kwargs) -> list[Message]: + async def _acting_step(self, assistant_message: Message, step: int, stage: str = "", **kwargs) -> list[Message]: if not assistant_message.tool_calls: return [] tool_list: list[BaseMemoryTool] = [] tool_result_messages: list[Message] = [] tool_dict = {t.tool_call.name: t for t in self.tools} + stage_prefix = f"-{stage}" if stage else "" for j, tool_call in enumerate(assistant_message.tool_calls): if tool_call.name not in tool_dict: - logger.warning(f"[{self.__class__.__name__}] unknown tool_call.name={tool_call.name}") + logger.warning(f"[{self.__class__.__name__}{stage_prefix}] unknown tool_call.name={tool_call.name}") continue logger.info( - f"[{self.__class__.__name__}] step{step + 1}.{j} " + f"[{self.__class__.__name__}{stage_prefix}] step{step + 1}.{j} " f"submit tool_calls={tool_call.name} argument={tool_call.arguments}", ) tool_copy: BaseMemoryTool = tool_dict[tool_call.name].copy() @@ -143,7 +146,7 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): self.memory_nodes.extend(op.memory_nodes) if hasattr(op, "messages") and op.messages: - self.messages.extend(op.messages) + self.tool_messages.extend(op.messages) tool_result = str(op.output) tool_message = Message( @@ -152,20 +155,27 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): tool_call_id=op.tool_call.id, ) tool_result_messages.append(tool_message) - logger.info(f"[{self.__class__.__name__}] step{step + 1}.{j} join tool_result={tool_result[:500]}...\n\n") + + # # Collect tool call information to meta_info + # tool_info = f"\n## Tool Call {step + 1}.{j + 1}: {op.tool_call.name}\n" + # tool_info += f"Arguments: {json.dumps(assistant_message.tool_calls[j].argument_dict, ensure_ascii=False)}\n" + # tool_info += f"Result: {tool_result}\n" + self.meta_info += tool_result + "\n" + + logger.info(f"[{self.__class__.__name__}{stage_prefix}] step{step + 1}.{j} join tool_result={tool_result[:2000]}...\n\n") return tool_result_messages - async def react(self, messages: list[Message]): + async def react(self, messages: list[Message], stage: str = ""): """Performs reasoning and acting steps until completion or max steps reached.""" success: bool = False for step in range(self.max_steps): - assistant_message, should_act = await self._reasoning_step(messages, step) + assistant_message, should_act = await self._reasoning_step(messages, step, stage=stage) if not should_act: success = True break - tool_result_messages = await self._acting_step(assistant_message, step) + tool_result_messages = await self._acting_step(assistant_message, step, stage=stage) messages.extend(tool_result_messages) return messages, success @@ -209,3 +219,8 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): def author(self) -> str: """Returns the LLM model name as the author identifier.""" return self.llm.model_name + + @property + def history_node(self): + """Returns the history node.""" + return self.context.get("history_node", None) \ No newline at end of file diff --git a/reme_ai/mem_agent/chat/remy_agent.py b/reme_ai/mem_agent/chat/remy_agent.py deleted file mode 100644 index 1eaa9ac7..00000000 --- a/reme_ai/mem_agent/chat/remy_agent.py +++ /dev/null @@ -1,51 +0,0 @@ -"""ReMy agent with identity and meta memory capabilities.""" - -from typing import List - -from ..base_memory_agent import BaseMemoryAgent -from ...core.context import C -from ...core.enumeration import Role -from ...core.schema import Message -from ...core.utils import get_now_time - - -@C.register_op() -class ReMyAgent(BaseMemoryAgent): - """Memory agent with identity awareness and meta memory retrieval.""" - - def __init__(self, enable_tool_memory: bool = True, enable_identity_memory: bool = True, **kwargs): - """Initialize ReMy agent with memory options.""" - super().__init__(**kwargs) - self.enable_tool_memory = enable_tool_memory - self.enable_identity_memory = enable_identity_memory - - @staticmethod - async def _read_identity_memory() -> str: - """Read and return identity memory as string.""" - from ...mem_tool import ReadIdentityMemory - - op = ReadIdentityMemory() - await op.call() - return str(op.output) - - async def _read_meta_memories(self) -> str: - """Read and return meta memories as string.""" - from ...mem_tool import ReadMetaMemory - - op = ReadMetaMemory( - enable_tool_memory=self.enable_tool_memory, - enable_identity_memory=self.enable_identity_memory, - ) - await op.call() - return str(op.output) - - async def build_messages(self) -> List[Message]: - """Build messages with system prompt and user messages.""" - system_prompt = self.prompt_format( - prompt_name="system_prompt", - now_time=get_now_time(), - identity_memory=await self._read_identity_memory(), - meta_memory_info=await self._read_meta_memories(), - ) - - return [Message(role=Role.SYSTEM, content=system_prompt)] + self.get_messages() diff --git a/reme_ai/mem_agent/chat/remy_agent.yaml b/reme_ai/mem_agent/chat/remy_agent.yaml deleted file mode 100644 index 5498b3cc..00000000 --- a/reme_ai/mem_agent/chat/remy_agent.yaml +++ /dev/null @@ -1,36 +0,0 @@ -tool: | - Conversational AI assistant with integrated memory capabilities. - Use this tool to engage in natural conversations with users while leveraging - stored identity and memory context. The agent can access historical information, - user preferences, and procedural knowledge through its memory system, and can - use various tools to accomplish tasks and answer questions. - -system_prompt: | - You are ReMy, an intelligent AI assistant with memory capabilities. - - ## Current Time - {now_time} - - ## Self-Awareness - {identity_memory} - - ## Available Meta Memories - Format: "- (): " - {meta_memory_info} - - ## Guiding Principles - 1. **Be Helpful and Accurate**: Provide clear and correct information. - 2. **Use Memory Wisely**: Retrieve relevant memories when they can improve your response. - 3. **Use Tools Appropriately**: Select the right tool for each task. - 4. **Stay Conversational**: Maintain a natural and friendly tone. - 5. **Seek Clarification**: Ask questions if the user’s intent is unclear. - 6. **Acknowledge Limitations**: Be honest about what you can and cannot do. - - ## How to Use the Memory Retrieval Tool - When using `vector_retrieve_memory` to search memories: - - Choose an appropriate `memory_type` and `memory_target` from the "Available Meta Memories" list above. - - Formulate a clear and specific query based on the information you need. - - **Important**: When retrieving tool-related memories (`memory_type` is "tool"), the query must use the tool’s exact name (not a description or a question). - - If retrieval results include a `ref_memory_id` and you need more details, use `read_history_memory` with the `ref_memory_id` as the `memory_id` parameter. - - If the initial retrieval yields no results, try rephrasing your query or using a different memory type. - - You may generate multiple queries with different phrasings or perspectives for the same memory type/target. diff --git a/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.py b/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.py index ee5c0aae..3b934172 100644 --- a/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.py +++ b/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.py @@ -12,7 +12,7 @@ from ...core.utils import format_messages @C.register_op() class ReMeRetrieverV2(BaseMemoryAgent): """Memory agent that autonomously retrieves memories from multiple angles. - + This retriever: - Directly queries memories based on user questions without time constraints - Tries multiple retrieval strategies: direct vector search, metadata filtering, partial filtering @@ -24,13 +24,13 @@ class ReMeRetrieverV2(BaseMemoryAgent): # Check if ReadHistory tool is available in the tools list tools = kwargs.get('tools', []) has_read_history = any(tool.__class__.__name__ == 'ReadHistory' for tool in tools) - + # Use simple prompt if ReadHistory is not available if not has_read_history: super().__init__(prompt_name="reme_retriever_v2_simple", **kwargs) else: super().__init__(**kwargs) - + self.meta_memories: list[dict] = meta_memories or [] async def _read_meta_memories(self) -> str: diff --git a/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.yaml b/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.yaml index 281796d6..755fd6b7 100644 --- a/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.yaml +++ b/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.yaml @@ -24,16 +24,16 @@ system_prompt: | 1. **Multi-Angle Vector Retrieval** (REQUIRED - at least 3 attempts): You must try AT LEAST 3 different retrieval approaches using `retrieve_memories`: - + a) **Direct Vector Search**: Use the user's question directly or with minimal reformulation - Query the most relevant memory_type and memory_target - Use straightforward query phrasing - + b) **Alternative Phrasing**: Reformulate the query from a different angle - Use synonyms or different expressions - Break down complex questions into simpler components - Try more specific or more general queries - + c) **Metadata-Filtered Search**: Add metadata filters to narrow down results - **Time-based filtering**: Use year/month/day metadata fields to filter by time periods * Example: {{"year": 2024}} for memories from 2024 @@ -41,33 +41,33 @@ system_prompt: | * Example: {{"year": 2024, "month": 5, "day": 15}} for memories from a specific date - Combine vector search with metadata constraints - Try partial metadata filtering if full filtering yields nothing (e.g., only year, or year+month) - + d) **Cross-Memory-Type Search**: If applicable, search across different memory types - Try different memory_type and memory_target combinations - Some information might be stored in unexpected memory categories - + e) **Keyword Extraction**: Extract key entities/concepts and search for them - Identify important names, places, concepts - Search for each key element separately - + 2. **Evaluate Retrieval Results** (After each attempt): - Review what memories were returned - Assess if they contain sufficient information to answer the question - If insufficient, identify what's missing and adjust your next query accordingly - Track which retrieval strategies you've already tried - + 3. **Persist Through Failures**: - DO NOT give up after 1-2 failed attempts - If a retrieval returns no results or irrelevant results, try a different approach - Consider that the information might be phrased differently than expected - Be creative with query reformulation - + 4. **Fallback to History Reading** (Only after 3+ vector retrieval attempts): - If after at least 3 different vector retrieval attempts you still lack sufficient information: * If any retrieved memories contain `ref_memory_id`, use `read_history` to read the original conversation * Use `read_history` with the `ref_memory_id` to get complete context * This can reveal details that weren't captured in the memory summaries - + 5. **Answer the Question**: - Once you have sufficient information, provide a direct answer based ONLY on retrieved memories - DO NOT fabricate, guess, or infer information not present in the memories @@ -95,30 +95,30 @@ system_prompt: | **Example 1: Simple Query** Attempt 1: Direct query "user's favorite food" → Result: No relevant memories found - + Attempt 2: Reformulated query "what does user like to eat" → Result: Some memories about meals, but not specific preferences - + Attempt 3: Keyword search "food preferences" with metadata filter → Result: Found relevant memory with ref_memory_id - + Attempt 4: Use read_history with ref_memory_id to get full context → Result: Found detailed conversation about favorite foods - + Answer: [Provide answer based on retrieved information] **Example 2: Time-based Query** Question: "What did the user do last summer?" - + Attempt 1: Direct query "user activities summer" with metadata {{"year": 2025, "month": [6, 7, 8]}} → Result: Found some vacation memories - + Attempt 2: Broader query "user summer vacation travel" with metadata {{"year": 2025}} → Result: Found additional travel-related memories - + Attempt 3: Use read_history for memories with ref_memory_id to get detailed context → Result: Complete picture of summer activities - + Answer: [Provide answer based on retrieved information] user_message: | diff --git a/reme_ai/mem_agent/retriever_v2/reme_retriever_v2_simple.yaml b/reme_ai/mem_agent/retriever_v2/reme_retriever_v2_simple.yaml index f8a4b7f5..cb9cc578 100644 --- a/reme_ai/mem_agent/retriever_v2/reme_retriever_v2_simple.yaml +++ b/reme_ai/mem_agent/retriever_v2/reme_retriever_v2_simple.yaml @@ -23,16 +23,16 @@ system_prompt: | 1. **Multi-Angle Vector Retrieval** (REQUIRED - at least 3 attempts): You must try AT LEAST 3 different retrieval approaches using `retrieve_memories`: - + a) **Direct Vector Search**: Use the user's question directly or with minimal reformulation - Query the most relevant memory_type and memory_target - Use straightforward query phrasing - + b) **Alternative Phrasing**: Reformulate the query from a different angle - Use synonyms or different expressions - Break down complex questions into simpler components - Try more specific or more general queries - + c) **Metadata-Filtered Search**: Add metadata filters to narrow down results - **Time-based filtering**: Use year/month/day metadata fields to filter by time periods * Example: {{"year": 2024}} for memories from 2024 @@ -40,27 +40,27 @@ system_prompt: | * Example: {{"year": 2024, "month": 5, "day": 15}} for memories from a specific date - Combine vector search with metadata constraints - Try partial metadata filtering if full filtering yields nothing (e.g., only year, or year+month) - + d) **Cross-Memory-Type Search**: If applicable, search across different memory types - Try different memory_type and memory_target combinations - Some information might be stored in unexpected memory categories - + e) **Keyword Extraction**: Extract key entities/concepts and search for them - Identify important names, places, concepts - Search for each key element separately - + 2. **Evaluate Retrieval Results** (After each attempt): - Review what memories were returned - Assess if they contain sufficient information to answer the question - If insufficient, identify what's missing and adjust your next query accordingly - Track which retrieval strategies you've already tried - + 3. **Persist Through Failures**: - DO NOT give up after 1-2 failed attempts - If a retrieval returns no results or irrelevant results, try a different approach - Consider that the information might be phrased differently than expected - Be creative with query reformulation - + 4. **Answer the Question**: - Once you have sufficient information, provide a direct answer based ONLY on retrieved memories - DO NOT fabricate, guess, or infer information not present in the memories @@ -88,27 +88,27 @@ system_prompt: | **Example 1: Simple Query** Attempt 1: Direct query "user's favorite food" → Result: No relevant memories found - + Attempt 2: Reformulated query "what does user like to eat" → Result: Some memories about meals, but not specific preferences - + Attempt 3: Keyword search "food preferences" with metadata filter → Result: Found relevant memory - + Answer: [Provide answer based on retrieved information] **Example 2: Time-based Query** Question: "What did the user do last summer?" - + Attempt 1: Direct query "user activities summer" with metadata {{"year": 2025, "month": [6, 7, 8]}} → Result: Found some vacation memories - + Attempt 2: Broader query "user summer vacation travel" with metadata {{"year": 2025}} → Result: Found additional travel-related memories - + Attempt 3: More specific queries about specific activities → Result: Complete picture of summer activities - + Answer: [Provide answer based on retrieved information] user_message: | diff --git a/reme_ai/mem_agent/summarizer/reme_summarizer.py b/reme_ai/mem_agent/summarizer/reme_summarizer.py index 9eb06d84..b2418db5 100644 --- a/reme_ai/mem_agent/summarizer/reme_summarizer.py +++ b/reme_ai/mem_agent/summarizer/reme_summarizer.py @@ -18,19 +18,19 @@ class ReMeSummarizer(BaseMemoryAgent): super().__init__(**kwargs) self.enable_identity_memory = enable_identity_memory self.meta_memories: list[dict] = meta_memories or [] - + # Check if AddMetaMemory is in tools self.enable_add_meta_memory = self._check_add_meta_memory_in_tools() def _check_add_meta_memory_in_tools(self) -> bool: """Check if AddMetaMemory tool is present in the tools list.""" from ...mem_tool import AddMetaMemory - + for tool in self.tools: if isinstance(tool, AddMetaMemory): return True return False - + def _build_tool_call(self) -> ToolCall: return ToolCall( **{ diff --git a/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.py b/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.py index 13395212..17bf4c66 100644 --- a/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.py +++ b/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.py @@ -12,7 +12,7 @@ class PersonalSummarizerV2(BaseMemoryAgent): memory_type: MemoryType = MemoryType.PERSONAL """Simplified personal memory summarizer that uses v2 memory tools. - + This summarizer follows a three-step workflow: 1. AddMemoryDrafts: Generate initial memory drafts from context 2. RetrieveRecentAndSimilarMemories: Retrieve similar and recent memories diff --git a/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2_simple.yaml b/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2_simple.yaml index acc5bd35..0a5347bb 100644 --- a/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2_simple.yaml +++ b/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2_simple.yaml @@ -10,7 +10,7 @@ system_prompt: | ## Context: {context} - + **Context Format Explanation**: The context contains formatted conversation messages in the following structure: - Each message is formatted as: `round{index} [{timestamp}] {role/name}: {content}` diff --git a/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.yaml b/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.yaml index 30792a08..58080cab 100644 --- a/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.yaml +++ b/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.yaml @@ -18,7 +18,7 @@ system_prompt: | 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: | 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..95907612 --- /dev/null +++ b/reme_ai/mem_agent/v3/personal_summarizer_v3.yaml @@ -0,0 +1,42 @@ +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 + - **Format**: Use third-person perspective to record what **{memory_target}** said, did, or expressed at specific times + - **Consolidation**: Merge related information under the same topic into ONE memory entry + - Group similar facts (e.g., multiple food preferences → one food preference entry) + - Avoid creating separate entries for closely related information + - 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..925bee68 --- /dev/null +++ b/reme_ai/mem_agent/v3/reme_retriever_v3.yaml @@ -0,0 +1,50 @@ +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 agent. Search for relevant memories to answer the user's question following this strategy: + + ## Available Meta Memories + Format: "- (): " + {meta_memory_info} + + ## User's Question + {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 + + **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 + - If nothing found after all three steps: State clearly "I don't know. " + - Be persistent: try multiple angles in each step before moving to the next + +user_message: | + 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..58080cab --- /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_agent/v4/__init__.py b/reme_ai/mem_agent/v4/__init__.py new file mode 100644 index 00000000..bf81501e --- /dev/null +++ b/reme_ai/mem_agent/v4/__init__.py @@ -0,0 +1,11 @@ +from .reme_summarizer_v4 import ReMeSummarizerV4 +from .reme_retriever_v4 import ReMeRetrieverV4 +from .personal_summarizer_v4 import PersonalSummarizerV4 +from .personal_retriever_v4 import PersonalRetrieverV4 + +__all__ = [ + "ReMeSummarizerV4", + "ReMeRetrieverV4", + "PersonalSummarizerV4", + "PersonalRetrieverV4", +] diff --git a/reme_ai/mem_agent/v4/personal_retriever_v4.py b/reme_ai/mem_agent/v4/personal_retriever_v4.py new file mode 100644 index 00000000..dedcf99a --- /dev/null +++ b/reme_ai/mem_agent/v4/personal_retriever_v4.py @@ -0,0 +1,52 @@ +from ..base_memory_agent import BaseMemoryAgent +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message +from ...core.utils import format_messages +from ...mem_tool.v4 import ReadUserProfile + + +class PersonalRetrieverV4(BaseMemoryAgent): + memory_type: MemoryType = MemoryType.PERSONAL + + async def build_messages(self) -> list[Message]: + context = self.context.query if self.context.get("query") else format_messages(self.context.messages) if self.context.get("messages") else None + if not context: + raise ValueError("input must have either `query` or `messages`") + + read_profile_tool = ReadUserProfile(show_ids="history") + await read_profile_tool.call(memory_type=self.memory_type.value, memory_target=self.memory_target) + self.context.user_profile = user_profile = read_profile_tool.output + + return [ + Message( + role=Role.USER, + content=self.prompt_format( + prompt_name="user_message", + memory_type=self.memory_type.value, + memory_target=self.memory_target, + user_profile=user_profile, + context=context, + )) + ] + + async def _acting_step(self, assistant_message: Message, step: int, **kwargs) -> list[Message]: + return await super()._acting_step( + assistant_message, + step, + memory_type=self.memory_type.value, + memory_target=self.memory_target, + **kwargs, + ) + + async def execute(self): + """Execute the retriever and determine success based on output markers.""" + await super().execute() + + # Check for memory found/not found markers in the output + if self.output: + if "" in self.output: + self.success = True + elif "" in self.output: + self.success = False + + self.meta_info = self.context.user_profile + "\n" + self.meta_info \ No newline at end of file diff --git a/reme_ai/mem_agent/v4/personal_retriever_v4.yaml b/reme_ai/mem_agent/v4/personal_retriever_v4.yaml new file mode 100644 index 00000000..85f7e013 --- /dev/null +++ b/reme_ai/mem_agent/v4/personal_retriever_v4.yaml @@ -0,0 +1,32 @@ +tool: | + Retrieve relevant personal memories to answer user questions through vector search and history reading. + +user_message: | + You are a memory agent managing **{memory_type}** memories about **{memory_target}**. + + ## User Profile + {user_profile} + + ## Question + {context} + + ## Task + Search for relevant memories to answer the question above. + + **Tool 1: Vector Search (`retrieve_memory`)** + - Try at least 3-5 different queries: + * Direct question + * Reformulated phrasings + * Entity-focused queries + * Different keyword combinations + - If no results: retry with different time ranges [start, end] in YYYYMMDD format + * Example: [20200101, 20200102] for 20200101 <= time <= 20200102 + * Single-sided: [0, 20200102] or [20200101, 99999999] + + **Tool 2: Read Context (`read_history`) - ONLY AFTER Tool 1** + - Use history_id from retrieved memories to read original conversations + - Read multiple if needed for complete context + + **Response** + - If found relevant memories: respond exactly `` + - If no memory found after thorough search: respond exactly `` diff --git a/reme_ai/mem_agent/v4/personal_summarizer_v4.py b/reme_ai/mem_agent/v4/personal_summarizer_v4.py new file mode 100644 index 00000000..c0e1d4f1 --- /dev/null +++ b/reme_ai/mem_agent/v4/personal_summarizer_v4.py @@ -0,0 +1,124 @@ +from loguru import logger + +from ..base_memory_agent import BaseMemoryAgent +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message, MemoryNode + + +class PersonalSummarizerV4(BaseMemoryAgent): + memory_type: MemoryType = MemoryType.PERSONAL + + async def build_messages_phase1(self) -> list[Message]: + """Build messages for phase 1: AddSummaryMemory""" + history_node: MemoryNode = self.context.history_node + messages = [ + Message( + role=Role.USER, + content=self.prompt_format( + prompt_name="user_message_phase1", + context=history_node.content, + memory_type=self.memory_type.value, + memory_target=self.memory_target, + )), + ] + return messages + + async def build_messages_phase2(self, user_profile: str) -> list[Message]: + """Build messages for phase 2: UpdateUserProfile""" + history_node: MemoryNode = self.context.history_node + messages = [ + Message( + role=Role.USER, + content=self.prompt_format( + prompt_name="user_message_phase2", + context=history_node.content, + memory_type=self.memory_type.value, + memory_target=self.memory_target, + user_profile=user_profile, + )), + ] + return messages + + async def _acting_step(self, assistant_message: Message, step: int, stage: str = "", **kwargs) -> list[Message]: + return await super()._acting_step( + assistant_message, + step, + stage=stage, + memory_type=self.memory_type.value, + memory_target=self.memory_target, + history_node=self.history_node, + author=self.author, + **kwargs, + ) + + async def execute(self): + """Execute in two phases: 1) AddSummaryMemory, 2) UpdateUserProfile""" + # Log available tools + for i, tool in enumerate(self.tools): + logger.info( + f"[{self.__class__.__name__}] step0.{i} " + f"tool_call={tool.tool_call.name}", + ) + + # Phase 1: AddSummaryMemory + logger.info(f"[{self.__class__.__name__}-S1] Starting Phase 1: AddSummaryMemory") + + # Filter tools for phase 1 (only AddSummaryMemory) + original_tools = self.tools.copy() + self.tools = [t for t in self.tools if t.tool_call.name == "add_summary_memory"] + + messages_phase1 = await self.build_messages_phase1() + for i, message in enumerate(messages_phase1): + logger.info( + f"[{self.__class__.__name__}-S1] phase1.step0.{i} {message.role} " + f"{message.simple_dump(enable_json_dump=True)}", + ) + + messages_phase1, success_phase1 = await self.react(messages_phase1, stage="S1") + if not success_phase1: + logger.warning(f"[{self.__class__.__name__}-S1] Phase 1 did not complete successfully") + + # Phase 2: Read user profile and UpdateUserProfile + logger.info(f"[{self.__class__.__name__}-S2] Starting Phase 2: UpdateUserProfile") + + # Restore original tools and get ReadUserProfile tool + self.tools = original_tools + read_profile_tool = next((t for t in self.tools if t.tool_call.name == "read_user_profile"), None) + + user_profile = "" + if read_profile_tool: + # Call ReadUserProfile to load current profile (only show profile_id, not history_id) + logger.info(f"[{self.__class__.__name__}-S2] Loading user profile with ReadUserProfile") + await read_profile_tool.call( + memory_type=self.memory_type.value, + memory_target=self.memory_target, + show_ids="profile", + ) + user_profile = str(read_profile_tool.output) + logger.info(f"[{self.__class__.__name__}-S2] User profile loaded: {user_profile}...") + else: + logger.warning(f"[{self.__class__.__name__}-S2] ReadUserProfile tool not found") + + # Filter tools for phase 2 (only UpdateUserProfile) + self.tools = [t for t in self.tools if t.tool_call.name == "update_user_profile"] + + messages_phase2 = await self.build_messages_phase2(user_profile) + for i, message in enumerate(messages_phase2): + logger.info( + f"[{self.__class__.__name__}-S2] phase2.step0.{i} {message.role} " + f"{message.simple_dump(enable_json_dump=True)}", + ) + + messages_phase2, success_phase2 = await self.react(messages_phase2, stage="S2") + + # Restore original tools + self.tools = original_tools + + # Set final output and messages + self.messages = messages_phase1 + messages_phase2 + self.success = success_phase1 and success_phase2 + + if self.success and messages_phase2: + self.output = messages_phase2[-1].content + else: + self.output = "Memory processing completed with issues." diff --git a/reme_ai/mem_agent/v4/personal_summarizer_v4.yaml b/reme_ai/mem_agent/v4/personal_summarizer_v4.yaml new file mode 100644 index 00000000..3663c992 --- /dev/null +++ b/reme_ai/mem_agent/v4/personal_summarizer_v4.yaml @@ -0,0 +1,49 @@ +tool: | + Extract and update personal memories about the user from conversation context. + Identify preferences, habits, background, relationships, and key facts. + +user_message_phase1: | + You are a memory agent managing **{memory_type}** memories about **{memory_target}**. + + ## Latest Conversation: + {context} + + Message format: `round [] : ` (timestamp: YYYY-MM-DD HH:MM:SS). + + **CRITICAL**: Extract ONLY explicitly stated information. DO NOT infer, assume, or fabricate. + + ## Task: Extract Memories with `AddSummaryMemory` + + Summarize all important information about **{memory_target}** + - Set `conversation_time` (format: 2020-01-01 00:00:00; use 0000-00-00 00:00:00 if unavailable) + + Extract personal memories from the conversation using `AddSummaryMemory`. + +# capturing complete contexts with preconditions, causes, and consequences +user_message_phase2: | + You are a memory agent managing **{memory_type}** memories about **{memory_target}**. + + ## Latest Conversation: + {context} + + Message format: `round [] : ` (timestamp: YYYY-MM-DD HH:MM:SS). + + **CRITICAL**: Extract ONLY explicitly stated information. DO NOT infer, assume, or fabricate. + + ## Current User Profile: + {user_profile} + + ## Task: Update Profile with `UpdateUserProfile` + + Synchronize profile with new information from the conversation: + - `profile_ids_to_delete`: Remove conflicting, or redundant entries (array of profile IDs). + - `profiles_to_add`: + - `conversation_time`: Time of conversation (format: `YYYY-MM-DD HH:MM:SS`, e.g., `2024-01-15 14:30:00`) + - `profile_content`: Complete, self-contained profile description with full context + + **Profile Requirements**: + - One user profile entry records one dimension of the user portrait, and MUST be complete and self-contained with all necessary context (preconditions, causes, and consequences) + - All profiles MUST be mutually exclusive (non-overlapping) and non-conflicting + - Profiles should collectively be comprehensive with no information loss + + Update user profile using `UpdateUserProfile` based on the conversation and current profile. diff --git a/reme_ai/mem_agent/v4/reme_retriever_v4.py b/reme_ai/mem_agent/v4/reme_retriever_v4.py new file mode 100644 index 00000000..48ab9f38 --- /dev/null +++ b/reme_ai/mem_agent/v4/reme_retriever_v4.py @@ -0,0 +1,117 @@ +from loguru import logger + +from ..base_memory_agent import BaseMemoryAgent +from ...core.enumeration import Role +from ...core.schema import Message +from ...core.utils import format_messages + + +class ReMeRetrieverV4(BaseMemoryAgent): + + def __init__(self, meta_memories: list[dict] | None = None, **kwargs): + super().__init__(**kwargs) + self.meta_memories: list[dict] = meta_memories or [] + self.meta_info_dict: dict[str, str] = {} + + async def _read_meta_memories(self) -> str: + from ...mem_tool import ReadMetaMemory + meta_memory_info = ReadMetaMemory().format_memory_metadata(self.meta_memories) + logger.info(f"meta_memory_info={meta_memory_info}") + return meta_memory_info + + async def build_messages(self) -> list[Message]: + if self.context.get("query"): + user_query = self.context.query + elif self.context.get("messages"): + user_query = format_messages(self.context.messages) + else: + raise ValueError("Input must have either `query` or `messages`") + + messages = [ + Message( + role=Role.SYSTEM, + content=self.prompt_format( + prompt_name="system_prompt", + meta_memory_info=await self._read_meta_memories(), + user_query=user_query, + ), + ), + Message( + role=Role.USER, + content=self.get_prompt("user_message"), + ), + ] + + return messages + + async def _acting_step(self, assistant_message: Message, step: int, **kwargs) -> list[Message]: + import asyncio + from ...mem_tool.v4 import HandsOff + + if not assistant_message.tool_calls: + return [] + + tool_list: list = [] + tool_result_messages: list[Message] = [] + tool_dict = {t.tool_call.name: t for t in self.tools} + stage_prefix = "" + + # Add required context parameters + kwargs["query"] = self.context.get("query", "") + kwargs["messages"] = self.context.get("messages", []) + + for j, tool_call in enumerate(assistant_message.tool_calls): + if tool_call.name not in tool_dict: + logger.warning(f"[{self.__class__.__name__}{stage_prefix}] unknown tool_call.name={tool_call.name}") + continue + + logger.info( + f"[{self.__class__.__name__}{stage_prefix}] step{step + 1}.{j} " + f"submit tool_calls={tool_call.name} argument={tool_call.arguments}", + ) + tool_copy = tool_dict[tool_call.name].copy() + tool_copy.tool_call.id = tool_call.id + tool_list.append(tool_copy) + kwargs.update(tool_call.argument_dict) + self.submit_async_task(tool_copy.call, retrieved_nodes=self.retrieved_nodes, **kwargs) + if self.tool_call_interval > 0: + await asyncio.sleep(self.tool_call_interval) + + await self.join_async_tasks() + + for j, op in enumerate(tool_list): + if op.memory_nodes: + self.memory_nodes.extend(op.memory_nodes) + + if hasattr(op, "messages") and op.messages: + self.tool_messages.extend(op.messages) + + # Collect meta_info_dict from HandsOff tool + if isinstance(op, HandsOff) and hasattr(op, "meta_info_dict"): + self.meta_info_dict.update(op.meta_info_dict) + logger.info(f"Collected meta_info_dict from HandsOff: {len(op.meta_info_dict)} entries") + + tool_result = str(op.output) + tool_message = Message( + role=Role.TOOL, + content=tool_result, + tool_call_id=op.tool_call.id, + ) + tool_result_messages.append(tool_message) + + self.meta_info += tool_result + "\n" + + logger.info(f"[{self.__class__.__name__}{stage_prefix}] step{step + 1}.{j} join tool_result={tool_result[:2000]}...\n\n") + + return tool_result_messages + + async def execute(self): + await super().execute() + + # Assemble meta_info_dict into output + if self.meta_info_dict: + output_parts = [] + for key, value in self.meta_info_dict.items(): + output_parts.append(f"## {key}\n{value}") + self.output = "\n\n".join(output_parts) + logger.info(f"Assembled output from meta_info_dict with {len(self.meta_info_dict)} entries") \ No newline at end of file diff --git a/reme_ai/mem_agent/v4/reme_retriever_v4.yaml b/reme_ai/mem_agent/v4/reme_retriever_v4.yaml new file mode 100644 index 00000000..7360febd --- /dev/null +++ b/reme_ai/mem_agent/v4/reme_retriever_v4.yaml @@ -0,0 +1,25 @@ +tool: | + Retrieve information from specialized memory agents to answer user queries. + +system_prompt: | + You are a Memory Retrieval Orchestrator responsible for querying specialized agents to answer user questions. + + # User Query + {user_query} + + ## Available Memory Agents + Each line indicates a specialized Memory Agent that stores and retrieves memories within a specific dimension (memory_type + memory_target). + Format: "- (): " + {meta_memory_info} + + ## Your Task + 1. Use the `hands_off` tool to retrieve information from relevant agents + - Specify `memory_type` and `memory_target` for each query + - The `memory_type` and `memory_target` must **exactly match** existing entries in the "Available Memory Agents" listed above + - Do NOT query agents that don't exist above + - You can query multiple agents if needed + 2. Answer the user query STRICTLY based on the `hands_off` results + 3. If the retrieved information is insufficient to answer the query, respond: "nothing found after thorough search." + +user_message: | + Please retrieve relevant information from the existing agents and provide an answer based on the results. diff --git a/reme_ai/mem_agent/v4/reme_summarizer_v4.py b/reme_ai/mem_agent/v4/reme_summarizer_v4.py new file mode 100644 index 00000000..a4069c85 --- /dev/null +++ b/reme_ai/mem_agent/v4/reme_summarizer_v4.py @@ -0,0 +1,63 @@ +from loguru import logger + +from ..base_memory_agent import BaseMemoryAgent +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message, MemoryNode +from ...core.utils import format_messages + + +class ReMeSummarizerV4(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: + from ...mem_tool import ReadMetaMemory + meta_memory_info = ReadMetaMemory().format_memory_metadata(self.meta_memories) + logger.info(f"meta_memory_info={meta_memory_info}") + return meta_memory_info + + async def build_messages(self) -> list[Message]: + self.context.messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages] + history_content = self.description + "\n" + format_messages(self.context.messages) + self.context.history_node = history_node = MemoryNode( + memory_type=MemoryType.HISTORY, + memory_target="", + when_to_use=history_content[:100], + content=history_content, + ref_memory_id="", + author=self.author, + metadata={}, + ) + + logger.info(f"Adding summary node: {history_node.model_dump_json(indent=2, exclude_none=True)}") + await self.vector_store.delete(history_node.memory_id) + await self.vector_store.insert([history_node.to_vector_node()]) + + messages = [ + Message( + role=Role.SYSTEM, + content=self.prompt_format( + prompt_name="system_prompt", + meta_memory_info=await self._read_meta_memories(), + context=history_node.content, + ), + ), + Message( + role=Role.USER, + content=self.get_prompt("user_message"), + ), + ] + + return messages + + async def _acting_step(self, assistant_message: Message, step: int, **kwargs) -> list[Message]: + return await super()._acting_step( + assistant_message, + step, + messages=self.context.messages, + history_node=self.context.history_node, + author=self.author, + **kwargs, + ) diff --git a/reme_ai/mem_agent/v4/reme_summarizer_v4.yaml b/reme_ai/mem_agent/v4/reme_summarizer_v4.yaml new file mode 100644 index 00000000..4cb6f54f --- /dev/null +++ b/reme_ai/mem_agent/v4/reme_summarizer_v4.yaml @@ -0,0 +1,26 @@ +tool: | + Orchestrate memory updates across specialized memory agents. + +system_prompt: | + You are a Memory Orchestrator responsible for routing memory tasks to specialized agents based on the context. + + # Context + {context} + + ## Available Memory Agents + Each line indicates 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 the `hands_off` tool to distribute memory tasks to specialized agents: + 1. Analyze the context and identify which memory dimensions require updates + 2. Specify `memory_type` and `memory_target` for each task + - The `memory_type` and `memory_target` must **exactly match** existing entries in the "Available Memory Agents" listed above + - Do NOT create new agents or use memory_type/memory_target combinations that don't exist above + 3. Multiple tasks can be specified to enable parallel processing by specialized agents + + Note: If the context contains no memorable information (e.g., simple greetings), return ``. + +user_message: | + Please analyze the context and route memory tasks to the appropriate existing agents. diff --git a/reme_ai/mem_agent/wk/reme_retriever_wk.yaml b/reme_ai/mem_agent/wk/reme_retriever_wk.yaml index 281796d6..755fd6b7 100644 --- a/reme_ai/mem_agent/wk/reme_retriever_wk.yaml +++ b/reme_ai/mem_agent/wk/reme_retriever_wk.yaml @@ -24,16 +24,16 @@ system_prompt: | 1. **Multi-Angle Vector Retrieval** (REQUIRED - at least 3 attempts): You must try AT LEAST 3 different retrieval approaches using `retrieve_memories`: - + a) **Direct Vector Search**: Use the user's question directly or with minimal reformulation - Query the most relevant memory_type and memory_target - Use straightforward query phrasing - + b) **Alternative Phrasing**: Reformulate the query from a different angle - Use synonyms or different expressions - Break down complex questions into simpler components - Try more specific or more general queries - + c) **Metadata-Filtered Search**: Add metadata filters to narrow down results - **Time-based filtering**: Use year/month/day metadata fields to filter by time periods * Example: {{"year": 2024}} for memories from 2024 @@ -41,33 +41,33 @@ system_prompt: | * Example: {{"year": 2024, "month": 5, "day": 15}} for memories from a specific date - Combine vector search with metadata constraints - Try partial metadata filtering if full filtering yields nothing (e.g., only year, or year+month) - + d) **Cross-Memory-Type Search**: If applicable, search across different memory types - Try different memory_type and memory_target combinations - Some information might be stored in unexpected memory categories - + e) **Keyword Extraction**: Extract key entities/concepts and search for them - Identify important names, places, concepts - Search for each key element separately - + 2. **Evaluate Retrieval Results** (After each attempt): - Review what memories were returned - Assess if they contain sufficient information to answer the question - If insufficient, identify what's missing and adjust your next query accordingly - Track which retrieval strategies you've already tried - + 3. **Persist Through Failures**: - DO NOT give up after 1-2 failed attempts - If a retrieval returns no results or irrelevant results, try a different approach - Consider that the information might be phrased differently than expected - Be creative with query reformulation - + 4. **Fallback to History Reading** (Only after 3+ vector retrieval attempts): - If after at least 3 different vector retrieval attempts you still lack sufficient information: * If any retrieved memories contain `ref_memory_id`, use `read_history` to read the original conversation * Use `read_history` with the `ref_memory_id` to get complete context * This can reveal details that weren't captured in the memory summaries - + 5. **Answer the Question**: - Once you have sufficient information, provide a direct answer based ONLY on retrieved memories - DO NOT fabricate, guess, or infer information not present in the memories @@ -95,30 +95,30 @@ system_prompt: | **Example 1: Simple Query** Attempt 1: Direct query "user's favorite food" → Result: No relevant memories found - + Attempt 2: Reformulated query "what does user like to eat" → Result: Some memories about meals, but not specific preferences - + Attempt 3: Keyword search "food preferences" with metadata filter → Result: Found relevant memory with ref_memory_id - + Attempt 4: Use read_history with ref_memory_id to get full context → Result: Found detailed conversation about favorite foods - + Answer: [Provide answer based on retrieved information] **Example 2: Time-based Query** Question: "What did the user do last summer?" - + Attempt 1: Direct query "user activities summer" with metadata {{"year": 2025, "month": [6, 7, 8]}} → Result: Found some vacation memories - + Attempt 2: Broader query "user summer vacation travel" with metadata {{"year": 2025}} → Result: Found additional travel-related memories - + Attempt 3: Use read_history for memories with ref_memory_id to get detailed context → Result: Complete picture of summer activities - + Answer: [Provide answer based on retrieved information] user_message: | diff --git a/reme_ai/mem_agent/wk/reme_summarizer_wk.yaml b/reme_ai/mem_agent/wk/reme_summarizer_wk.yaml index 30792a08..58080cab 100644 --- a/reme_ai/mem_agent/wk/reme_summarizer_wk.yaml +++ b/reme_ai/mem_agent/wk/reme_summarizer_wk.yaml @@ -18,7 +18,7 @@ system_prompt: | 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: | diff --git a/reme_ai/mem_tool/base_memory_tool.py b/reme_ai/mem_tool/base_memory_tool.py deleted file mode 100644 index 6accfe2a..00000000 --- a/reme_ai/mem_tool/base_memory_tool.py +++ /dev/null @@ -1,124 +0,0 @@ -"""Base class for memory tool""" - -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 -from ..core.utils import CacheHandler - - -class BaseMemoryTool(BaseOp, metaclass=ABCMeta): - """Base class for memory tool""" - - def __init__( - self, - enable_multiple: bool = True, - enable_thinking_params: bool = False, - meta_memory_path: str = "./meta_memory", - **kwargs, - ): - super().__init__(**kwargs) - self.enable_multiple: bool = enable_multiple - self.enable_thinking_params: bool = enable_thinking_params - self.meta_memory_path: str = meta_memory_path - self.memory_nodes: list[MemoryNode | str] = [] - - def _build_parameters(self) -> dict: - return {} - - def _build_multiple_parameters(self) -> dict: - return {} - - def _build_tool_description(self) -> str: - """Build tool description.""" - return self.get_prompt("tool" + ("_multiple" if self.enable_multiple else "")) - - def _build_tool_call(self) -> ToolCall: - tool_call_params: dict = { - "description": self._build_tool_description(), - } - - if self.enable_multiple: - parameters = self._build_multiple_parameters() - else: - parameters = self._build_parameters() - - if parameters: - tool_call_params["parameters"] = parameters - - if self.enable_thinking_params and "thinking" not in parameters["properties"]: - parameters["properties"] = { - "thinking": { - "type": "string", - "description": "Your complete and detailed thinking process about how to fill in each parameter", - }, - **parameters["properties"], - } - parameters["required"] = ["thinking", *parameters["required"]] - - return ToolCall(**tool_call_params) - - @property - def meta_memory(self) -> CacheHandler: - """Create the meta memory cache handler.""" - return CacheHandler(Path(self.meta_memory_path) / self.vector_store.collection_name) - - @property - def memory_type(self) -> MemoryType: - """Get the memory type from context.""" - return MemoryType(self.context.get("memory_type")) - - @property - def memory_target(self) -> str: - """Get the memory target from context.""" - return self.context.get("memory_target", "") - - @property - def ref_memory_id(self) -> str: - """Get the reference memory ID from context.""" - return self.context.get("ref_memory_id", "") - - @property - def messages_formated(self) -> str: - """Get the formated messages from context.""" - return self.context.get("messages_formated", "") - - @property - def retrieved_nodes(self) -> list[MemoryNode]: - """Get the retrieved nodes from context.""" - return self.context.get("retrieved_nodes") - - @property - def author(self) -> str: - """Get the author from context.""" - return self.context.get("author", "") - - def _build_memory_node( - self, - memory_content: str, - memory_type: MemoryType | None = None, - memory_target: str = "", - ref_memory_id: str = "", - when_to_use: str = "", - author: str = "", - metadata: dict | None = None, - ) -> MemoryNode: - """Build MemoryNode from content, when_to_use, and metadata.""" - node = MemoryNode( - memory_type=memory_type or self.memory_type, - memory_target=memory_target or self.memory_target, - when_to_use=when_to_use or "", - content=memory_content, - ref_memory_id=ref_memory_id or self.ref_memory_id, - author=author or self.author, - metadata=metadata or {}, - ) - - # logger.opt(depth=1).info( - # f"[{self.__class__.__name__}] build node={node.model_dump_json(indent=2, exclude_none=True)}", - # ) - return node diff --git a/reme_ai/mem_tool/history/add_history_memory.py b/reme_ai/mem_tool/history/add_history_memory.py deleted file mode 100644 index a92deca5..00000000 --- a/reme_ai/mem_tool/history/add_history_memory.py +++ /dev/null @@ -1,50 +0,0 @@ -"""Add history memory operation.""" - -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ...core.context import C -from ...core.enumeration import MemoryType -from ...core.schema import ToolCall, Message -from ...core.utils import format_messages - - -@C.register_op() -class AddHistoryMemory(BaseMemoryTool): - """Add history memory from conversation messages.""" - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": self.get_prompt("tool"), - "parameters": { - "type": "object", - "properties": { - "messages": { - "type": "array", - "description": self.get_prompt("messages"), - "items": {"type": "object"}, - }, - }, - "required": ["messages"], - }, - }, - ) - - async def execute(self): - messages: list[Message | dict] = self.context.get("messages", []) - if not messages: - self.output = "No messages provided for addition." - return - - messages = [Message(**m) if isinstance(m, dict) else m for m in messages] - memory_content = format_messages(messages) - memory_node = self._build_memory_node(memory_content=memory_content, memory_type=MemoryType.HISTORY) - vector_node = memory_node.to_vector_node() - - await self.vector_store.delete(vector_ids=[vector_node.vector_id]) - await self.vector_store.insert(nodes=[vector_node]) - self.memory_nodes.append(memory_node) - - self.output = "Successfully added history memory to vector_store." - logger.info(self.output) diff --git a/reme_ai/mem_tool/history/add_history_memory.yaml b/reme_ai/mem_tool/history/add_history_memory.yaml deleted file mode 100644 index 18fd9163..00000000 --- a/reme_ai/mem_tool/history/add_history_memory.yaml +++ /dev/null @@ -1,14 +0,0 @@ -tool: | - Add history memory from conversation messages. - -tool_multiple: | - Add multiple history memories in a single operation. - -messages: | - List of message objects with 'role' and 'content' fields. - -metadata: | - Optional metadata (time, session_id, topic, etc.). - -histories: | - List of history objects, each with messages and optional metadata. diff --git a/reme_ai/mem_tool/history/read_history_memory.py b/reme_ai/mem_tool/history/read_history_memory.py deleted file mode 100644 index def2ff24..00000000 --- a/reme_ai/mem_tool/history/read_history_memory.py +++ /dev/null @@ -1,65 +0,0 @@ -"""Read history memory operation.""" - -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ...core.context import C -from ...core.schema import MemoryNode - - -@C.register_op() -class ReadHistoryMemory(BaseMemoryTool): - """Read history memories by IDs.""" - - def _build_parameters(self) -> dict: - return { - "type": "object", - "properties": { - "ref_memory_id": { - "type": "string", - "description": self.get_prompt("ref_memory_id"), - }, - }, - "required": ["ref_memory_id"], - } - - def _build_multiple_parameters(self) -> dict: - return { - "type": "object", - "properties": { - "ref_memory_ids": { - "type": "array", - "description": self.get_prompt("ref_memory_ids"), - "items": {"type": "string"}, - }, - }, - "required": ["ref_memory_ids"], - } - - async def execute(self): - if self.enable_multiple: - ref_memory_ids: list[str] = self.context.get("ref_memory_ids", []) - else: - ref_memory_id = self.context.get("ref_memory_id", "") - ref_memory_ids: list[str] = [ref_memory_id] if ref_memory_id else [] - - # Remove empty IDs and duplicates - ref_memory_ids = [mid for mid in ref_memory_ids if mid] - ref_memory_ids = list(dict.fromkeys(ref_memory_ids)) # Remove duplicates while preserving order - - if not ref_memory_ids: - self.output = "No valid reference memory IDs provided for reading." - logger.warning(self.output) - return - - # Query original history dialogues by ref_memory_id - nodes = await self.vector_store.get(vector_ids=ref_memory_ids) - - if not nodes: - self.output = "No history memories found with the provided reference IDs." - logger.warning(self.output) - return - - memories: list[MemoryNode] = [MemoryNode.from_vector_node(n) for n in nodes] - self.output = "---\n".join([m.content for m in memories]) - logger.info(f"Successfully read {len(memories)} history memories by reference IDs.") diff --git a/reme_ai/mem_tool/history/read_history_memory.yaml b/reme_ai/mem_tool/history/read_history_memory.yaml deleted file mode 100644 index 24d99b18..00000000 --- a/reme_ai/mem_tool/history/read_history_memory.yaml +++ /dev/null @@ -1,11 +0,0 @@ -tool: | - Read original history dialogue by reference memory ID. - -tool_multiple: | - Read multiple original history dialogues by reference memory IDs. - -ref_memory_id: | - Reference memory ID to query the original history dialogue. - -ref_memory_ids: | - List of reference memory IDs to query the original history dialogues. Please provide unique IDs without duplicates. diff --git a/reme_ai/mem_tool/identity/read_identity_memory.py b/reme_ai/mem_tool/identity/read_identity_memory.py deleted file mode 100644 index bd9f8031..00000000 --- a/reme_ai/mem_tool/identity/read_identity_memory.py +++ /dev/null @@ -1,27 +0,0 @@ -"""Read identity memory operation.""" - -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ...core.context import C - - -@C.register_op() -class ReadIdentityMemory(BaseMemoryTool): - """Read identity memory for agent self-cognition.""" - - def __init__(self, **kwargs): - kwargs["enable_multiple"] = False - super().__init__(**kwargs) - - def _build_parameters(self) -> dict: - return { - "type": "object", - "properties": {}, - "required": [], - } - - async def execute(self): - identity_memory = self.meta_memory.load("identity_memory") or "" - self.output = identity_memory or "No identity memory found." - logger.info(self.output) diff --git a/reme_ai/mem_tool/identity/read_identity_memory.yaml b/reme_ai/mem_tool/identity/read_identity_memory.yaml deleted file mode 100644 index 6e84e4bf..00000000 --- a/reme_ai/mem_tool/identity/read_identity_memory.yaml +++ /dev/null @@ -1,3 +0,0 @@ -tool: | - Read the identity memory for the agent. - Retrieve self-cognition information such as identity, role, personality, or current state. diff --git a/reme_ai/mem_tool/identity/update_identity_memory.py b/reme_ai/mem_tool/identity/update_identity_memory.py deleted file mode 100644 index b0211242..00000000 --- a/reme_ai/mem_tool/identity/update_identity_memory.py +++ /dev/null @@ -1,39 +0,0 @@ -"""Update identity memory operation.""" - -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ...core.context import C - - -@C.register_op() -class UpdateIdentityMemory(BaseMemoryTool): - """Update identity memory for agent self-cognition.""" - - def __init__(self, **kwargs): - kwargs["enable_multiple"] = False - super().__init__(**kwargs) - - def _build_parameters(self) -> dict: - return { - "type": "object", - "properties": { - "identity_memory": { - "type": "string", - "description": self.get_prompt("identity_memory"), - }, - }, - "required": ["identity_memory"], - } - - async def execute(self): - identity_memory = self.context.get("identity_memory", "") - - if not identity_memory: - self.output = "No valid identity memory provided for update." - logger.warning(self.output) - return - - self.meta_memory.save("identity_memory", identity_memory) - self.output = "Successfully updated identity memory." - logger.info(self.output) diff --git a/reme_ai/mem_tool/identity/update_identity_memory.yaml b/reme_ai/mem_tool/identity/update_identity_memory.yaml deleted file mode 100644 index 48033bfe..00000000 --- a/reme_ai/mem_tool/identity/update_identity_memory.yaml +++ /dev/null @@ -1,7 +0,0 @@ -tool: | - Update the identity memory for the agent. - Store self-cognition information such as identity, role, personality, or current state. - -identity_memory: | - The identity memory content to store. - Should be a clear statement capturing the agent's self-cognition or current state. diff --git a/reme_ai/mem_tool/meta/add_meta_memory.py b/reme_ai/mem_tool/meta/add_meta_memory.py deleted file mode 100644 index d7b41254..00000000 --- a/reme_ai/mem_tool/meta/add_meta_memory.py +++ /dev/null @@ -1,121 +0,0 @@ -"""Add meta memory operation for adding memory metadata.""" - -import json - -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ...core.context import C -from ...core.enumeration import MemoryType - - -@C.register_op() -class AddMetaMemory(BaseMemoryTool): - """Add memory metadata (memory_type and memory_target) to meta storage. - - Supports single/multiple addition modes via `enable_multiple` parameter. - """ - - def _build_item_schema(self) -> tuple[dict, list[str]]: - """Build shared schema properties and required fields for meta memory items. - - Returns: - Tuple of (properties dict, required fields list). - """ - properties = { - "memory_type": { - "type": "string", - "description": self.get_prompt("memory_type"), - "enum": [MemoryType.PERSONAL.value, MemoryType.PROCEDURAL.value], - }, - "memory_target": { - "type": "string", - "description": self.get_prompt("memory_target"), - }, - } - required = ["memory_type", "memory_target"] - return properties, required - - def _build_parameters(self) -> dict: - """Build input schema for single meta memory addition.""" - properties, required = self._build_item_schema() - return { - "type": "object", - "properties": properties, - "required": required, - } - - def _build_multiple_parameters(self) -> dict: - """Build input schema for multiple meta memory addition.""" - item_properties, required_fields = self._build_item_schema() - return { - "type": "object", - "properties": { - "meta_memories": { - "type": "array", - "description": self.get_prompt("meta_memories"), - "items": { - "type": "object", - "properties": item_properties, - "required": required_fields, - }, - }, - }, - "required": ["meta_memories"], - } - - def _load_meta_memories(self) -> list[dict]: - """Load existing meta memories from cache.""" - return self.meta_memory.load("meta_memories") or [] - - def _save_meta_memories(self, memories: list[dict]) -> bool: - """Save meta memories to cache.""" - return self.meta_memory.save("meta_memories", memories) - - @staticmethod - def _filter_memory_type_target(memory_type: str, memory_target: str, existing_set: set) -> bool: - result = ( - memory_type in [MemoryType.PERSONAL.value, MemoryType.PROCEDURAL.value] - and memory_target - and (memory_type, memory_target) not in existing_set - ) - if result: - existing_set.add((memory_type, memory_target)) - return result - - async def execute(self): - """Execute addition: load existing, merge with new, and save. - - Duplicates (same memory_type and memory_target) are skipped. - """ - existing_memories: list[dict] = self._load_meta_memories() - existing_set = {(m["memory_type"], m["memory_target"]) for m in existing_memories} - - # Build new memories to add based on mode - new_memories: list[dict] = [] - if self.enable_multiple: - meta_memories: list[dict] = self.context.get("meta_memories", []) - for mem in meta_memories: - memory_type = mem.get("memory_type", "") - memory_target = mem.get("memory_target", "") - if self._filter_memory_type_target(memory_type, memory_target, existing_set): - new_memories.append({"memory_type": memory_type, "memory_target": memory_target}) - - else: - memory_type = self.context.get("memory_type", "") - memory_target = self.context.get("memory_target", "") - if self._filter_memory_type_target(memory_type, memory_target, existing_set): - new_memories.append({"memory_type": memory_type, "memory_target": memory_target}) - - if not new_memories: - self.output = "No new meta memories to add (all entries already exist or invalid)." - return - - # Merge and save - all_memories = existing_memories + new_memories - self._save_meta_memories(all_memories) - - # Format output - added_str = json.dumps(new_memories, ensure_ascii=False) - self.output = f"Successfully added {len(new_memories)} meta memory entries: {added_str}" - logger.info(self.output) diff --git a/reme_ai/mem_tool/meta/add_meta_memory.yaml b/reme_ai/mem_tool/meta/add_meta_memory.yaml deleted file mode 100644 index d54c3345..00000000 --- a/reme_ai/mem_tool/meta/add_meta_memory.yaml +++ /dev/null @@ -1,26 +0,0 @@ -tool: | - Add a memory metadata entry to register a new memory type and target. - IMPORTANT: Before using this tool, verify that the Main Agent's Meta Memory does NOT already contain the same () combination. Only create new entries if they don't exist. - Use this tool to define what types of memories should be tracked, such as: - - Personal memories: "John", "Alice" (person-specific preferences and context) - - Procedural memories: "deployment_process", "code_review_steps" (how-to knowledge) - -tool_multiple: | - Add multiple memory metadata entries to register multiple memory types and targets at once. - Before using this tool, verify that the Main Agent's Meta Memory does NOT already contain the same () combinations. Only create new entries for those that don't exist. - Use this tool to define multiple memory tracking categories in a single operation. - Each entry specifies a memory_type and memory_target for organizing different memory domains. - -meta_memories: | - A list of memory metadata entries to add. Each entry contains memory_type and memory_target. - -memory_type: | - The type of memory to register. Valid values are: personal, procedural. - - personal: Person-specific memory storing preferences and context about specific individuals - - procedural: Procedural memory storing how-to knowledge and step-by-step processes - -memory_target: | - The target identifier for this memory category. - Examples: - - For personal memory: person's name (e.g., "John", "Alice") - - For procedural memory: domain or topic name (e.g., "deployment", "code_review") diff --git a/reme_ai/mem_tool/meta/read_meta_memory.py b/reme_ai/mem_tool/meta/read_meta_memory.py deleted file mode 100644 index 07ad1ecf..00000000 --- a/reme_ai/mem_tool/meta/read_meta_memory.py +++ /dev/null @@ -1,97 +0,0 @@ -"""Read meta memory operation for retrieving memory metadata.""" - -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ...core.context import C -from ...core.enumeration import MemoryType - - -@C.register_op() -class ReadMetaMemory(BaseMemoryTool): - """Read memory metadata (memory_type and memory_target) from meta storage. - - This operation reads stored memory metadata and optionally includes - TOOL and IDENTITY type memories. - """ - - def __init__( - self, - enable_identity_memory: bool = False, - **kwargs, - ): - """Initialize ReadMetaMemory. - - Args: - enable_identity_memory: Include IDENTITY type meta memory. Defaults to False. - **kwargs: Additional arguments for BaseMemoryTool. - """ - kwargs["enable_multiple"] = False - super().__init__(**kwargs) - self.enable_identity_memory = enable_identity_memory - - def _build_parameters(self) -> dict: - """Build input schema for reading meta memory. - - No input parameters required for reading. - """ - return { - "type": "object", - "properties": {}, - "required": [], - } - - def _load_meta_memories(self) -> list[dict[str, str]]: - """Load meta memories from cache and apply filters.""" - result = self.meta_memory.load("meta_memories") - all_memories = result if result is not None else [] - - filtered_memories = [] - for m in all_memories: - if m.get("memory_type") in [MemoryType.PERSONAL.value, MemoryType.PROCEDURAL.value]: - filtered_memories.append(m) - - if self.enable_identity_memory: - filtered_memories.append( - { - "memory_type": MemoryType.IDENTITY.value, - "memory_target": "self", - }, - ) - - return filtered_memories - - def format_memory_metadata(self, memories: list[dict[str, str]]) -> str: - """Format memory metadata into a readable string. - - Args: - memories: List of memory metadata entries. - - Returns: - str: Formatted memory metadata string. - """ - if not memories: - return "" - - lines = [] - for memory in memories: - memory_type = memory["memory_type"] - memory_target = memory["memory_target"] - description = self.get_prompt(f"type_{memory_type}") - lines.append(f"- {memory_type}({memory_target}): {description}") - - return "\n".join(lines) - - async def execute(self): - """Execute the read meta memory operation. - - Reads memory metadata from cache storage and formats output. - """ - memories = self._load_meta_memories() - - if memories: - self.output = self.format_memory_metadata(memories) - logger.info(f"Retrieved {len(memories)} meta memory entries") - else: - self.output = "No memory metadata found." - logger.info(self.output) diff --git a/reme_ai/mem_tool/meta/read_meta_memory.yaml b/reme_ai/mem_tool/meta/read_meta_memory.yaml deleted file mode 100644 index dff25278..00000000 --- a/reme_ai/mem_tool/meta/read_meta_memory.yaml +++ /dev/null @@ -1,16 +0,0 @@ -tool: | - Read the memory metadata registry to see what types of memories are being tracked. - Use this tool to retrieve all registered memory types and their targets. - This helps understand what memory categories are available for storing and retrieving information. - -type_identity: | - Self-cognition memory storing agent's identity, personality, and current state. - -type_personal: | - Person-specific memory storing preferences and context about specific individuals. - -type_procedural: | - Procedural memory storing how-to knowledge and step-by-step processes. - -type_tool: | - Tool memory storing tool usage patterns, success rates, token consumption, and latency. 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..f239d5d1 --- /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/v2/read_history.py b/reme_ai/mem_tool/v2/read_history.py index 15aaada3..7bb3b830 100644 --- a/reme_ai/mem_tool/v2/read_history.py +++ b/reme_ai/mem_tool/v2/read_history.py @@ -10,13 +10,13 @@ from ...core.schema import MemoryNode @C.register_op() class ReadHistory(BaseMemoryTool): """Read original history dialogue by reference memory ID. - + Only supports single memory read (enable_multiple=False). """ def __init__(self, **kwargs): """Initialize ReadHistory. - + Args: **kwargs: Additional args for BaseMemoryTool. """ diff --git a/reme_ai/mem_tool/v2/retrieve_memories.yaml b/reme_ai/mem_tool/v2/retrieve_memories.yaml index f83e11cb..0e65d246 100644 --- a/reme_ai/mem_tool/v2/retrieve_memories.yaml +++ b/reme_ai/mem_tool/v2/retrieve_memories.yaml @@ -9,7 +9,7 @@ tool_multiple: | This prevents redundant information in subsequent retrievals. memory_type: | - The type of memory to search for. + The type of memory to search for. You MUST select one of the memory_type values that are explicitly provided in the Available Meta-Memories. memory_target: | diff --git a/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.yaml b/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.yaml index ed91357e..d5791a5d 100644 --- a/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.yaml +++ b/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.yaml @@ -1,16 +1,16 @@ tool_multiple: | Retrieve memories using both time-based and multiple vector similarity searches. - + This tool combines two retrieval strategies: 1. First retrieves the most recent memories based on modification time (recent top {recent_top_k}) 2. Then retrieves semantically similar memories for each of your queries (similar top {similar_top_k} per query) - + This is useful when you need to search for different types of information in a single operation, while also considering recent context. - + The results are automatically deduplicated, so you get a combined set of both recent and relevant memories without duplicates. - + Note: Within the same session, this tool automatically deduplicates results across multiple calls. If you call this tool multiple times, only new memories (not previously retrieved) will be returned. This prevents redundant information in subsequent retrievals. 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..ee488639 --- /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": "string", + "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..3dba2bcf --- /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 n: n.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..46879e2b --- /dev/null +++ b/reme_ai/mem_tool/v3/update_user_profile.py @@ -0,0 +1,99 @@ +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.schema.memory_node import MemoryNode + + +class UpdateUserProfile(BaseMemoryTool): + + def __init__(self, **kwargs): + kwargs["enable_multiple"] = True + super().__init__(**kwargs) + + def _build_tool_description(self) -> str: + return "Update user profile." + + def _build_multiple_parameters(self) -> dict: + return { + "type": "object", + "properties": { + "profile_ids_to_delete": { + "type": "array", + "description": "profile_ids_to_delete", + "items": {"type": "string"}, + }, + "profiles_to_add": { + "type": "array", + "description": "profiles_to_add", + "items": { + "type": "object", + "properties": { + "profile_content": { + "type": "string", + "description": "profile_content", + }, + "conversation_time": { + "type": "string", + "description": "conversation_time, e.g. '2020-01-01 00:00:00'", + }, + }, + "required": ["profile_content", "conversation_time"], + }, + }, + }, + "required": ["profile_ids_to_delete", "profiles_to_add"], + } + + async def execute(self): + 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 + + cache_key = f"{self.memory_type}_{self.memory_target}" + 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", "") + conversation_time = mem.get("conversation_time", "") + new_memory_nodes.append(self._build_memory_node( + memory_content=profile_content, + when_to_use="", + metadata={"conversation_time": conversation_time} + )) + logger.info(f"Added {len(new_memory_nodes)} 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/v4/__init__.py b/reme_ai/mem_tool/v4/__init__.py new file mode 100644 index 00000000..58b04aaf --- /dev/null +++ b/reme_ai/mem_tool/v4/__init__.py @@ -0,0 +1,15 @@ +from .add_summary_memory import AddSummaryMemory +from .hands_off import HandsOff +from .read_history import ReadHistory +from .read_user_profile import ReadUserProfile +from .retrieve_memory import RetrieveMemory +from .update_user_profile import UpdateUserProfile + +__all__ = [ + "AddSummaryMemory", + "HandsOff", + "ReadHistory", + "ReadUserProfile", + "RetrieveMemory", + "UpdateUserProfile", +] diff --git a/reme_ai/mem_tool/v4/add_summary_memory.py b/reme_ai/mem_tool/v4/add_summary_memory.py new file mode 100644 index 00000000..cc4be602 --- /dev/null +++ b/reme_ai/mem_tool/v4/add_summary_memory.py @@ -0,0 +1,63 @@ +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.schema import MemoryNode + + +class AddSummaryMemory(BaseMemoryTool): + + def __init__(self, **kwargs): + kwargs["enable_multiple"] = False + super().__init__(**kwargs) + + def _build_tool_description(self) -> str: + return "Add a summary memory to the vector store for future retrieval." + + @staticmethod + def _build_item_schema() -> tuple[dict, list[str]]: + properties = { + "conversation_time": {"type": "string", "description": "conversation time, e.g. '2020-01-01 00:00:00'"}, + "summary_memory": {"type": "string", "description": "summary_memory"}, + } + return properties, ["conversation_time", "summary_memory"] + + def _build_parameters(self) -> dict: + properties, required = self._build_item_schema() + return { + "type": "object", + "properties": properties, + "required": required, + } + + async def execute(self): + summary_memory = self.context.get("summary_memory", "") + conversation_time = self.context.get("conversation_time", "") + + if not summary_memory: + self.output = "No summary_memory provided for addition." + return + + metadata: dict = {"conversation_time": conversation_time} + try: + metadata["time_int"] = int(conversation_time.split(" ")[0].replace("-", "")) + except Exception: + pass + + memory_node = MemoryNode( + memory_type=self.memory_type, + memory_target=self.memory_target, + when_to_use="", + content=summary_memory, + ref_memory_id=self.history_node.memory_id, + author=self.author, + metadata=metadata, + ) + + vector_node = memory_node.to_vector_node() + vector_id = vector_node.vector_id + await self.vector_store.delete(vector_ids=[vector_id]) + await self.vector_store.insert(nodes=[vector_node]) + self.memory_nodes.append(memory_node) + + self.output = f"Successfully added summary memory to vector_store." + logger.info(self.output) diff --git a/reme_ai/mem_tool/v4/hands_off.py b/reme_ai/mem_tool/v4/hands_off.py new file mode 100644 index 00000000..fd818b4f --- /dev/null +++ b/reme_ai/mem_tool/v4/hands_off.py @@ -0,0 +1,115 @@ +from typing import TYPE_CHECKING + +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.enumeration import MemoryType +from ...core.schema import Message + +if TYPE_CHECKING: + from ...mem_agent import BaseMemoryAgent + + +class HandsOff(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] = [] + self.meta_info_dict: dict[str, str] = {} + + @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 "Distribute memory tasks to appropriate memory agents." + + def _build_multiple_parameters(self) -> dict: + return { + "type": "object", + "properties": { + "memory_tasks": { + "type": "array", + "description": "List of memory tasks to distribute to specific agents", + "items": { + "type": "object", + "properties": { + "memory_type": { + "type": "string", + "description": "Type of memory to handle", + "enum": [k.value for k in self.memory_agent_dict], + }, + "memory_target": { + "type": "string", + "description": "Target or context for the memory operation", + }, + }, + "required": ["memory_type", "memory_target"], + }, + }, + }, + "required": ["memory_tasks"], + } + + async def execute(self): + tasks = [] + seen = set() + for task in self.context.get("memory_tasks", []): + memory_type = MemoryType(task.get("memory_type", "")) + memory_target = task.get("memory_target", "") + + # Deduplicate tasks with same memory_type and memory_target + task_key = (memory_type, memory_target) + if task_key in seen: + logger.info(f"Skipping duplicate task: memory_type={memory_type.value}, memory_target={memory_target}") + continue + seen.add(task_key) + + tasks.append({ + "memory_type": memory_type, + "memory_target": memory_target, + }) + + 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"] + + 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, + memory_type=memory_type, + memory_target=memory_target, + query=self.context.get("query", ""), + messages=self.context.get("messages", []), + description=self.context.get("description"), + history_node=self.context.get("history_node"), + ) + + await self.join_async_tasks() + + results = [] + for i, (agent, memory_type, memory_target) in enumerate(agent_list): + if agent.memory_nodes: + self.memory_nodes.extend(agent.memory_nodes) + if agent.messages: + self.messages.extend(agent.messages) + if agent.meta_info: + self.meta_info_dict[f"{memory_type.value} {memory_target}"] = agent.meta_info + + results.append(f"{memory_type.value} {memory_target} agent result: {agent.output}") + + self.output = "\n".join(results) + logger.info(f"Completed {len(results)} hands-off task(s):\n{self.output}") diff --git a/reme_ai/mem_tool/v4/retrieve_memory.py b/reme_ai/mem_tool/v4/retrieve_memory.py new file mode 100644 index 00000000..7a902c71 --- /dev/null +++ b/reme_ai/mem_tool/v4/retrieve_memory.py @@ -0,0 +1,103 @@ +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": { + "query": { + "type": "string", + "description": "query", + }, + "time_range": { + "type": "string", + "description": "time_range(optional), e.g. [20200101, 20200101]", + }, + }, + "required": ["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: + query = query_item.get("query") + time_range = query_item.get("time_range", "") + + filter_dict: dict = { + "memory_type": self.memory_type.value, + "memory_target": self.memory_target, + } + + if time_range: + # Handle different time_range formats + if isinstance(time_range, str): + try: + time_range = json.loads(time_range) + except json.JSONDecodeError: + # If it's a plain string like "20250907", treat it as a single date + time_range = time_range + + # Convert to list format [start, end] + if isinstance(time_range, (list, tuple)): + if len(time_range) == 1: + # Single element list, use it for both start and end + filter_dict["time_int"] = [int(time_range[0]), int(time_range[0])] + else: + # Two element list/tuple + filter_dict["time_int"] = [int(time_range[0]), int(time_range[1])] + else: + # Single value (int or string), use it for both start and end + filter_dict["time_int"] = [int(time_range), int(time_range)] + logger.info(f"memory_type={self.memory_type} memory_target={self.memory_target} query={query} " + f"filter_dict={filter_dict}") + + 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: + output = [] + for node in new_memory_nodes: + line = "" + if "conversation_time" in node.metadata and node.metadata["conversation_time"]: + line += f"conversation_time={node.metadata['conversation_time']} " + line += node.content.strip() + " " + if node.ref_memory_id: + line += f"history_id={node.ref_memory_id} " + output.append(line.strip()) + self.output = "### Extracted Memories\n" + "\n".join(output) + + logger.info(f"Retrieved {len(memory_nodes)} memory_nodes, {len(new_memory_nodes)} new after deduplication") 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..f6336395 --- /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..6ec678c9 100644 --- a/reme_ai/reme.py +++ b/reme_ai/reme.py @@ -1,18 +1,29 @@ """ReMe classes for simplified configuration and execution.""" -from .core.application import Application -from .core.config import ReMeConfigParser -from .core.context import C -from .core.embedding import BaseEmbeddingModel -from .core.enumeration import Role -from .core.llm import BaseLLM -from .core.schema import Message -from .core.utils import singleton -from .core.vector_store import BaseVectorStore +from .core_old.application import Application +from .core_old.config import ReMeConfigParser +from .core_old.context import C +from .core_old.embedding import BaseEmbeddingModel +from .core_old.enumeration import Role +from .core_old.llm import BaseLLM +from .core_old.schema import Message +from .core_old.utils import singleton +from .core_old.vector_store import BaseVectorStore 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_agent.v4 import ( + PersonalSummarizerV4, + PersonalRetrieverV4, + ReMeRetrieverV4, + ReMeSummarizerV4, +) from .mem_tool import ( HandsOffTool, ReadHistoryMemory, @@ -24,15 +35,29 @@ 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, +) +from .mem_tool.v4 import ( + AddSummaryMemory as AddSummaryMemoryV4, + HandsOff as HandsOffV4, + ReadHistory as ReadHistoryV4, + ReadUserProfile as ReadUserProfileV4, + RetrieveMemory as RetrieveMemoryV4, + UpdateUserProfile as UpdateUserProfileV4, +) -@singleton class ReMe(Application): """Simplified ReMe application that auto-initializes the service context.""" @@ -151,7 +176,6 @@ class ReMe(Application): except Exception as e: print(f"Warning: reme_summarizer.call failed: {e}") return [] - else: raise NotImplementedError @@ -202,18 +226,17 @@ class ReMe(Application): except Exception as e: print(f"Warning: reme_retriever.call failed: {e}") return "error, not retrieved" - else: raise NotImplementedError async def summary_v2( - self, - messages: list[dict], - description: str = "", - user_id: str = "", - assistant_id: str = "", - **kwargs, + self, + messages: list[dict], + description: str = "", + user_id: str = "", + assistant_id: str = "", + **kwargs, ): """Summarizes messages using V2 workflow with simplified tools.""" @@ -267,14 +290,14 @@ class ReMe(Application): raise NotImplementedError async def retrieve_v2( - self, - query: str = "", - messages: list[dict] | None = None, - description: str = "", - user_id: str = "", - assistant_id: str = "", - top_k: int = 20, - **kwargs, + 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 V2 workflow with autonomous retrieval.""" @@ -314,3 +337,166 @@ 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(enable_thinking_params=True), + ReadUserProfile(enable_thinking_params=True, add_memory_type_target=False), + UpdateUserProfile(enable_thinking_params=True), + ], + ) + + 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(enable_thinking_params=True, add_memory_type_target=True), + RetrieveMemory(enable_thinking_params=True, top_k=top_k), + ReadHistoryV3(enable_thinking_params=True), + ], + ) + + # 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 + + async def summary_v4( + self, + messages: list[dict], + description: str = "", + user_id: str = "", + assistant_id: str = "", + enable_thinking_params: bool = False, + **kwargs, + ): + """Summarizes messages using V4 workflow with simplified memory management.""" + + if user_id: + meta_memories = [ + { + "memory_type": "personal", + "memory_target": user_id, + }, + ] + messages = self._prepare_messages(messages, user_id, assistant_id) + + personal_summarizer_v4 = PersonalSummarizerV4( + tools=[ + AddSummaryMemoryV4(enable_thinking_params=enable_thinking_params), + ReadUserProfileV4(enable_thinking_params=enable_thinking_params), + UpdateUserProfileV4(enable_thinking_params=enable_thinking_params), + ], + ) + + reme_summarizer_v4 = ReMeSummarizerV4( + meta_memories=meta_memories, + tools=[HandsOffV4(memory_agents=[personal_summarizer_v4])], + ) + + await reme_summarizer_v4.call(messages=messages, description=description, **kwargs) + return reme_summarizer_v4.memory_nodes, reme_summarizer_v4.tool_messages, reme_summarizer_v4.success + + else: + raise NotImplementedError + + async def retrieve_v4( + self, + query: str = "", + messages: list[dict] | None = None, + description: str = "", + user_id: str = "", + assistant_id: str = "", + top_k: int = 20, + enable_thinking_params: bool = False, + **kwargs, + ): + """Retrieves relevant memories using V4 workflow with enhanced retrieval.""" + + if user_id: + messages = self._prepare_messages(messages, user_id, assistant_id) + + meta_memories = [ + { + "memory_type": "personal", + "memory_target": user_id, + }, + ] + + personal_retriever_v4 = PersonalRetrieverV4( + tools=[ + RetrieveMemoryV4(enable_thinking_params=enable_thinking_params, top_k=top_k), + ReadHistoryV4(enable_thinking_params=enable_thinking_params), + ], + ) + + reme_retriever_v4 = ReMeRetrieverV4( + meta_memories=meta_memories, + tools=[HandsOffV4(memory_agents=[personal_retriever_v4])], + ) + + await reme_retriever_v4.call(query=query, messages=messages, description=description, **kwargs) + return reme_retriever_v4.output, reme_retriever_v4.tool_messages, reme_retriever_v4.success + + else: + raise NotImplementedError diff --git a/reme_ai/tool/__init__.py b/reme_ai/tool/__init__.py deleted file mode 100644 index 653cbed4..00000000 --- a/reme_ai/tool/__init__.py +++ /dev/null @@ -1,9 +0,0 @@ -"""tool""" - -from . import execute -from . import search - -__all__ = [ - "execute", - "search", -] diff --git a/reme_ai/tool/execute/__init__.py b/reme_ai/tool/execute/__init__.py deleted file mode 100644 index f78a73ac..00000000 --- a/reme_ai/tool/execute/__init__.py +++ /dev/null @@ -1,9 +0,0 @@ -"""execute tool""" - -from .execute_code import ExecuteCode -from .execute_shell import ExecuteShell - -__all__ = [ - "ExecuteCode", - "ExecuteShell", -] diff --git a/test/http_client_test.py b/test/http_client_test.py deleted file mode 100644 index acd1254c..00000000 --- a/test/http_client_test.py +++ /dev/null @@ -1,165 +0,0 @@ -import asyncio -import json - -import aiohttp - -base_url = "http://0.0.0.0:8002" - - -async def run1(session): - workspace_id = "default1" - - async with session.post( - f"{base_url}/vector_store", - json={ - "action": "delete", - "workspace_id": workspace_id, - }, - headers={"Content-Type": "application/json"}, - ) as response: - result = await response.json() - print(json.dumps(result, ensure_ascii=False)) - - trajectories = [ - { - "task_id": "t1", - "messages": [ - {"role": "user", "content": "搜索可以使用websearch工具"}, - ], - "score": 1, - }, - { - "task_id": "t1", - "messages": [ - {"role": "user", "content": "搜索可以使用code工具"}, - ], - "score": 0, - }, - ] - - async with session.post( - # f"{base_url}/summary_task_memory", - f"{base_url}/summary_task_memory_simple", - json={ - "trajectories": trajectories, - "workspace_id": workspace_id, - }, - headers={"Content-Type": "application/json"}, - ) as response: - result = await response.json() - print(json.dumps(result, ensure_ascii=False)) - - await asyncio.sleep(2) - - async with session.post( - # f"{base_url}/retrieve_task_memory", - f"{base_url}/retrieve_task_memory_simple", - json={ - "query": "茅台怎么样?", - "workspace_id": workspace_id, - }, - headers={"Content-Type": "application/json"}, - ) as response: - result = await response.json() - print(json.dumps(result, ensure_ascii=False)) - - -async def run2(session): - workspace_id = "default2" - - async with session.post( - f"{base_url}/vector_store", - json={ - "action": "delete", - "workspace_id": workspace_id, - }, - headers={"Content-Type": "application/json"}, - ) as response: - result = await response.json() - print(json.dumps(result, ensure_ascii=False)) - - messages = [ - {"role": "user", "content": "我喜欢吃西瓜🍉"}, - {"role": "user", "content": "昨天吃了苹果,很好吃"}, - {"role": "user", "content": "我不太喜欢吃西瓜"}, - {"role": "user", "content": "上周我去了日本,得了肠胃炎"}, - {"role": "user", "content": "这周只能在家里,喝粥"}, - ] - - async with session.post( - f"{base_url}/summary_personal_memory", - json={ - "messages": messages, - "workspace_id": workspace_id, - }, - headers={"Content-Type": "application/json"}, - ) as response: - result = await response.json() - print(json.dumps(result, ensure_ascii=False)) - - await asyncio.sleep(2) - - async with session.post( - f"{base_url}/retrieve_personal_memory", - json={ - "query": "你知道我喜欢吃什么?", - "workspace_id": workspace_id, - }, - headers={"Content-Type": "application/json"}, - ) as response: - result = await response.json() - print(json.dumps(result, ensure_ascii=False)) - - -async def run3(session): - workspace_id = "default2" - - async with session.post( - f"{base_url}/add_tool_call_result", - json={ - "tool_call_results": [ - {"a": 1}, - {"a": 2}, - ], - "workspace_id": workspace_id, - }, - headers={"Content-Type": "application/json"}, - ) as response: - result = await response.json() - print(json.dumps(result, ensure_ascii=False)) - - -async def run4(session): - workspace_id = "default4" - - async with session.post( - f"{base_url}/agentic_retrieve", - json={ - "messages": [ - {"role": "user", "content": "hello" * 10000}, - ], - "workspace_id": workspace_id, - "context_manage_mode": "auto", - "keep_recent_count": 0, - "max_total_tokens": 10000, - }, - headers={"Content-Type": "application/json"}, - ) as response: - result = await response.json() - print(json.dumps(result, ensure_ascii=False)) - - -async def main(): - - async with aiohttp.ClientSession() as session: - # 获取工具列表 - print("获取工具列表...") - - # await run1(session) - # await run2(session) - # await run3(session) - await run4(session) - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/test/mcp_client_test.py b/test/mcp_client_test.py deleted file mode 100644 index 9bb08b8b..00000000 --- a/test/mcp_client_test.py +++ /dev/null @@ -1,45 +0,0 @@ -from fastmcp import Client -from mcp.types import CallToolResult - - -async def main(): - async with Client("http://0.0.0.0:8002/sse/") as client: - tools = await client.list_tools() - for tool in tools: - print(tool.model_dump_json()) - - workspace_id = "default" - - result: CallToolResult = await client.call_tool( - "retrieve_task_memory_simple", - arguments={ - "query": "茅台怎么样?", - "workspace_id": workspace_id, - }, - ) - print(result.content) - - trajectories = [ - { - "task_id": "t1", - "messages": [ - {"role": "user", "content": "今天天气不错"}, - ], - "score": 0.9, - }, - ] - - result: CallToolResult = await client.call_tool( - "summary_task_memory_simple", - arguments={ - "trajectories": trajectories, - "workspace_id": workspace_id, - }, - ) - print(result.content) - - -if __name__ == "__main__": - import asyncio - - asyncio.run(main()) diff --git a/test/record_audio.py b/test/record_audio.py deleted file mode 100644 index d9446813..00000000 --- a/test/record_audio.py +++ /dev/null @@ -1,153 +0,0 @@ -#!/usr/bin/env python3 -""" -macOS 麦克风录音脚本 -需要安装: pip install pyaudio wave -""" - -import pyaudio -import wave -import sys -import os -from datetime import datetime - - -class AudioRecorder: - """macOS 音频录制器""" - - def __init__(self, output_dir="recordings"): - """ - 初始化录音器 - - Args: - output_dir: 录音文件保存目录 - """ - self.output_dir = output_dir - self.chunk = 1024 # 每次读取的音频块大小 - self.format = pyaudio.paInt16 # 16位深度 - self.channels = 1 # 单声道 - self.rate = 44100 # 采样率 44.1kHz - - # 创建输出目录 - if not os.path.exists(output_dir): - os.makedirs(output_dir) - - def record(self, duration=5, filename=None): - """ - 录制音频 - - Args: - duration: 录制时长(秒) - filename: 输出文件名,如果为None则自动生成 - - Returns: - str: 保存的文件路径 - """ - # 生成文件名 - if filename is None: - timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") - filename = f"recording_{timestamp}.wav" - - filepath = os.path.join(self.output_dir, filename) - - # 初始化PyAudio - audio = pyaudio.PyAudio() - - try: - # 打开音频流(这会触发macOS的麦克风权限请求) - print("正在请求麦克风权限...") - stream = audio.open( - format=self.format, - channels=self.channels, - rate=self.rate, - input=True, - frames_per_buffer=self.chunk - ) - - print(f"开始录音,时长: {duration} 秒") - print("录音中...") - - frames = [] - - # 录制音频 - for i in range(0, int(self.rate / self.chunk * duration)): - data = stream.read(self.chunk) - frames.append(data) - - # 显示进度 - progress = (i + 1) / (self.rate / self.chunk * duration) * 100 - sys.stdout.write(f"\r进度: {progress:.1f}%") - sys.stdout.flush() - - print("\n录音完成!") - - # 停止并关闭流 - stream.stop_stream() - stream.close() - - # 保存为WAV文件 - print(f"正在保存到: {filepath}") - wf = wave.open(filepath, 'wb') - wf.setnchannels(self.channels) - wf.setsampwidth(audio.get_sample_size(self.format)) - wf.setframerate(self.rate) - wf.writeframes(b''.join(frames)) - wf.close() - - print(f"✓ 文件已保存: {filepath}") - return filepath - - except Exception as e: - print(f"\n错误: {e}") - print("\n提示:") - print("1. 请确保已安装 pyaudio: pip install pyaudio") - print("2. 在macOS上,首次运行会弹出权限请求对话框") - print("3. 如果权限被拒绝,请前往 系统偏好设置 > 安全性与隐私 > 隐私 > 麦克风") - return None - - finally: - audio.terminate() - - def record_interactive(self): - """交互式录音""" - print("=" * 50) - print("macOS 麦克风录音工具") - print("=" * 50) - - try: - duration = input("\n请输入录音时长(秒,默认5秒): ").strip() - duration = int(duration) if duration else 5 - - filename = input("请输入文件名(留空自动生成): ").strip() - filename = filename if filename else None - if filename and not filename.endswith('.wav'): - filename += '.wav' - - print() - self.record(duration=duration, filename=filename) - - except KeyboardInterrupt: - print("\n\n录音已取消") - except ValueError: - print("输入无效,请输入数字") - - -def main(): - """主函数""" - recorder = AudioRecorder() - - if len(sys.argv) > 1: - # 命令行模式 - try: - duration = int(sys.argv[1]) - filename = sys.argv[2] if len(sys.argv) > 2 else None - recorder.record(duration=duration, filename=filename) - except ValueError: - print("用法: python record_audio.py [时长(秒)] [文件名(可选)]") - print("示例: python record_audio.py 10 my_recording.wav") - else: - # 交互式模式 - recorder.record_interactive() - - -if __name__ == "__main__": - main() diff --git a/test/test1.py b/test/test1.py deleted file mode 100644 index bb73331a..00000000 --- a/test/test1.py +++ /dev/null @@ -1,12 +0,0 @@ -# 2025年半年报点评:Q2业绩同比增长,CPU、DCU业务进展顺利 -# https://data.eastmoney.com/report/info/AP202508061722561937.html -# -# https://pdf.dfcfw.com/pdf/H3_AP202508061722561937_1.pdf - -import requests - -headers = { - "User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64)", -} -url = requests.get("https://data.eastmoney.com/report/stock.jshtml", headers=headers) -print(url.text) diff --git a/test/test2.py b/test/test2.py deleted file mode 100644 index acd329c3..00000000 --- a/test/test2.py +++ /dev/null @@ -1,393 +0,0 @@ -import json -import os -import random -import re -from datetime import datetime, timedelta -from io import BytesIO -from time import sleep -from urllib.parse import urljoin - -import pycurl -import requests -from PyPDF2 import PdfReader - -# 全局配置 -BASE_URL = "https://reportapi.eastmoney.com/report/list" -DETAIL_BASE_URL = "https://data.eastmoney.com/report/info/" - -# 读取config.json获取stock_code -with open("config.json", "r", encoding="utf-8") as f: - config = json.load(f) -STOCK_CODE = config.get("stock_code", "600519") -MIN_PAGES = config.get("min_pages", 20) -DOWNLOAD_DIR = config.get("download_dir", "reports_pdf") -YEARS_AGO = config.get("years_ago", 2) -os.makedirs(DOWNLOAD_DIR, exist_ok=True) - -# 随机User-Agent列表 -USER_AGENTS = [ - "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/91.0.4472.124 Safari/537.36", - "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/91.0.4472.124 Safari/537.36", - "Mozilla/5.0 (Windows NT 10.0; Win64; x64; rv:89.0) Gecko/20100101 Firefox/89.0", -] - - -def get_random_user_agent(): - """获取随机User-Agent""" - import random - - return random.choice(USER_AGENTS) - - -def fetch_jsonp_data(page_no=1): - """ - 获取研究报告列表数据 - :param page_no: 页码 - :return: 解析后的数据字典 - """ - # 计算日期 - today = datetime.today() - end_time = today.strftime("%Y-%m-%d") - begin_time = (today - timedelta(days=365 * YEARS_AGO)).strftime("%Y-%m-%d") - - # 检查是否存在已保存的原始数据 - raw_data_dir = "raw_data" - raw_data_file = os.path.join(raw_data_dir, f"page_{page_no}_{STOCK_CODE}_{begin_time}_{end_time}.json") - - if os.path.exists(raw_data_file): - print(f"使用已保存的原始数据: {raw_data_file}") - try: - with open(raw_data_file, "r", encoding="utf-8") as f: - return json.load(f) - except Exception as e: - print(f"读取已保存数据失败: {e}") - - params = { - "cb": "datatable6333112", - "pageNo": page_no, - "pageSize": 50, - "code": STOCK_CODE, - "industryCode": "*", - "industry": "*", - "rating": "*", - "ratingchange": "*", - "beginTime": begin_time, - "endTime": end_time, - "fields": "", - "qType": 0, - "p": page_no, - "pageNum": page_no, - "pageNumber": page_no, - "_": int(time.time() * 1000), # 使用当前时间戳 - } - headers = { - "User-Agent": get_random_user_agent(), - "Referer": "https://data.eastmoney.com/", - } - try: - response = requests.get(BASE_URL, params=params, headers=headers) - response.raise_for_status() - # 提取JSON部分 - json_str = re.search(r"\((.*)\)", response.text).group(1) - data = json.loads(json_str) - - # 保存原始数据到本地 - if not os.path.exists(raw_data_dir): - os.makedirs(raw_data_dir, exist_ok=True) - - with open(raw_data_file, "w", encoding="utf-8") as f: - json.dump(data, f, ensure_ascii=False, indent=2) - - print(f"原始数据已保存: {raw_data_file}") - return data - except Exception as e: - print(f"获取第{page_no}页数据失败: {e}") - return None - - -def get_report_detail(info_code): - """ - 获取研究报告详情页内容 - :param info_code: 报告ID - :return: 详情页HTML内容 - """ - # 检查是否存在已保存的详情页HTML - detail_data_dir = "detail_data" - detail_html_file = os.path.join(detail_data_dir, f"detail_{info_code}.html") - - if os.path.exists(detail_html_file): - print(f"使用已保存的详情页HTML: {detail_html_file}") - try: - with open(detail_html_file, "r", encoding="utf-8") as f: - return f.read() - except Exception as e: - print(f"读取已保存详情页失败: {e}") - - url = urljoin(DETAIL_BASE_URL, f"{info_code}.html") - headers = { - "User-Agent": get_random_user_agent(), - "Referer": "https://data.eastmoney.com/", - } - - try: - response = requests.get(url, headers=headers) - response.raise_for_status() - - # 保存详情页HTML原始数据 - if not os.path.exists(detail_data_dir): - os.makedirs(detail_data_dir, exist_ok=True) - - with open(detail_html_file, "w", encoding="utf-8") as f: - f.write(response.text) - - print(f"详情页HTML已保存: {detail_html_file}") - return response.text - except Exception as e: - print(f"获取报告详情{info_code}失败: {e}") - return None - - -def parse_detail_page(html, info_code): - """ - 解析详情页获取PDF下载链接及相关信息 - :param html: 详情页HTML - :param info_code: 报告ID - :return: dict,包含PDF下载URL及命名所需字段 - """ - try: - # 使用正则提取zwinfo变量 - match = re.search(r"var zwinfo\s*=\s*({.*?});", html, re.DOTALL) - if not match: - return None - zwinfo = json.loads(match.group(1)) - - # 保存解析后的zwinfo数据 - detail_data_dir = "detail_data" - zwinfo_file = os.path.join(detail_data_dir, f"zwinfo_{info_code}.json") - with open(zwinfo_file, "w", encoding="utf-8") as f: - json.dump(zwinfo, f, ensure_ascii=False, indent=2) - - print(f"zwinfo数据已保存: {zwinfo_file}") - - # 提取所需字段 - return { - "attach_url": zwinfo.get("attach_url"), - "notice_title": zwinfo.get("notice_title", ""), - "short_name": zwinfo.get("short_name", ""), - "notice_date": zwinfo.get("notice_date", ""), - "source_sample_name": zwinfo.get("source_sample_name", ""), - "attach_pages": zwinfo.get("attach_pages", ""), - } - except Exception as e: - print(f"解析详情页失败: {e}") - return None - - -def is_pdf_complete(pdf_path, expected_pages): - """ - 检查PDF页数是否与预期一致 - :param pdf_path: PDF文件路径 - :param expected_pages: 预期页数(int) - :return: bool - """ - try: - with open(pdf_path, "rb") as f: - reader = PdfReader(f) - actual_pages = len(reader.pages) - return actual_pages == expected_pages, actual_pages - except Exception as e: - print(f"读取PDF页数失败: {e}") - return False, 0 - - -def download_pdf(pdf_url, filename): - """ - 使用pycurl下载PDF文件(模拟curl请求) - - 参数: - pdf_url (str): PDF文件的URL - filename (str): 保存文件名(不含路径) - - 返回: - bool: 是否下载成功 - """ - save_path = os.path.join(DOWNLOAD_DIR, filename) - buffer = BytesIO() - c = pycurl.Curl() - - try: - # 设置curl选项 - c.setopt(pycurl.URL, pdf_url) - c.setopt(pycurl.WRITEDATA, buffer) - c.setopt(pycurl.FOLLOWLOCATION, True) - c.setopt(pycurl.MAXREDIRS, 5) - c.setopt(pycurl.CONNECTTIMEOUT, 30) - c.setopt(pycurl.TIMEOUT, 300) - - # 设置防爬虫headers - headers = [ - f"User-Agent: {get_random_user_agent()}", - "Accept: text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8", - "Referer: https://data.eastmoney.com/", - "Accept-Language: zh-CN,zh;q=0.9", - ] - c.setopt(pycurl.HTTPHEADER, headers) - - # 执行下载 - c.perform() - - # 验证响应 - if c.getinfo(pycurl.HTTP_CODE) != 200: - print(f"下载失败 HTTP {c.getinfo(pycurl.HTTP_CODE)}") - return False - - # 保存文件 - with open(save_path, "wb") as f: - f.write(buffer.getvalue()) - - print(f"✓ 成功下载 {filename}") - return True - - except pycurl.error as e: - errno, errstr = e.args - print(f"pycurl错误({errno}): {errstr}") - return False - except Exception as e: - print(f"下载异常: {str(e)}") - return False - finally: - c.close() - buffer.close() - - -def process_all_reports(): - """处理所有研究报告""" - # 获取第一页数据 - first_page_data = fetch_jsonp_data(1) - if not first_page_data: - return - - total_page = first_page_data.get("TotalPage", 1) - total_reports = first_page_data.get("hits", 0) - print(f"共发现{total_reports}篇研究报告,{total_page}页") - - # 处理所有页面 - for page in range(1, total_page + 1): - print(f"\n正在处理第{page}/{total_page}页...") - # 获取当前页数据 - if page == 1: - page_data = first_page_data - else: - page_data = fetch_jsonp_data(page) - if not page_data: - continue - # 处理每篇报告 - report_list = page_data.get("data", []) - random.shuffle(report_list) - for report in report_list: - info_code = report.get("infoCode") - if not info_code: - continue - - # 检查页数,只有大于20页的才下载 - attach_pages = report.get("attachPages", 0) - try: - attach_pages = int(attach_pages) - except (ValueError, TypeError): - attach_pages = 0 - - if attach_pages < MIN_PAGES: - print(f"跳过页数不足的报告: {report.get('title')} (页数: {attach_pages})") - continue - - print(f"\n处理报告: {report.get('title')} [{info_code}] (页数: {attach_pages})") - # 获取详情页 - detail_html = get_report_detail(info_code) - if not detail_html: - continue - # 解析PDF链接及命名信息 - detail_info = parse_detail_page(detail_html, info_code) - if not detail_info or not detail_info.get("attach_url"): - print("未找到PDF链接") - continue - # 组装文件名,避免重复拼接 - notice_title = detail_info.get("notice_title", "").strip().replace("/", "_") - short_name = detail_info.get("short_name", "").strip().replace("/", "_") - notice_date = detail_info.get("notice_date", "").replace("-", "")[:8] # 只取年月日 - source_sample_name = detail_info.get("source_sample_name", "").strip().replace("/", "_") - - filename_parts = [] - filename_parts.append(notice_date) - # 判断source_sample_name是否已在notice_title中 - if source_sample_name and source_sample_name not in notice_title: - filename_parts.append(source_sample_name) - # 判断short_name是否已在notice_title中 - if short_name and short_name not in notice_title: - filename_parts.append(short_name) - filename_parts.append(notice_title) - # 分离文件名和目录 - pdf_filename = f"{'_'.join(filename_parts)}.pdf" - pdf_subdir = f"{short_name}" - - # 判断是否为深度报告(页数大于20页) - if attach_pages >= 20: - pdf_subdir = f"{short_name}/深度报告" - - pdf_full_path = os.path.join(DOWNLOAD_DIR, pdf_subdir, pdf_filename) - - # 检查并创建目录 - pdf_dir = os.path.join(DOWNLOAD_DIR, pdf_subdir) - if not os.path.exists(pdf_dir): - os.makedirs(pdf_dir, exist_ok=True) - print(f"创建目录: {pdf_dir}") - - # 检查文件是否已存在 - if os.path.exists(pdf_full_path): - print(f"文件已存在,跳过下载: {pdf_full_path}") - continue - - # 下载PDF并校验页数,最多重试3次 - max_retries = 5 - for attempt in range(1, max_retries + 1): - download_pdf(detail_info["attach_url"], os.path.join(pdf_subdir, pdf_filename)) - # 校验PDF页数 - try: - expected_pages = int(detail_info.get("attach_pages", 0)) - except Exception: - expected_pages = 0 - is_complete = True - actual_pages = 0 - if expected_pages > 0: - is_complete, actual_pages = is_pdf_complete(pdf_full_path, expected_pages) - if is_complete: - print(f"✓ PDF页数校验通过:{actual_pages}页") - break - else: - print( - f"✗ PDF页数不符:实际{actual_pages}页,预期{expected_pages}页,正在重试({attempt}/{max_retries})...", - ) - # 删除不完整文件 - try: - os.remove(pdf_full_path) - except Exception: - pass - sleep(1) - else: - break - sleep(60 * attempt) - - else: - print(f"!!! PDF多次下载后仍不完整:{pdf_full_path}") - # 礼貌性延迟 - sleep(30) - - -if __name__ == "__main__": - import time - - start_time = time.time() - - process_all_reports() - - end_time = time.time() - print(f"\n全部完成,耗时: {end_time - start_time:.2f}秒") diff --git a/test/test3.py b/test/test3.py deleted file mode 100644 index 2b0021a8..00000000 --- a/test/test3.py +++ /dev/null @@ -1,92 +0,0 @@ -import os -from io import BytesIO - -import pycurl -from PyPDF2 import PdfReader - -DOWNLOAD_DIR = "./" - -# 随机User-Agent列表 -USER_AGENTS = [ - "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/91.0.4472.124 Safari/537.36", - "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/91.0.4472.124 Safari/537.36", - "Mozilla/5.0 (Windows NT 10.0; Win64; x64; rv:89.0) Gecko/20100101 Firefox/89.0", -] - - -def get_random_user_agent(): - """获取随机User-Agent""" - import random - - return random.choice(USER_AGENTS) - - -def download_pdf(pdf_url, filename): - """ - 使用pycurl下载PDF文件(模拟curl请求) - - 参数: - pdf_url (str): PDF文件的URL - filename (str): 保存文件名(不含路径) - - 返回: - bool: 是否下载成功 - """ - save_path = os.path.join(DOWNLOAD_DIR, filename) - buffer = BytesIO() - c = pycurl.Curl() - - try: - # 设置curl选项 - c.setopt(pycurl.URL, pdf_url) - c.setopt(pycurl.WRITEDATA, buffer) - c.setopt(pycurl.FOLLOWLOCATION, True) - c.setopt(pycurl.MAXREDIRS, 5) - c.setopt(pycurl.CONNECTTIMEOUT, 30) - c.setopt(pycurl.TIMEOUT, 300) - - # 设置防爬虫headers - headers = [ - f"User-Agent: {get_random_user_agent()}", - "Accept: text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8", - "Referer: https://data.eastmoney.com/", - "Accept-Language: zh-CN,zh;q=0.9", - ] - c.setopt(pycurl.HTTPHEADER, headers) - - # 执行下载 - c.perform() - - # 验证响应 - if c.getinfo(pycurl.HTTP_CODE) != 200: - print(f"下载失败 HTTP {c.getinfo(pycurl.HTTP_CODE)}") - return False - - # 保存文件 - with open(save_path, "wb") as f: - f.write(buffer.getvalue()) - - print(f"✓ 成功下载 {filename}") - return True - - except pycurl.error as e: - errno, errstr = e.args - print(f"pycurl错误({errno}): {errstr}") - return False - except Exception as e: - print(f"下载异常: {str(e)}") - return False - finally: - c.close() - buffer.close() - - -if __name__ == "__main__": - url_list = [ - "https://pdf.dfcfw.com/pdf/H3_AP202508061722531920_1.pdf?1754495126000.pdf", - ] - - url_list = [x.split("?")[0] for x in url_list] - for url in url_list: - name = url.split("_")[1] - download_pdf(url, f"{name}.pdf") diff --git a/test/test4.py b/test/test4.py deleted file mode 100644 index bc1b703e..00000000 --- a/test/test4.py +++ /dev/null @@ -1,67 +0,0 @@ -#!/usr/bin/env python3 -# -*- coding: utf-8 -*- - - -def analyze_corrupted_text(text): - """分析乱码文本的字节构成""" - print(f"分析文本: {text}") - print(f"文本长度: {len(text)}") - - # 显示每个字符的Unicode码点 - print("字符分析:") - for i, char in enumerate(text[:20]): # 只显示前20个字符 - print(f" {i}: '{char}' -> U+{ord(char):04X}") - - # 尝试不同的编码方式 - print("\n编码尝试:") - - try: - # 方法1: Latin1 -> UTF-8 - bytes_latin1 = text.encode("latin1") - result_utf8 = bytes_latin1.decode("utf-8") - print(f"Latin1->UTF-8: {result_utf8}") - except Exception as e: - print(f"Latin1->UTF-8 失败: {e}") - - try: - # 方法2: Latin1 -> GBK - bytes_latin1 = text.encode("latin1") - result_gbk = bytes_latin1.decode("gbk") - print(f"Latin1->GBK: {result_gbk}") - except Exception as e: - print(f"Latin1->GBK 失败: {e}") - - try: - # 方法3: CP1252 -> UTF-8 - bytes_cp1252 = text.encode("cp1252") - result_utf8 = bytes_cp1252.decode("utf-8") - print(f"CP1252->UTF-8: {result_utf8}") - except Exception as e: - print(f"CP1252->UTF-8 失败: {e}") - - # 显示原始字节 - try: - raw_bytes = text.encode("latin1") - print(f"\n原始字节 (Latin1): {raw_bytes}") - print(f"字节十六进制: {raw_bytes.hex()}") - except Exception as e: - print(f"获取原始字节失败: {e}") - - -def main(): - """调试主函数""" - test_texts = [ - "为ä»ä¹è¯´æçå»è¯è¿å¥ä¸­æå¸å±æç¹ï¼", - "åçäºâäºä¸âæé´ä¸­å½ç»æµå¤è¯éªçä¹è§å¤æ­", - "æçç§ææ°ï¼HSTECH.HIï¼åº¦æ¼æ¶4.45%", - ] - - for i, text in enumerate(test_texts, 1): - print(f"\n{'=' * 60}") - print(f"测试 {i}") - print("=" * 60) - analyze_corrupted_text(text) - - -if __name__ == "__main__": - main() diff --git a/test/test5.py b/test/test5.py deleted file mode 100644 index 6220633e..00000000 --- a/test/test5.py +++ /dev/null @@ -1,7 +0,0 @@ -import tiktoken - -enc = tiktoken.get_encoding("o200k_base") - -# r = enc.encode("我爱吃西瓜,你说啥") -r = enc.encode("hello world aaaaaaaaaaaa") -print(len(r)) diff --git a/test/test6.py b/test/test6.py deleted file mode 100644 index 7e555898..00000000 --- a/test/test6.py +++ /dev/null @@ -1,15 +0,0 @@ -import tiktoken - - -def count_tokens(text: str) -> int: - """计算给定文本在指定模型下的 token 数量""" - encoding = tiktoken.get_encoding("o200k_base") - tokens = encoding.encode(text) - return len(tokens) - - -# 示例使用 -text = "你好,世界!Hello, world!" -token_count = count_tokens(text) -print(f"Token 数量: {token_count}") -print(len(text) / 4) diff --git a/test_op/test_agentic_retrieve_op.py b/test/test_agentic_retrieve_op.py similarity index 100% rename from test_op/test_agentic_retrieve_op.py rename to test/test_agentic_retrieve_op.py diff --git a/test_op/test_message_compact_op.py b/test/test_message_compact_op.py similarity index 100% rename from test_op/test_message_compact_op.py rename to test/test_message_compact_op.py diff --git a/test_op/test_message_compress_op.py b/test/test_message_compress_op.py similarity index 100% rename from test_op/test_message_compress_op.py rename to test/test_message_compress_op.py diff --git a/test_op/test_message_offload_op.py b/test/test_message_offload_op.py similarity index 100% rename from test_op/test_message_offload_op.py rename to test/test_message_offload_op.py diff --git a/tests/test_reme.py b/test/test_reme.py similarity index 97% rename from tests/test_reme.py rename to test/test_reme.py index c9e6843e..3279da2a 100644 --- a/tests/test_reme.py +++ b/test/test_reme.py @@ -2,8 +2,9 @@ import asyncio -from reme_ai.core.schema import VectorNode, MemoryNode -from reme_ai.reme import ReMe +from reme.reme import ReMe + +from reme.core.schema import VectorNode, MemoryNode reme = ReMe( vector_store={"collection_name": "reme"}, diff --git a/test/test_update_insight_op.py b/test/test_update_insight_op.py deleted file mode 100644 index 6505b6a4..00000000 --- a/test/test_update_insight_op.py +++ /dev/null @@ -1,128 +0,0 @@ -#!/usr/bin/env python3 -""" -Simple test script to verify the UpdateInsightOp implementation. -This is a basic validation test to ensure the class structure is correct. -""" - -import sys - -sys.path.append("/Users/yuli/workspace/MemoryScope") - - -def test_update_insight_op_import(): - """Test that we can import the UpdateInsightOp class""" - try: - from reme_ai.summary.personal.update_insight_op import UpdateInsightOp - - print("✓ Successfully imported UpdateInsightOp") - return True - except ImportError as e: - print(f"✗ Failed to import UpdateInsightOp: {e}") - return False - - -def test_personal_memory_import(): - """Test that we can import PersonalMemory""" - try: - from reme_ai.schema.memory import PersonalMemory - - print("✓ Successfully imported PersonalMemory") - return True - except ImportError as e: - print(f"✗ Failed to import PersonalMemory: {e}") - return False - - -def test_op_utils_import(): - """Test that we can import the utility functions""" - try: - from reme_ai.utils.op_utils import parse_update_insight_response - - print("✓ Successfully imported parse_update_insight_response") - return True - except ImportError as e: - print(f"✗ Failed to import parse_update_insight_response: {e}") - return False - - -def test_personal_memory_creation(): - """Test PersonalMemory creation with reflection_subject""" - try: - from reme_ai.schema.memory import PersonalMemory - - memory = PersonalMemory( - workspace_id="test_workspace", - content="User likes playing basketball", - target="test_user", - reflection_subject="hobbies", - author="test_system", - ) - - print(f"✓ Created PersonalMemory: {memory.content}") - print(f" - Memory ID: {memory.memory_id}") - print(f" - Target: {memory.target}") - print(f" - Reflection Subject: {memory.reflection_subject}") - return True - except Exception as e: - print(f"✗ Failed to create PersonalMemory: {e}") - return False - - -def test_parse_update_insight_response(): - """Test the parse_update_insight_response function""" - try: - from reme_ai.utils.op_utils import parse_update_insight_response - - # Test Chinese format - chinese_response = "思考:用户喜欢篮球和足球\ntest_user的资料:<喜欢篮球和足球>" - result_zh = parse_update_insight_response(chinese_response, "zh") - print(f"✓ Parsed Chinese response: '{result_zh}'") - - # Test English format - english_response = ( - "Thoughts: User likes basketball and football\ntest_user's profile: " - ) - result_en = parse_update_insight_response(english_response, "en") - print(f"✓ Parsed English response: '{result_en}'") - - return True - except Exception as e: - print(f"✗ Failed to test parse_update_insight_response: {e}") - return False - - -def main(): - """Run all tests""" - print("Running UpdateInsightOp validation tests...\n") - - tests = [ - test_personal_memory_import, - test_op_utils_import, - test_update_insight_op_import, - test_personal_memory_creation, - test_parse_update_insight_response, - ] - - passed = 0 - total = len(tests) - - for test in tests: - print(f"\nRunning {test.__name__}:") - if test(): - passed += 1 - print() - - print("=" * 50) - print(f"Test Results: {passed}/{total} passed") - - if passed == total: - print("🎉 All tests passed! The UpdateInsightOp implementation looks good.") - else: - print("⚠️ Some tests failed. Please check the implementation.") - - return passed == total - - -if __name__ == "__main__": - success = main() - sys.exit(0 if success else 1) diff --git a/tests/test_base_context.py b/tests/test_base_context.py index 316a6796..61a355c0 100644 --- a/tests/test_base_context.py +++ b/tests/test_base_context.py @@ -4,7 +4,8 @@ Ensures attribute-style and dict-style access work interchangeably. """ import pickle -from reme_ai.core.context import BaseContext + +from reme.core.context import BaseContext def test_attribute_access(): diff --git a/tests/test_cache_handler.py b/tests/test_cache_handler.py index ddcac86f..f1facc0f 100644 --- a/tests/test_cache_handler.py +++ b/tests/test_cache_handler.py @@ -9,7 +9,7 @@ from pathlib import Path import pandas as pd from loguru import logger -from reme_ai.core.utils.cache_handler import CacheHandler +from reme.core.utils.cache_handler import CacheHandler def run_tests(): diff --git a/tests/test_embedding.py b/tests/test_embedding.py index d1c404f5..867037f4 100644 --- a/tests/test_embedding.py +++ b/tests/test_embedding.py @@ -14,16 +14,16 @@ Usage: # flake8: noqa: E402 # pylint: disable=C0413 -import asyncio import argparse +import asyncio from typing import Type, List -from reme_ai.core.utils import load_env +from reme.core.utils import load_env load_env() -from reme_ai.core.embedding import OpenAIEmbeddingModel, BaseEmbeddingModel -from reme_ai.core.schema import VectorNode +from reme.core.embedding import OpenAIEmbeddingModel, BaseEmbeddingModel +from reme.core.schema import VectorNode def get_embedding_model(model_class: Type[BaseEmbeddingModel]) -> BaseEmbeddingModel: diff --git a/tests/test_embedding_sync.py b/tests/test_embedding_sync.py index 361a42b3..34ddd350 100644 --- a/tests/test_embedding_sync.py +++ b/tests/test_embedding_sync.py @@ -17,12 +17,12 @@ Usage: import argparse from typing import Type, List -from reme_ai.core.utils import load_env +from reme.core.utils import load_env load_env() -from reme_ai.core.embedding import OpenAIEmbeddingModelSync, BaseEmbeddingModel -from reme_ai.core.schema import VectorNode +from reme.core.embedding import OpenAIEmbeddingModelSync, BaseEmbeddingModel +from reme.core.schema import VectorNode def get_embedding_model(model_class: Type[BaseEmbeddingModel]) -> BaseEmbeddingModel: diff --git a/tests/test_llm.py b/tests/test_llm.py index 12c6eca3..5e1755fe 100644 --- a/tests/test_llm.py +++ b/tests/test_llm.py @@ -14,17 +14,17 @@ Usage: # flake8: noqa: E402 # pylint: disable=C0413 -import asyncio import argparse +import asyncio from typing import Type -from reme_ai.core.utils import load_env +from reme.core.utils import load_env load_env() -from reme_ai.core.llm import OpenAILLM, LiteLLM, BaseLLM -from reme_ai.core.schema import Message, ToolCall -from reme_ai.core.enumeration import Role, ChunkEnum +from reme.core.llm import OpenAILLM, LiteLLM, BaseLLM +from reme.core.schema import Message, ToolCall +from reme.core.enumeration import Role, ChunkEnum def get_llm(llm_class: Type[BaseLLM]) -> BaseLLM: diff --git a/tests/test_llm_sync.py b/tests/test_llm_sync.py index 98751f87..0b3c8cce 100644 --- a/tests/test_llm_sync.py +++ b/tests/test_llm_sync.py @@ -17,13 +17,13 @@ Usage: import argparse from typing import Type -from reme_ai.core.utils import load_env +from reme.core.utils import load_env load_env() -from reme_ai.core.llm import OpenAILLMSync, LiteLLMSync, BaseLLM -from reme_ai.core.schema import Message, ToolCall -from reme_ai.core.enumeration import Role, ChunkEnum +from reme.core.llm import OpenAILLMSync, LiteLLMSync, BaseLLM +from reme.core.schema import Message, ToolCall +from reme.core.enumeration import Role, ChunkEnum def get_llm(llm_class: Type[BaseLLM]) -> BaseLLM: diff --git a/tests/test_logo.py b/tests/test_logo.py index eeede81e..12b757c4 100644 --- a/tests/test_logo.py +++ b/tests/test_logo.py @@ -1,9 +1,9 @@ """test logo""" -from reme_ai.core.schema import ServiceConfig, MCPConfig +from reme.core.schema import ServiceConfig, MCPConfig if __name__ == "__main__": - from reme_ai.core.utils import print_logo + from reme.core.utils import print_logo c = ServiceConfig(app_name="reme", backend="mcp", mcp=MCPConfig(transport="sse")) print_logo(service_config=c) diff --git a/tests/test_mcp_client.py b/tests/test_mcp_client.py index d2fef40e..a369d7cf 100644 --- a/tests/test_mcp_client.py +++ b/tests/test_mcp_client.py @@ -5,7 +5,7 @@ import asyncio import json -from reme_ai.core.utils import MCPClient +from reme.core.utils import MCPClient async def main(): diff --git a/tests/test_mcp_server.py b/tests/test_mcp_server.py index 67f9c542..f5744f99 100644 --- a/tests/test_mcp_server.py +++ b/tests/test_mcp_server.py @@ -5,8 +5,8 @@ from typing import Any from fastmcp import FastMCP from fastmcp.tools import FunctionTool -from reme_ai.core.schema import ToolCall -from reme_ai.core.utils import create_pydantic_model +from reme.core.schema import ToolCall +from reme.core.utils import create_pydantic_model mcp = FastMCP("DynamicSchemaServer", port=8010) diff --git a/tests/test_message.py b/tests/test_message.py index 141174c5..7b7e44bc 100644 --- a/tests/test_message.py +++ b/tests/test_message.py @@ -4,8 +4,8 @@ import unittest from mcp.types import Tool -from reme_ai.core.enumeration import Role -from reme_ai.core.schema import ToolAttr, ToolCall, ContentBlock, Message +from reme.core.enumeration import Role +from reme.core.schema import ToolAttr, ToolCall, ContentBlock, Message class TestModelDefinitions(unittest.TestCase): diff --git a/tests/test_op_composition.py b/tests/test_op_composition.py deleted file mode 100644 index 8d32b51a..00000000 --- a/tests/test_op_composition.py +++ /dev/null @@ -1,325 +0,0 @@ -""" -Unit tests for BaseOp and operator composition (>>, <<, |). -Tests asynchronous execution mode. -""" - -import asyncio - -from reme_ai.core.op import BaseOp -from reme_ai.core.schema import ToolCall, ToolAttr - - -class AddOp(BaseOp): - """Simple operator that adds a value to a number in context.""" - - def __init__(self, value: int = 1, **kwargs): - super().__init__(**kwargs) - self.value = value - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "name": self.name, - "description": f"Add {self.value} to input", - "parameters": ToolAttr( - **{ - "type": "object", - "properties": { - "number": {"type": "integer", "description": "Input number"}, - }, - "required": ["number"], - }, - ), - }, - ) - - async def execute(self): - """Async execution: add value to input number.""" - self.context["number"] += self.value - self.output = self.context["number"] - - -class MultiplyOp(BaseOp): - """Simple operator that multiplies a number in context.""" - - def __init__(self, factor: int = 2, **kwargs): - super().__init__(**kwargs) - self.factor = factor - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "name": self.name, - "description": f"Multiply by {self.factor}", - "parameters": ToolAttr( - **{ - "type": "object", - "properties": { - "number": {"type": "integer", "description": "Input number"}, - }, - "required": ["number"], - }, - ), - }, - ) - - async def execute(self): - """Async execution: multiply input number.""" - self.context["number"] *= self.factor - self.output = self.context["number"] - - -class AppendOp(BaseOp): - """Operator that appends a value to a list in context.""" - - def __init__(self, value: str = "", **kwargs): - super().__init__(**kwargs) - self.value = value - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "name": self.name, - "description": f"Append {self.value} to list", - "parameters": ToolAttr( - **{ - "type": "object", - "properties": { - "items": {"type": "array", "description": "List of items"}, - }, - "required": ["items"], - }, - ), - }, - ) - - async def execute(self): - """Async execution: append value to list.""" - self.context["items"].append(self.value) - self.output = self.context["items"] - - -async def test_basic_async_call(): - """Test basic asynchronous operator execution.""" - op = AddOp(value=5, name="add_5") - await op.call(number=10) - number = op.context["number"] - assert number == 15, f"Expected context result 15, got {number}" - print("✓ test_basic_async_call passed") - - -async def test_sequential_composition_async(): - """Test >> operator for sequential composition in async mode.""" - add_op = AddOp(value=5, name="add_5") - multiply_op = MultiplyOp(factor=2, name="multiply_2") - composed = add_op >> multiply_op - await composed.call(number=10) - - # (10 + 5) * 2 = 30 - assert composed.context["number"] == 30, f"Expected 30, got {composed.context['number']}" - print("✓ test_sequential_composition_async passed") - - -async def test_parallel_composition_async(): - """Test | operator for parallel composition in async mode.""" - append_a = AppendOp(value="A", name="append_a") - append_b = AppendOp(value="B", name="append_b") - append_c = AppendOp(value="C", name="append_c") - - composed = append_a | append_b | append_c - - await composed.call(items=[]) - - # All should append to the list - items = composed.context["items"] - assert len(items) == 3, f"Expected 3 items, got {len(items)}" - assert set(items) == {"A", "B", "C"}, f"Expected A,B,C, got {items}" - print("✓ test_parallel_composition_async passed") - - -async def test_add_sub_ops_async(): - """Test << operator for adding sub-operations in async mode.""" - parent_op = BaseOp(name="parent") - child1 = AddOp(value=5, name="child1") - child2 = MultiplyOp(factor=2, name="child2") - - _ = parent_op << child1 - _ = parent_op << child2 - - assert len(parent_op.sub_ops) == 2, f"Expected 2 sub_ops, got {len(parent_op.sub_ops)}" - sub_op_names = [op.name for op in parent_op.sub_ops] - assert "child1" in sub_op_names, "child1 not in sub_ops" - assert "child2" in sub_op_names, "child2 not in sub_ops" - print("✓ test_add_sub_ops_async passed") - - -async def test_add_sub_ops_dict(): - """Test << operator with dictionary of operations.""" - parent_op = BaseOp(name="parent") - ops_dict = { - "add": AddOp(value=5, name="add"), - "multiply": MultiplyOp(factor=2, name="multiply"), - } - - _ = parent_op << ops_dict - - assert len(parent_op.sub_ops) == 2, f"Expected 2 ops_dict, got {len(parent_op.sub_ops)}" - sub_op_names = [op.name for op in parent_op.sub_ops] - assert "add" in sub_op_names, "add not in ops_dict" - assert "multiply" in sub_op_names, "multiply not in ops_dict" - print("✓ test_add_sub_ops_dict passed") - - -async def test_add_sub_ops_list(): - """Test << operator with list of operations.""" - parent_op = BaseOp(name="parent") - sub_ops = [ - AddOp(value=5, name="add"), - MultiplyOp(factor=2, name="multiply"), - ] - - _ = parent_op << sub_ops - - assert len(parent_op.sub_ops) == 2, f"Expected 2 sub_ops, got {len(parent_op.sub_ops)}" - sub_op_names = [op.name for op in parent_op.sub_ops] - assert "add" in sub_op_names, "add not in sub_ops" - assert "multiply" in sub_op_names, "multiply not in sub_ops" - print("✓ test_add_sub_ops_list passed") - - -async def test_mixed_composition_async(): - """Test mixing >> and | operators in async mode.""" - # (add_5 >> multiply_2) | (add_10 >> multiply_3) - seq1 = AddOp(value=5, name="add_5") >> MultiplyOp(factor=2, name="multiply_2") - seq2 = AddOp(value=10, name="add_10") >> MultiplyOp(factor=3, name="multiply_3") - - composed = seq1 | seq2 - - await composed.call(number=10) - - # Both sequences execute in parallel with shared context - # seq1: (10 + 5) * 2 = 30 - # seq2: (30 + 10) * 3 = 120 (builds on seq1's result due to shared context) - # The exact result depends on execution order and timing - # With current implementation, result is 120 - assert composed.context["number"] == 120, f"Expected 120, got {composed.context['number']}" - print("✓ test_mixed_composition_async passed") - - -async def test_op_copy(): - """Test operator copy functionality.""" - original = AddOp(value=5, name="original") - copy_op = original.copy(name="copy") - - assert copy_op.name == "copy", f"Expected name 'copy', got {copy_op.name}" - assert copy_op.value == 5, f"Expected value 5, got {copy_op.value}" - assert copy_op is not original, "Copy should be a different object" - print("✓ test_op_copy passed") - - -async def test_input_mapping(): - """Test input_mapping parameter.""" - op = AddOp( - value=5, - name="add_5", - input_mapping={"x": "number"}, # Map x to number - ) - - await op.call(x=10) # Input is 'x' not 'number' - - assert op.context["number"] == 15, f"Expected number=15, got {op.context['number']}" - print("✓ test_input_mapping passed") - - -async def test_output_mapping(): - """Test output_mapping parameter.""" - op = AddOp( - value=5, - name="add_5", - output_mapping={"number": "final_result"}, # Map number to final_result - ) - - await op.call(number=10) - - assert op.context["number"] == 15, f"Expected number=15, got {op.context['number']}" - assert op.context["final_result"] == 15, f"Expected final_result=15, got {op.context['final_result']}" - print("✓ test_output_mapping passed") - - -async def test_validation_missing_required(): - """Test that missing required inputs raise an error.""" - op = AddOp(value=5, name="add_5", raise_exception=True) - - try: - await op.call() # Missing 'number' field - assert False, "Should have raised ValueError for missing required input" - except ValueError as e: - assert "number" in str(e), f"Expected error about 'number', got: {e}" - print("✓ test_validation_missing_required passed") - - -async def test_max_retries(): - """Test max_retries parameter with failing operation.""" - - class FailingOp(BaseOp): - """An operation that always fails.""" - - def __init__(self, **kwargs): - super().__init__(**kwargs) - self.attempt_count = 0 - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "name": self.name, - "description": "Always fails", - "parameters": ToolAttr(**{"type": "object", "properties": {}}), - "output": ToolAttr( - **{ - "type": "object", - "properties": { - "result": ToolAttr(**{"type": "string", "description": "Result"}), - }, - }, - ), - }, - ) - - async def execute(self): - self.attempt_count += 1 - raise RuntimeError(f"Attempt {self.attempt_count} failed") - - op = FailingOp(max_retries=3, name="failing") - - await op.call() - - assert op.attempt_count == 3, f"Expected 3 attempts, got {op.attempt_count}" - print("✓ test_max_retries passed") - - -async def async_main(): - """Run all async tests.""" - await test_basic_async_call() - await test_sequential_composition_async() - await test_parallel_composition_async() - await test_add_sub_ops_async() - await test_add_sub_ops_dict() - await test_add_sub_ops_list() - await test_mixed_composition_async() - await test_op_copy() - await test_input_mapping() - await test_output_mapping() - await test_validation_missing_required() - await test_max_retries() - - -if __name__ == "__main__": - print("Running BaseOp composition tests...\n") - - # Async tests - print("=== Asynchronous Tests ===") - asyncio.run(async_main()) - - print("\n" + "=" * 50) - print("All tests passed! ✓") - print("=" * 50) diff --git a/tests/test_timer.py b/tests/test_timer.py index c9714e38..98b1d816 100644 --- a/tests/test_timer.py +++ b/tests/test_timer.py @@ -7,7 +7,7 @@ import time from loguru import logger -from reme_ai.core.utils import timer +from reme.core.utils import timer @timer diff --git a/tests/test_token_counter.py b/tests/test_token_counter.py index 3c44a298..2fd8bdfd 100644 --- a/tests/test_token_counter.py +++ b/tests/test_token_counter.py @@ -14,9 +14,9 @@ Usage: import argparse from typing import Type, List -from reme_ai.core.enumeration import Role -from reme_ai.core.schema import Message, ToolCall -from reme_ai.core.token_counter import BaseTokenCounter, OpenAITokenCounter, HFTokenCounter +from reme.core.enumeration import Role +from reme.core.schema import Message, ToolCall +from reme.core.token_counter import BaseTokenCounter, OpenAITokenCounter, HFTokenCounter def get_token_counter(counter_class: Type[BaseTokenCounter], **kwargs) -> BaseTokenCounter: diff --git a/tests/test_tool.py b/tests/test_tool.py index 9db5a3ed..8d1fb232 100644 --- a/tests/test_tool.py +++ b/tests/test_tool.py @@ -8,9 +8,9 @@ search tools (Dashscope, Mock, Tavily) and execution tools (Code, Shell). import asyncio -from reme_ai.reme import ReMe +from reme.reme_app import ReMeApp -ReMe() +app = ReMeApp() def test_search(): @@ -19,7 +19,7 @@ def test_search(): Tests DashscopeSearch, MockSearch, and TavilySearch operations with a sample query to verify they work correctly. """ - from reme_ai.tool.search import DashscopeSearch, MockSearch, TavilySearch + from reme.tool.search import DashscopeSearch, MockSearch, TavilySearch query = "今天杭州的天气如何?" @@ -32,8 +32,8 @@ def test_search(): print(f"Testing {op.__class__.__name__}") print("=" * 60) print(f"Query: {query}") - asyncio.run(op.call(query=query)) - print(f"Output:\n{op.output}") + output = asyncio.run(op.call(query=query, service_context=app.service_context)) + print(f"Output:\n{output}") def test_execute(): @@ -43,7 +43,7 @@ def test_execute(): including successful execution, syntax errors, runtime errors, and invalid commands to verify error handling. """ - from reme_ai.tool.execute import ExecuteCode, ExecuteShell + from reme.tool.gallery import ExecuteCode, ExecuteShell # Test ExecuteCode print("\n" + "=" * 60) @@ -53,8 +53,8 @@ def test_execute(): op = ExecuteCode() code_to_execute = "print('hello world')" print(f"Executing Python code: {code_to_execute}") - asyncio.run(op.call(code=code_to_execute)) - print(f"Output:\n{op.output}") + output = asyncio.run(op.call(code=code_to_execute)) + print(f"Output:\n{output}") # Test ExecuteCode with more complex code print("\n" + "=" * 60) @@ -64,8 +64,8 @@ def test_execute(): op = ExecuteCode() code_to_execute = "result = sum(range(1, 11))\nprint(f'Sum of 1-10: {result}')" print(f"Executing Python code:\n{code_to_execute}") - asyncio.run(op.call(code=code_to_execute)) - print(f"Output:\n{op.output}") + output = asyncio.run(op.call(code=code_to_execute)) + print(f"Output:\n{output}") # Test ExecuteShell print("\n" + "=" * 60) @@ -75,8 +75,8 @@ def test_execute(): op = ExecuteShell() command = "ls" print(f"Executing shell command: {command}") - asyncio.run(op.call(command=command)) - print(f"Output:\n{op.output}") + output = asyncio.run(op.call(command=command)) + print(f"Output:\n{output}") # Test ExecuteShell with echo print("\n" + "=" * 60) @@ -86,8 +86,8 @@ def test_execute(): op = ExecuteShell() command = "echo 'Hello from shell!'" print(f"Executing shell command: {command}") - asyncio.run(op.call(command=command)) - print(f"Output:\n{op.output}") + output = asyncio.run(op.call(command=command)) + print(f"Output:\n{output}") # Test ExecuteCode with error (syntax error) print("\n" + "=" * 60) @@ -97,8 +97,8 @@ def test_execute(): op = ExecuteCode() code_to_execute = "print('missing closing quote)" print(f"Executing Python code with syntax error:\n{code_to_execute}") - asyncio.run(op.call(code=code_to_execute)) - print(f"Output:\n{op.output}") + output = asyncio.run(op.call(code=code_to_execute)) + print(f"Output:\n{output}") # Test ExecuteCode with runtime error print("\n" + "=" * 60) @@ -108,8 +108,8 @@ def test_execute(): op = ExecuteCode() code_to_execute = "x = 1 / 0" print(f"Executing Python code with runtime error:\n{code_to_execute}") - asyncio.run(op.call(code=code_to_execute)) - print(f"Output:\n{op.output}") + output = asyncio.run(op.call(code=code_to_execute)) + print(f"Output:\n{output}") # Test ExecuteCode with undefined variable print("\n" + "=" * 60) @@ -119,8 +119,8 @@ def test_execute(): op = ExecuteCode() code_to_execute = "print(undefined_variable)" print(f"Executing Python code with undefined variable:\n{code_to_execute}") - asyncio.run(op.call(code=code_to_execute)) - print(f"Output:\n{op.output}") + output = asyncio.run(op.call(code=code_to_execute)) + print(f"Output:\n{output}") # Test ExecuteShell with invalid command print("\n" + "=" * 60) @@ -130,8 +130,8 @@ def test_execute(): op = ExecuteShell() command = "this_command_does_not_exist" print(f"Executing invalid shell command: {command}") - asyncio.run(op.call(command=command)) - print(f"Output:\n{op.output}") + output = asyncio.run(op.call(command=command)) + print(f"Output:\n{output}") # Test ExecuteShell with command that returns non-zero exit code print("\n" + "=" * 60) @@ -141,8 +141,8 @@ def test_execute(): op = ExecuteShell() command = "ls /nonexistent_directory_12345" print(f"Executing shell command that should fail: {command}") - asyncio.run(op.call(command=command)) - print(f"Output:\n{op.output}") + output = asyncio.run(op.call(command=command)) + print(f"Output:\n{output}") print("\n" + "=" * 60) print("All tests completed!") @@ -155,11 +155,11 @@ def test_simple_chat(): Tests the SimpleChat agent with a basic query to verify it can process and respond to user input. """ - from reme_ai.mem_agent.chat import SimpleChat + from reme.agent.chat import SimpleChat op = SimpleChat() - asyncio.run(op.call(query="你好")) - print(op.output) + output = asyncio.run(op.call(query="你好", service_context=app.service_context)) + print(output) async def test_stream_chat(): @@ -168,13 +168,13 @@ async def test_stream_chat(): Tests the StreamChat agent with a query to verify it can process and stream responses in real-time using async operations. """ - from reme_ai.mem_agent.chat import StreamChat - from reme_ai.core.utils import execute_stream_task - from reme_ai.core.context import RuntimeContext + from reme.agent.chat import StreamChat + from reme.core.utils import execute_stream_task + from reme.core.context import RuntimeContext from asyncio import Queue op = StreamChat() - context = RuntimeContext(query="你好,详细介绍一下自己", stream_queue=Queue()) + context = RuntimeContext(query="你好,详细介绍一下自己", stream_queue=Queue(), service_context=app.service_context) async def task(): await op.call(context) @@ -192,5 +192,5 @@ async def test_stream_chat(): if __name__ == "__main__": # test_search() # test_execute() - # test_simple_chat() + test_simple_chat() asyncio.run(test_stream_chat()) diff --git a/tests/test_tool_call.py b/tests/test_tool_call.py index 30c0a37e..7be3079e 100644 --- a/tests/test_tool_call.py +++ b/tests/test_tool_call.py @@ -2,7 +2,7 @@ import json -from reme_ai.core.schema.tool_call import ToolCall +from reme.core.schema.tool_call import ToolCall def test_simple_schema(): diff --git a/tests/test_vector_store.py b/tests/test_vector_store.py index 2f4c1f4d..14ed1c1d 100644 --- a/tests/test_vector_store.py +++ b/tests/test_vector_store.py @@ -18,14 +18,16 @@ Usage: import argparse import asyncio import shutil +from concurrent.futures import ThreadPoolExecutor from pathlib import Path from typing import List from loguru import logger -from reme_ai.core.embedding import OpenAIEmbeddingModel -from reme_ai.core.schema import VectorNode -from reme_ai.core.vector_store import ( +from reme.core.embedding import OpenAIEmbeddingModel +from reme.core.schema import VectorNode +from reme.core.utils import load_env +from reme.core.vector_store import ( BaseVectorStore, ChromaVectorStore, LocalVectorStore, @@ -34,6 +36,7 @@ from reme_ai.core.vector_store import ( QdrantVectorStore, ) +load_env() # ==================== Configuration ==================== @@ -199,7 +202,7 @@ def create_vector_store(store_type: str, collection_name: str) -> BaseVectorStor """Create a vector store instance based on type. Args: - store_type: Type of vector store ("local", "es", or "qdrant") + store_type: Type of vector store ("local", "es", "pgvector", "qdrant", or "chroma") collection_name: Name of the collection Returns: @@ -213,16 +216,21 @@ def create_vector_store(store_type: str, collection_name: str) -> BaseVectorStor dimensions=config.EMBEDDING_DIMENSIONS, ) + # Create thread pool executor for vector stores + thread_pool = ThreadPoolExecutor(max_workers=4) + if store_type == "local": return LocalVectorStore( collection_name=collection_name, embedding_model=embedding_model, + thread_pool=thread_pool, root_path=config.LOCAL_ROOT_PATH, ) elif store_type == "es": return ESVectorStore( collection_name=collection_name, embedding_model=embedding_model, + thread_pool=thread_pool, hosts=config.ES_HOSTS, basic_auth=config.ES_BASIC_AUTH, ) @@ -230,6 +238,7 @@ def create_vector_store(store_type: str, collection_name: str) -> BaseVectorStor return QdrantVectorStore( collection_name=collection_name, embedding_model=embedding_model, + thread_pool=thread_pool, path=config.QDRANT_PATH, host=config.QDRANT_HOST, port=config.QDRANT_PORT, @@ -242,6 +251,7 @@ def create_vector_store(store_type: str, collection_name: str) -> BaseVectorStor return PGVectorStore( collection_name=collection_name, embedding_model=embedding_model, + thread_pool=thread_pool, dsn=config.PG_DSN, min_size=config.PG_MIN_SIZE, max_size=config.PG_MAX_SIZE, @@ -252,6 +262,7 @@ def create_vector_store(store_type: str, collection_name: str) -> BaseVectorStor return ChromaVectorStore( collection_name=collection_name, embedding_model=embedding_model, + thread_pool=thread_pool, path=config.CHROMA_PATH, host=config.CHROMA_HOST, port=config.CHROMA_PORT, @@ -350,35 +361,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 +399,17 @@ 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 +804,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 +814,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 +1132,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 +1180,327 @@ 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 start_time <= 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 base_timestamp + 8000 <= ts <= base_timestamp + 12000, "Timestamp out of range" + assert 65 <= 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 60 <= 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 60 <= 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 + 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: + embedding_model = OpenAIEmbeddingModel() + thread_pool = ThreadPoolExecutor(max_workers=4) + + # This should raise ValueError due to invalid table name + try: + _ = PGVectorStore( + collection_name="test'; DROP TABLE users; --", + embedding_model=embedding_model, + thread_pool=thread_pool, + ) + 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: + _ = 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 +1669,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 +1690,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 ==========