From e6ad682ede8058389d9cde4f76cfb19c3b49d4b2 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sat, 17 Jan 2026 01:15:49 +0800 Subject: [PATCH 01/19] feat(core): add ReMe V3 implementation with optional MCP client and enhanced filtering --- bench/halumem/analyze_dataset_stats.py | 45 +- bench/halumem/eval_reme_simple_v3.py | 668 ++++++++++++++++++ reme_ai/core/context/prompt_handler.py | 2 +- reme_ai/core/schema/tool_call.py | 6 +- reme_ai/core/utils/__init__.py | 10 +- reme_ai/core/utils/cache_handler.py | 20 +- .../core/vector_store/chroma_vector_store.py | 46 +- reme_ai/core/vector_store/es_vector_store.py | 28 +- .../core/vector_store/local_vector_store.py | 21 +- reme_ai/core/vector_store/pgvector_store.py | 72 +- .../core/vector_store/qdrant_vector_store.py | 64 +- reme_ai/mem_agent/v3/__init__.py | 9 + .../mem_agent/v3/personal_summarizer_v3.py | 69 ++ .../mem_agent/v3/personal_summarizer_v3.yaml | 38 + reme_ai/mem_agent/v3/reme_retriever_v3.py | 44 ++ reme_ai/mem_agent/v3/reme_retriever_v3.yaml | 53 ++ reme_ai/mem_agent/v3/reme_summarizer_v3.py | 88 +++ reme_ai/mem_agent/v3/reme_summarizer_v3.yaml | 25 + reme_ai/mem_tool/base_memory_tool.py | 2 - reme_ai/mem_tool/read_local_memories.py | 54 ++ reme_ai/mem_tool/read_local_memories.yaml | 8 + reme_ai/mem_tool/v3/__init__.py | 15 + reme_ai/mem_tool/v3/add_memory.py | 67 ++ reme_ai/mem_tool/v3/read_history.py | 38 + reme_ai/mem_tool/v3/read_user_profile.py | 65 ++ reme_ai/mem_tool/v3/retrieve_memory.py | 84 +++ reme_ai/mem_tool/v3/summary_and_hands_off.py | 140 ++++ reme_ai/mem_tool/v3/update_user_profile.py | 118 ++++ reme_ai/mem_tool/write_local_memories.py | 57 ++ reme_ai/mem_tool/write_local_memories.yaml | 5 + reme_ai/reme.py | 97 ++- tests/test_vector_store.py | 417 +++++++++-- 32 files changed, 2369 insertions(+), 106 deletions(-) create mode 100644 bench/halumem/eval_reme_simple_v3.py create mode 100644 reme_ai/mem_agent/v3/__init__.py create mode 100644 reme_ai/mem_agent/v3/personal_summarizer_v3.py create mode 100644 reme_ai/mem_agent/v3/personal_summarizer_v3.yaml create mode 100644 reme_ai/mem_agent/v3/reme_retriever_v3.py create mode 100644 reme_ai/mem_agent/v3/reme_retriever_v3.yaml create mode 100644 reme_ai/mem_agent/v3/reme_summarizer_v3.py create mode 100644 reme_ai/mem_agent/v3/reme_summarizer_v3.yaml create mode 100644 reme_ai/mem_tool/read_local_memories.py create mode 100644 reme_ai/mem_tool/read_local_memories.yaml create mode 100644 reme_ai/mem_tool/v3/__init__.py create mode 100644 reme_ai/mem_tool/v3/add_memory.py create mode 100644 reme_ai/mem_tool/v3/read_history.py create mode 100644 reme_ai/mem_tool/v3/read_user_profile.py create mode 100644 reme_ai/mem_tool/v3/retrieve_memory.py create mode 100644 reme_ai/mem_tool/v3/summary_and_hands_off.py create mode 100644 reme_ai/mem_tool/v3/update_user_profile.py create mode 100644 reme_ai/mem_tool/write_local_memories.py create mode 100644 reme_ai/mem_tool/write_local_memories.yaml diff --git a/bench/halumem/analyze_dataset_stats.py b/bench/halumem/analyze_dataset_stats.py index 6695ad7f..d5a213b4 100644 --- a/bench/halumem/analyze_dataset_stats.py +++ b/bench/halumem/analyze_dataset_stats.py @@ -29,6 +29,7 @@ class UserStats: dialogues_per_session: list[int] # 每个 session 的对话数量 dialogue_lengths_per_session: list[int] # 每个 session 的对话总长度(字符数) num_chunks_after_split: int # 按 5000 字符分割后的 chunk 数量 + session_time_ranges: list[tuple[Any, Any]] # 每个 session 的 (开始时间, 结束时间) @dataclass @@ -169,6 +170,7 @@ class DatasetAnalyzer: dialogues_per_session = [] dialogue_lengths_per_session = [] + session_time_ranges = [] total_chunks = 0 for session in sessions: @@ -179,6 +181,11 @@ class DatasetAnalyzer: dialogues_per_session.append(num_dialogues) dialogue_lengths_per_session.append(dialogue_length) + # 收集 session 的时间范围 + start_time = session.get("start_time", None) + end_time = session.get("end_time", None) + session_time_ranges.append((start_time, end_time)) + # 计算这个 session 分割后的 chunk 数量 num_chunks = self.split_session_into_chunks(dialogue, max_length=5000) total_chunks += num_chunks @@ -202,7 +209,8 @@ class DatasetAnalyzer: num_sessions=len(sessions), dialogues_per_session=dialogues_per_session, dialogue_lengths_per_session=dialogue_lengths_per_session, - num_chunks_after_split=total_chunks + num_chunks_after_split=total_chunks, + session_time_ranges=session_time_ranges ) self.user_stats_list.append(user_stats) @@ -392,6 +400,32 @@ class DatasetAnalyzer: print(f" 平均每 Session 对话长度: {avg_length:.2f} 字符") print() + def print_first_user_session_times(self): + """打印第一个用户的每个 session 的时间范围""" + if not self.user_stats_list: + print("\n没有用户数据") + return + + first_user = self.user_stats_list[0] + + print("\n" + "=" * 80) + print(f"第一个用户的 Session 时间统计") + print("=" * 80 + "\n") + print(f"用户名: {first_user.user_name}") + print(f"UUID: {first_user.uuid}") + print(f"总 Session 数: {first_user.num_sessions}\n") + + print("-" * 80) + print(f"{'Session #':<12} {'开始时间':<30} {'结束时间':<30}") + print("-" * 80) + + for idx, (start_time, end_time) in enumerate(first_user.session_time_ranges, 1): + start_str = str(start_time) if start_time is not None else "无" + end_str = str(end_time) if end_time is not None else "无" + print(f"{idx:<12} {start_str:<30} {end_str:<30}") + + print("=" * 80) + def print_user_split_summary(self): """打印每个用户的分割统计摘要(表格形式)""" print("\n" + "=" * 80) @@ -469,7 +503,11 @@ class DatasetAnalyzer: if u.dialogue_lengths_per_session else 0 ), "dialogues_per_session": u.dialogues_per_session, - "dialogue_lengths_per_session": u.dialogue_lengths_per_session + "dialogue_lengths_per_session": u.dialogue_lengths_per_session, + "session_time_ranges": [ + {"start_time": start, "end_time": end} + for start, end in u.session_time_ranges + ] } for u in self.user_stats_list ] @@ -498,6 +536,9 @@ def main(data_path: str, output_path: str = None, show_per_user: bool = False): # 打印摘要 analyzer.print_summary(stats) + # 打印第一个用户的 session 时间统计 + analyzer.print_first_user_session_times() + # 打印每个用户的分割统计摘要(始终显示) analyzer.print_user_split_summary() diff --git a/bench/halumem/eval_reme_simple_v3.py b/bench/halumem/eval_reme_simple_v3.py new file mode 100644 index 00000000..62b1fe64 --- /dev/null +++ b/bench/halumem/eval_reme_simple_v3.py @@ -0,0 +1,668 @@ +""" +HaluMem Benchmark Evaluator for ReMe V3 - Question Answering + +A modular evaluation pipeline that: +1. Loads HaluMem benchmark data +2. Processes user sessions through ReMe V3 (summarization + retrieval) +3. Evaluates question answering performance +4. Generates comprehensive metrics + +Usage: + python bench/halumem/eval_reme_simple_v3.py \ + --data_path /Users/yuli/workspace/HaluMem/data/HaluMem-Medium.jsonl \ + --top_k 20 --user_num 100 --max_concurrency 20 +""" + +import asyncio +import json +import os +import re +import time +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +from loguru import logger + +from eval_tools import evaluation_for_question2 +from reme_ai.core.enumeration import MemoryType +from reme_ai.core.schema import MemoryNode +from reme_ai.reme import ReMe + + +# ==================== Configuration ==================== + +@dataclass +class EvalConfig: + """Evaluation configuration parameters.""" + data_path: str + top_k: int = 20 + user_num: int = 1 + max_concurrency: int = 2 + batch_size: int = 20 + output_dir: str = "bench_results/reme_simple_v3" + + +# ==================== Utilities ==================== + +class DataLoader: + """Handles loading and parsing of HaluMem data.""" + + @staticmethod + def load_jsonl(file_path: str) -> list[dict]: + """Load all entries from a JSONL file.""" + with open(file_path, "r", encoding="utf-8") as f: + return [json.loads(line.strip()) for line in f if line.strip()] + + @staticmethod + def extract_user_name(persona_info: str) -> str: + """Extract user name from persona info string.""" + match = re.search(r"Name:\s*(.*?); Gender:", persona_info) + if not match: + raise ValueError(f"No name found in persona_info: {persona_info}") + return match.group(1).strip() + + @staticmethod + def format_dialogue_messages(dialogue: list[dict]) -> list[dict]: + """Format dialogue into ReMe message format with conversation_time (user messages only).""" + return [ + { + "role": turn["role"], + "content": turn["content"], + "time_created": datetime.strptime( + turn["timestamp"], "%b %d, %Y, %H:%M:%S" + ) + .replace(tzinfo=timezone.utc) + .strftime("%Y-%m-%d %H:%M:%S"), + } + for turn in dialogue + if turn["role"] == "user" # Only include user messages + ] + + @staticmethod + def format_dialogue_for_eval(dialogue: list[dict], user_name: str = None) -> str: + """Format dialogue into string for evaluation.""" + formatted_turns = [] + for turn in dialogue: + timestamp = datetime.strptime( + turn["timestamp"], "%b %d, %Y, %H:%M:%S" + ).replace(tzinfo=timezone.utc).strftime("%Y-%m-%d %H:%M:%S") + + # Use user_name if role is 'user' and user_name is provided + role = user_name if turn['role'] == 'user' and user_name else turn['role'] + + formatted_turns.append( + f"Role: {role}\n" + f"Content: {turn['content']}\n" + f"Time: {timestamp}" + ) + return "\n\n".join(formatted_turns) + + +class FileManager: + """Manages file I/O operations.""" + + def __init__(self, base_dir: str): + self.base_dir = Path(base_dir) + self.tmp_dir = self.base_dir / "tmp" + self.tmp_dir.mkdir(parents=True, exist_ok=True) + + def get_user_dir(self, user_name: str) -> Path: + """Get the directory path for a user.""" + user_dir = self.tmp_dir / user_name + user_dir.mkdir(parents=True, exist_ok=True) + return user_dir + + def get_session_file(self, user_name: str, session_id: int) -> Path: + """Get the file path for a specific session.""" + return self.get_user_dir(user_name) / f"session_{session_id}.json" + + def save_session(self, user_name: str, session_id: int, data: dict): + """Save session data to file.""" + file_path = self.get_session_file(user_name, session_id) + with open(file_path, "w", encoding="utf-8") as f: + json.dump(data, f, ensure_ascii=False, indent=2) + logger.info(f"✅ Saved session {session_id} to {file_path}") + + def load_session(self, user_name: str, session_id: int) -> dict | None: + """Load session data from file.""" + file_path = self.get_session_file(user_name, session_id) + if not file_path.exists(): + return None + with open(file_path, "r", encoding="utf-8") as f: + return json.load(f) + + def user_has_cache(self, user_name: str) -> bool: + """Check if user has cached results.""" + user_dir = self.get_user_dir(user_name) + return any(f.name.startswith("session_") and f.suffix == ".json" + for f in user_dir.iterdir()) + + def combine_results(self, output_file: str): + """Combine all user session files into a single JSONL file.""" + with open(output_file, "w", encoding="utf-8") as f_out: + for user_dir in self.tmp_dir.iterdir(): + if not user_dir.is_dir(): + continue + + session_files = sorted([ + f for f in user_dir.iterdir() + if f.name.startswith("session_") and f.suffix == ".json" + ]) + + if not session_files: + continue + + # Load first session to get user metadata + with open(session_files[0], "r", encoding="utf-8") as f_in: + first_session = json.load(f_in) + + user_data = { + "uuid": first_session["uuid"], + "user_name": first_session["user_name"], + "sessions": [] + } + + # Load all sessions + for session_file in session_files: + with open(session_file, "r", encoding="utf-8") as f_in: + session_data = json.load(f_in) + # Remove redundant user metadata + session_data.pop("uuid", None) + session_data.pop("user_name", None) + user_data["sessions"].append(session_data) + + f_out.write(json.dumps(user_data, ensure_ascii=False) + "\n") + + +# ==================== Memory Operations ==================== + +class MemoryProcessor: + """Handles ReMe V3 memory operations.""" + + def __init__(self, reme: ReMe): + self.reme = reme + + async def add_memories( + self, + user_id: str, + messages: list[dict], + batch_size: int = 10000 + ) -> tuple[list[str], list[list[dict]], float]: + """ + Add memories in batches using ReMe V3 and return extracted memory contents. + + Returns: + tuple: (extracted_memories, agent_messages, total_duration_ms) + """ + added_memories: list[MemoryNode] = [] + deleted_memories: list[str] = [] + all_agent_messages: list = [] + total_duration_ms = 0 + + for i in range(0, len(messages), batch_size): + batch = messages[i:i + batch_size] + start = time.time() + + # Use summary_v3 instead of summary_v2 + memory_nodes, agent_messages, success = await self.reme.summary_v3( + messages=batch, + user_id=user_id + ) + + duration_ms = (time.time() - start) * 1000 + total_duration_ms += duration_ms + + # Save agent messages for this batch + if agent_messages: + all_agent_messages.extend(agent_messages) + + if memory_nodes: + for node in memory_nodes: + if isinstance(node, MemoryNode) and node.memory_type == MemoryType.HISTORY: + continue + + if isinstance(node, MemoryNode): + added_memories.append(node) + + if isinstance(node, str): + deleted_memories.append(node) + + extracted_memories = deleted_memories + extracted_memories += ["[delete]" + n.format_memory() for n in added_memories if n.memory_id in deleted_memories] + extracted_memories += ["[add]" + n.format_memory() for n in added_memories if n.memory_id not in deleted_memories] + return extracted_memories, all_agent_messages, total_duration_ms + + async def search_memory( + self, + query: str, + user_id: str, + top_k: int = 20 + ) -> tuple[str, list, float]: + """ + Search memory using ReMe V3 and return response. + + Returns: + tuple: (response, agent_messages, duration_ms) + """ + start = time.time() + + # Use retrieve_v3 instead of retrieve_v2 + response, agent_messages, success = await self.reme.retrieve_v3( + query=query, + user_id=user_id, + top_k=top_k + ) + + duration_ms = (time.time() - start) * 1000 + return response, agent_messages, duration_ms + + +# ==================== Evaluation ==================== + +class QuestionAnsweringEvaluator: + """Evaluates question answering performance.""" + + def __init__(self, memory_processor: MemoryProcessor, top_k: int): + self.memory_processor = memory_processor + self.top_k = top_k + + async def evaluate_questions( + self, + questions: list[dict], + user_name: str, + uuid: str, + session_id: int, + formatted_dialogue: str + ) -> list[dict]: + """Evaluate all questions for a session.""" + results = [] + + for qa in questions: + # Search memory for answer using V3 + response, agent_messages, duration_ms = await self.memory_processor.search_memory( + query=qa["question"], + user_id=user_name, + top_k=self.top_k + ) + + # Evaluate response + evidence_text = "\n".join([e["memory_content"] for e in qa["evidence"]]) + eval_result = await evaluation_for_question2( + qa["question"], + qa["answer"], + evidence_text, + response, + formatted_dialogue + ) + + # Build result record + qa_result = { + **qa, + "uuid": uuid, + "session_id": session_id, + "system_response": response, + "retrieve_messages": [m.model_dump() for m in agent_messages], + "search_duration_ms": duration_ms, + "result_type": eval_result.get("evaluation_result"), + "question_answering_reasoning": eval_result.get("reasoning", "") + } + results.append(qa_result) + + return results + + +class MetricsAggregator: + """Aggregates evaluation metrics.""" + + @staticmethod + def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: + """Compute question answering metrics.""" + total = len(qa_records) + if total == 0: + return { + "correct_qa_ratio(all)": 0, + "hallucination_qa_ratio(all)": 0, + "omission_qa_ratio(all)": 0, + "correct_qa_ratio(valid)": 0, + "hallucination_qa_ratio(valid)": 0, + "omission_qa_ratio(valid)": 0, + "qa_valid_num": 0, + "qa_num": 0 + } + + correct = 0 + hallucination = 0 + omission = 0 + valid = 0 + + for qa in qa_records: + result_type = qa.get("result_type", "") + + if result_type in ["Correct", "Hallucination", "Omission"]: + valid += 1 + if result_type == "Correct": + correct += 1 + elif result_type == "Hallucination": + hallucination += 1 + elif result_type == "Omission": + omission += 1 + + metrics = { + "correct_qa_ratio(all)": correct / total, + "hallucination_qa_ratio(all)": hallucination / total, + "omission_qa_ratio(all)": omission / total, + "qa_valid_num": valid, + "qa_num": total + } + + if valid > 0: + metrics.update({ + "correct_qa_ratio(valid)": correct / valid, + "hallucination_qa_ratio(valid)": hallucination / valid, + "omission_qa_ratio(valid)": omission / valid + }) + else: + metrics.update({ + "correct_qa_ratio(valid)": 0, + "hallucination_qa_ratio(valid)": 0, + "omission_qa_ratio(valid)": 0 + }) + + return metrics + + @staticmethod + def compute_time_metrics(eval_results_file: str) -> dict[str, float]: + """Compute timing metrics from evaluation results.""" + add_duration = 0 + search_duration = 0 + + with open(eval_results_file, "r", encoding="utf-8") as f: + for line in f: + if not line.strip(): + continue + user_data = json.loads(line) + + for session in user_data["sessions"]: + add_duration += session.get("add_dialogue_duration_ms", 0) + + eval_results = session.get("evaluation_results", {}) + for qa in eval_results.get("question_answering_records", []): + search_duration += qa.get("search_duration_ms", 0) + + # Convert to minutes + return { + "add_dialogue_duration_time": add_duration / 1000 / 60, + "search_memory_duration_time": search_duration / 1000 / 60, + "total_duration_time": (add_duration + search_duration) / 1000 / 60 + } + + +# ==================== Main Pipeline ==================== + +class HaluMemEvaluatorV3: + """Main evaluator orchestrating the entire ReMe V3 pipeline.""" + + def __init__(self, config: EvalConfig): + self.config = config + self.reme = ReMe() + self.file_manager = FileManager(config.output_dir) + self.memory_processor = MemoryProcessor(self.reme) + self.qa_evaluator = QuestionAnsweringEvaluator( + self.memory_processor, + config.top_k + ) + self.data_loader = DataLoader() + + async def process_session( + self, + session: dict, + session_id: int, + user_name: str, + uuid: str + ) -> dict: + """Process a single session using ReMe V3.""" + session_data = { + "uuid": uuid, + "user_name": user_name, + "session_id": session_id, + "memory_points": session["memory_points"] + } + + # Skip generated QA sessions + if session.get("is_generated_qa_session", False): + session_data["is_generated_qa_session"] = True + return session_data + + # Format and add dialogue to memory using V3 + dialogue = session["dialogue"] + formatted_messages = self.data_loader.format_dialogue_messages(dialogue) + + extracted_memories, agent_messages, duration_ms = await self.memory_processor.add_memories( + user_id=user_name, + messages=formatted_messages, + batch_size=self.config.batch_size + ) + + session_data.update({ + "dialogue": dialogue, + "extracted_memories": extracted_memories, + "summary_messages": [m.model_dump() for m in agent_messages], + "add_dialogue_duration_ms": duration_ms + }) + + # Evaluate questions if present + if "questions" in session: + formatted_dialogue = self.data_loader.format_dialogue_for_eval(dialogue, user_name) + qa_results = await self.qa_evaluator.evaluate_questions( + questions=session["questions"], + user_name=user_name, + uuid=uuid, + session_id=session_id, + formatted_dialogue=formatted_dialogue + ) + + session_data["evaluation_results"] = { + "question_answering_records": qa_results + } + + return session_data + + async def process_user(self, user_data: dict) -> dict: + """Process all sessions for a user.""" + user_name = self.data_loader.extract_user_name(user_data["persona_info"]) + uuid = user_data["uuid"] + + logger.info(f"Processing user: {user_name}") + + for idx, session in enumerate(user_data["sessions"]): + logger.info(f" Session {idx + 1}/{len(user_data['sessions'])}") + + session_data = await self.process_session( + session=session, + session_id=idx, + user_name=user_name, + uuid=uuid + ) + + self.file_manager.save_session(user_name, idx, session_data) + + return {"uuid": uuid, "user_name": user_name, "status": "ok"} + + async def run_evaluation(self): + """Run the complete evaluation pipeline using ReMe V3.""" + start_time = time.time() + + # Clear existing data + await self.reme.vector_store.delete_all() + + # Load user data + all_users = self.data_loader.load_jsonl(self.config.data_path) + users_to_process = all_users[:self.config.user_num] + + print("\n" + "=" * 80) + print("HALUMEM EVALUATION - REME V3 - QUESTION ANSWERING") + print(f"Users: {len(users_to_process)} | Concurrency: {self.config.max_concurrency}") + print("=" * 80 + "\n") + + # Process users with concurrency control + semaphore = asyncio.Semaphore(self.config.max_concurrency) + + async def process_with_cache_check(idx: int, user_data: dict): + async with semaphore: + user_name = self.data_loader.extract_user_name(user_data["persona_info"]) + + # Check cache + if self.file_manager.user_has_cache(user_name): + print(f"⚡ [{idx}/{len(users_to_process)}] Skipping {user_name} (cached)") + return {"user_name": user_name, "status": "cached"} + + print(f"🔄 [{idx}/{len(users_to_process)}] Processing {user_name}...") + result = await self.process_user(user_data) + print(f"✅ [{idx}/{len(users_to_process)}] Completed {user_name}") + return result + + tasks = [ + process_with_cache_check(idx, user) + for idx, user in enumerate(users_to_process, 1) + ] + await asyncio.gather(*tasks) + + # Combine results + output_file = os.path.join(self.config.output_dir, "eval_results.jsonl") + self.file_manager.combine_results(output_file) + + elapsed = time.time() - start_time + print(f"\n✅ Processing completed in {elapsed:.2f}s") + print(f"📁 Results: {output_file}\n") + + # Aggregate metrics + await self.aggregate_and_report(output_file) + + async def aggregate_and_report(self, results_file: str): + """Aggregate results and generate final report.""" + print("=" * 80) + print("AGGREGATING METRICS") + print("=" * 80 + "\n") + + # Collect all QA records + qa_records = [] + with open(results_file, "r", encoding="utf-8") as f: + for line in f: + if not line.strip(): + continue + user_data = json.loads(line) + + for session in user_data["sessions"]: + if session.get("is_generated_qa_session"): + continue + + eval_results = session.get("evaluation_results", {}) + qa_records.extend( + eval_results.get("question_answering_records", []) + ) + + # Compute metrics + qa_metrics = MetricsAggregator.compute_qa_metrics(qa_records) + time_metrics = MetricsAggregator.compute_time_metrics(results_file) + + final_results = { + "overall_score": { + "question_answering": qa_metrics, + "time_consuming": time_metrics + }, + "question_answering_records": qa_records + } + + # Save final report + report_file = os.path.join(self.config.output_dir, "eval_statistics.json") + with open(report_file, "w", encoding="utf-8") as f: + json.dump(final_results, f, ensure_ascii=False, indent=4) + + print(f"📊 Statistics saved to: {report_file}\n") + + # Print summary + self._print_summary(qa_metrics, time_metrics) + + def _print_summary(self, qa_metrics: dict, time_metrics: dict): + """Print evaluation summary.""" + print("=" * 80) + print("EVALUATION SUMMARY - REME V3") + print("=" * 80 + "\n") + + print("📊 Question Answering:") + print(f" Correct (all): {qa_metrics['correct_qa_ratio(all)']:.4f}") + print(f" Hallucination (all): {qa_metrics['hallucination_qa_ratio(all)']:.4f}") + print(f" Omission (all): {qa_metrics['omission_qa_ratio(all)']:.4f}") + print(f" Correct (valid): {qa_metrics['correct_qa_ratio(valid)']:.4f}") + print(f" Hallucination (valid): {qa_metrics['hallucination_qa_ratio(valid)']:.4f}") + print(f" Omission (valid): {qa_metrics['omission_qa_ratio(valid)']:.4f}") + print(f" Valid/Total: {qa_metrics['qa_valid_num']}/{qa_metrics['qa_num']}") + + print(f"\n⏱️ Time Metrics:") + print(f" Memory Addition: {time_metrics['add_dialogue_duration_time']:.2f} min") + print(f" Memory Search: {time_metrics['search_memory_duration_time']:.2f} min") + print(f" Total: {time_metrics['total_duration_time']:.2f} min") + print("\n" + "=" * 80) + + +# ==================== Entry Point ==================== + +def main( + data_path: str, + top_k: int = 20, + user_num: int = 1, + max_concurrency: int = 2 +): + """Main entry point for ReMe V3 evaluation.""" + config = EvalConfig( + data_path=data_path, + top_k=top_k, + user_num=user_num, + max_concurrency=max_concurrency + ) + + evaluator = HaluMemEvaluatorV3(config) + asyncio.run(evaluator.run_evaluation()) + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser( + description="Evaluate ReMe V3 on HaluMem benchmark (Question Answering)" + ) + parser.add_argument( + "--data_path", + type=str, + required=True, + help="Path to HaluMem JSONL file" + ) + parser.add_argument( + "--top_k", + type=int, + default=20, + help="Number of memories to retrieve (default: 20)" + ) + parser.add_argument( + "--user_num", + type=int, + default=1, + help="Number of users to evaluate (default: 1)" + ) + parser.add_argument( + "--max_concurrency", + type=int, + default=2, + help="Maximum concurrent user processing (default: 2)" + ) + + args = parser.parse_args() + + main( + data_path=args.data_path, + top_k=args.top_k, + user_num=args.user_num, + max_concurrency=args.max_concurrency + ) diff --git a/reme_ai/core/context/prompt_handler.py b/reme_ai/core/context/prompt_handler.py index e48f98ac..b428b163 100644 --- a/reme_ai/core/context/prompt_handler.py +++ b/reme_ai/core/context/prompt_handler.py @@ -55,7 +55,7 @@ class PromptHandler(BaseContext): key += "_" + self.language.strip() assert key in self, f"prompt_name={key} not found." - return self[key] + return self[key].strip() def prompt_format(self, prompt_name: str, **kwargs) -> str: """Format a prompt by filtering flagged lines and filling template variables.""" diff --git a/reme_ai/core/schema/tool_call.py b/reme_ai/core/schema/tool_call.py index 0a96f8c4..035c63f1 100644 --- a/reme_ai/core/schema/tool_call.py +++ b/reme_ai/core/schema/tool_call.py @@ -41,14 +41,14 @@ class ToolAttr(BaseModel): if self.enum: res["enum"] = self.enum - if self.type == "object" and self.properties: + if self.type == "object" and self.properties is not None: res["properties"] = { k: v.simple_input_dump() if isinstance(v, ToolAttr) else v for k, v in self.properties.items() } - if self.required: + if self.required is not None: res["required"] = self.required - if self.type == "array" and self.items: + if self.type == "array" and self.items is not None: res["items"] = self.items.simple_input_dump() if isinstance(self.items, ToolAttr) else self.items return res diff --git a/reme_ai/core/utils/__init__.py b/reme_ai/core/utils/__init__.py index 242f7be9..23396f97 100644 --- a/reme_ai/core/utils/__init__.py +++ b/reme_ai/core/utils/__init__.py @@ -9,7 +9,15 @@ from .http_client import HttpClient from .llm_utils import extract_content, format_messages, deduplicate_memories from .logger_utils import init_logger from .logo_utils import print_logo -from .mcp_client import MCPClient + +# Make MCPClient import optional to avoid breaking if MCP dependencies are not available +try: + from .mcp_client import MCPClient + _HAS_MCP = True +except ImportError: + MCPClient = None + _HAS_MCP = False + from .pydantic_config_parser import PydanticConfigParser from .pydantic_utils import create_pydantic_model from .singleton import singleton diff --git a/reme_ai/core/utils/cache_handler.py b/reme_ai/core/utils/cache_handler.py index 70c8585a..f3b0072f 100644 --- a/reme_ai/core/utils/cache_handler.py +++ b/reme_ai/core/utils/cache_handler.py @@ -15,7 +15,7 @@ class CacheHandler: _EXTENSIONS = { pd.DataFrame: ".csv", dict: ".json", - list: ".json", + list: ".jsonl", str: ".txt", } @@ -76,11 +76,17 @@ class CacheHandler: data.to_csv(path, index=kwargs.get("index", False), encoding="utf-8") return {"row_count": len(data), "file_size": path.stat().st_size} - if dtype in (dict, list): + if dtype is dict: with open(path, "w", encoding="utf-8") as f: json.dump(data, f, ensure_ascii=False, indent=2) return {"item_count": len(data), "file_size": path.stat().st_size} + if dtype is list: + with open(path, "w", encoding="utf-8") as f: + for item in data: + f.write(json.dumps(item, ensure_ascii=False) + "\n") + return {"item_count": len(data), "file_size": path.stat().st_size} + if dtype is str: path.write_text(data, encoding=kwargs.get("encoding", "utf-8")) return {"char_count": len(data), "file_size": path.stat().st_size} @@ -92,9 +98,17 @@ class CacheHandler: """Execute type-specific load operations.""" if type_name == "DataFrame": return pd.read_csv(path, encoding=kwargs.get("encoding", "utf-8")) - if type_name in ("dict", "list"): + if type_name == "dict": with open(path, "r", encoding="utf-8") as f: return json.load(f) + if type_name == "list": + result = [] + with open(path, "r", encoding="utf-8") as f: + for line in f: + line = line.strip() + if line: + result.append(json.loads(line)) + return result if type_name == "str": return path.read_text(encoding=kwargs.get("encoding", "utf-8")) raise ValueError(f"Unknown data type in metadata: {type_name}") diff --git a/reme_ai/core/vector_store/chroma_vector_store.py b/reme_ai/core/vector_store/chroma_vector_store.py index 7baa5eab..3a483c23 100644 --- a/reme_ai/core/vector_store/chroma_vector_store.py +++ b/reme_ai/core/vector_store/chroma_vector_store.py @@ -117,14 +117,33 @@ class ChromaVectorStore(BaseVectorStore): @staticmethod def _generate_where_clause(filters: dict | None) -> dict | None: - """Convert the universal filter format to a ChromaDB-compatible where clause.""" + """Convert the universal filter format to a ChromaDB-compatible where clause. + + Supports two filter formats: + 1. Range query: {"field": [start_value, end_value]} - filters for field >= start_value AND field <= end_value + 2. Exact match: {"field": value} - filters for field == value + """ if not filters: return None - def convert_condition(k: str, v: Any) -> dict | None: - """Convert a single filter condition to ChromaDB operator format.""" + def convert_condition(k: str, v: Any) -> dict | list | None: + """Convert a single filter condition to ChromaDB operator format. + + Returns: + - dict for simple conditions + - list of dicts for range queries (which need to be wrapped in $and) + - None for wildcard filters + """ if v == "*": return None + # New syntax: [start, end] represents a range query + if isinstance(v, list) and len(v) == 2: + # Range query: field >= v[0] AND field <= v[1] + # ChromaDB requires separate conditions combined with $and + return [ + {k: {"$gte": v[0]}}, + {k: {"$lte": v[1]}} + ] if isinstance(v, dict): chroma_condition = {} for op, val in v.items(): @@ -141,8 +160,7 @@ class ChromaVectorStore(BaseVectorStore): chroma_op = mapping.get(op, "$eq") chroma_condition[k] = {chroma_op: val} return chroma_condition - if isinstance(v, list): - return {k: {"$in": v}} + # Exact match for non-list values return {k: {"$eq": v}} processed_filters = [] @@ -155,7 +173,11 @@ class ChromaVectorStore(BaseVectorStore): for sub_key, sub_value in condition.items(): converted = convert_condition(sub_key, sub_value) if converted: - or_condition.update(converted) + if isinstance(converted, list): + # Range query in OR condition - need to wrap in $and + or_conditions.append({"$and": converted}) + else: + or_condition.update(converted) if or_condition: or_conditions.append(or_condition) if len(or_conditions) > 1: @@ -168,13 +190,21 @@ class ChromaVectorStore(BaseVectorStore): for sub_key, sub_value in condition.items(): converted = convert_condition(sub_key, sub_value) if converted: - processed_filters.append(converted) + if isinstance(converted, list): + # Range query - add each condition separately + processed_filters.extend(converted) + else: + processed_filters.append(converted) elif key == "$not": continue else: converted = convert_condition(key, value) if converted: - processed_filters.append(converted) + if isinstance(converted, list): + # Range query - add each condition separately + processed_filters.extend(converted) + else: + processed_filters.append(converted) if not processed_filters: return None diff --git a/reme_ai/core/vector_store/es_vector_store.py b/reme_ai/core/vector_store/es_vector_store.py index 16226749..39989277 100644 --- a/reme_ai/core/vector_store/es_vector_store.py +++ b/reme_ai/core/vector_store/es_vector_store.py @@ -262,9 +262,19 @@ class ESVectorStore(BaseVectorStore): if filters: filter_conditions = [] for key, value in filters.items(): - if isinstance(value, list): - filter_conditions.append({"terms": {f"metadata.{key}": value}}) + # New syntax: [start, end] represents a range query + if isinstance(value, list) and len(value) == 2: + # Range query: field >= value[0] AND field <= value[1] + filter_conditions.append({ + "range": { + f"metadata.{key}": { + "gte": value[0], + "lte": value[1] + } + } + }) else: + # Exact match filter_conditions.append({"term": {f"metadata.{key}": value}}) search_query["knn"]["filter"] = {"bool": {"must": filter_conditions}} @@ -448,9 +458,19 @@ class ESVectorStore(BaseVectorStore): if filters: filter_conditions = [] for key, value in filters.items(): - if isinstance(value, list): - filter_conditions.append({"terms": {f"metadata.{key}": value}}) + # New syntax: [start, end] represents a range query + if isinstance(value, list) and len(value) == 2: + # Range query: field >= value[0] AND field <= value[1] + filter_conditions.append({ + "range": { + f"metadata.{key}": { + "gte": value[0], + "lte": value[1] + } + } + }) else: + # Exact match filter_conditions.append({"term": {f"metadata.{key}": value}}) query["query"] = {"bool": {"must": filter_conditions}} diff --git a/reme_ai/core/vector_store/local_vector_store.py b/reme_ai/core/vector_store/local_vector_store.py index cce3cae2..cfeec74e 100644 --- a/reme_ai/core/vector_store/local_vector_store.py +++ b/reme_ai/core/vector_store/local_vector_store.py @@ -91,17 +91,32 @@ class LocalVectorStore(BaseVectorStore): @staticmethod def _match_filters(node: VectorNode, filters: dict | None) -> bool: - """Check if a vector node matches the provided metadata filters.""" + """Check if a vector node matches the provided metadata filters. + + Supports two filter formats: + 1. Range query: {"field": [start_value, end_value]} - filters for field >= start_value AND field <= end_value + 2. Exact match: {"field": value} - filters for field == value + """ if not filters: return True for key, value in filters.items(): node_value = node.metadata.get(key) - if isinstance(value, list): - if node_value not in value: + # New syntax: [start, end] represents a range query + if isinstance(value, list) and len(value) == 2: + # Range query: field >= value[0] AND field <= value[1] + if node_value is None: + return False + try: + # Try numeric comparison + if not (value[0] <= node_value <= value[1]): + return False + except TypeError: + # If comparison fails, the filter doesn't match return False else: + # Exact match if node_value != value: return False diff --git a/reme_ai/core/vector_store/pgvector_store.py b/reme_ai/core/vector_store/pgvector_store.py index a23c84b9..576d8e6c 100644 --- a/reme_ai/core/vector_store/pgvector_store.py +++ b/reme_ai/core/vector_store/pgvector_store.py @@ -1,6 +1,7 @@ """PostgreSQL pgvector implementation for vector storage and retrieval.""" import json +import re from typing import Any from loguru import logger @@ -25,6 +26,25 @@ except ImportError as e: class PGVectorStore(BaseVectorStore): """Vector store implementation using PostgreSQL and pgvector for efficient similarity search.""" + @staticmethod + def _validate_table_name(name: str) -> None: + """Validate table name to prevent SQL injection. + + PostgreSQL table names must: + - Contain only alphanumeric characters and underscores + - Not start with a digit + - Be between 1 and 63 characters + """ + if not name: + raise ValueError("Table name cannot be empty") + if len(name) > 63: + raise ValueError(f"Table name too long: {len(name)} characters (max 63)") + if not re.match(r'^[a-zA-Z_][a-zA-Z0-9_]*$', name): + raise ValueError( + f"Invalid table name: {name}. Must start with letter or underscore, " + "and contain only alphanumeric characters and underscores." + ) + def __init__( self, collection_name: str, @@ -47,6 +67,9 @@ class PGVectorStore(BaseVectorStore): "PGVector requires extra dependencies. Install with `pip install asyncpg pgvector`", ) from _ASYNCPG_IMPORT_ERROR + # Validate collection name to prevent SQL injection + self._validate_table_name(collection_name) + super().__init__(collection_name=collection_name, embedding_model=embedding_model, **kwargs) self.dsn = dsn @@ -106,6 +129,7 @@ class PGVectorStore(BaseVectorStore): async def create_collection(self, collection_name: str, **kwargs): """Create a new PostgreSQL table with vector support and appropriate indexing.""" + self._validate_table_name(collection_name) pool = await self._get_pool() dimensions = kwargs.get("dimensions", self.embedding_model_dims) @@ -150,6 +174,7 @@ class PGVectorStore(BaseVectorStore): async def delete_collection(self, collection_name: str, **kwargs): """Remove the specified collection table from the database.""" + self._validate_table_name(collection_name) pool = await self._get_pool() async with pool.acquire() as conn: await conn.execute(f"DROP TABLE IF EXISTS {collection_name}") @@ -157,6 +182,7 @@ class PGVectorStore(BaseVectorStore): async def copy_collection(self, collection_name: str, **kwargs): """Duplicate the structure and content of the current collection to a new table.""" + self._validate_table_name(collection_name) pool = await self._get_pool() async with pool.acquire() as conn: @@ -252,7 +278,14 @@ class PGVectorStore(BaseVectorStore): @staticmethod def _build_filter_clause(filters: dict | None) -> tuple[str, list]: - """Generate an SQL WHERE clause and parameter list from a filter dictionary.""" + """Generate an SQL WHERE clause and parameter list from a filter dictionary. + + Supports two filter formats: + 1. Range query: {"field": [start_value, end_value]} - filters for field >= start_value AND field <= end_value + 2. Exact match: {"field": value} - filters for field == value + + Range queries support both numeric and string (e.g., timestamp strings) comparisons. + """ if not filters: return "", [] @@ -261,12 +294,28 @@ class PGVectorStore(BaseVectorStore): param_idx = 1 for key, value in filters.items(): - if isinstance(value, list): - placeholders = ", ".join([f"${param_idx + i}" for i in range(len(value))]) - conditions.append(f"metadata->>'{key}' IN ({placeholders})") - params.extend([str(v) for v in value]) - param_idx += len(value) + # Sanitize key to prevent SQL injection (only allow alphanumeric and underscore) + if not key.replace('_', '').replace('.', '').isalnum(): + raise ValueError(f"Invalid metadata key: {key}. Only alphanumeric characters, underscore and dot are allowed.") + + # New syntax: [start, end] represents a range query + if isinstance(value, list) and len(value) == 2: + # Range query: field >= value[0] AND field <= value[1] + # Try numeric comparison first, fall back to text comparison if needed + if isinstance(value[0], (int, float)) and isinstance(value[1], (int, float)): + # Numeric range query + conditions.append( + f"(metadata->>'{key}')::numeric >= ${param_idx} AND (metadata->>'{key}')::numeric <= ${param_idx + 1}" + ) + else: + # Text range query (works for strings, timestamps, etc.) + conditions.append( + f"metadata->>'{key}' >= ${param_idx} AND metadata->>'{key}' <= ${param_idx + 1}" + ) + params.extend([value[0], value[1]]) + param_idx += 2 else: + # Exact match conditions.append(f"metadata->>'{key}' = ${param_idx}") params.append(str(value)) param_idx += 1 @@ -290,11 +339,14 @@ class PGVectorStore(BaseVectorStore): filter_clause, filter_params = self._build_filter_clause(filters) + # Adjust parameter indices in filter clause to account for $1 being used by vector_str if filter_clause: - for i in range(len(filter_params)): - old_idx = i + 1 - new_idx = i + 2 - filter_clause = filter_clause.replace(f"${old_idx}", f"${new_idx}", 1) + # Replace from highest index to lowest to avoid conflicts + for i in range(len(filter_params), 0, -1): + old_placeholder = f"${i}" + new_placeholder = f"${i + 1}" + # Use word boundary to ensure we only replace exact matches (e.g., $1 not $10) + filter_clause = re.sub(rf'\${i}\b', new_placeholder, filter_clause) async with pool.acquire() as conn: sql = f""" diff --git a/reme_ai/core/vector_store/qdrant_vector_store.py b/reme_ai/core/vector_store/qdrant_vector_store.py index 1ac4db64..5d9fa4f5 100644 --- a/reme_ai/core/vector_store/qdrant_vector_store.py +++ b/reme_ai/core/vector_store/qdrant_vector_store.py @@ -246,29 +246,65 @@ class QdrantVectorStore(BaseVectorStore): @staticmethod def _create_filter(filters: dict) -> Filter | None: - """Convert a dictionary of filter conditions into a Qdrant Filter object.""" + """Convert a dictionary of filter conditions into a Qdrant Filter object. + + Supports two filter formats: + 1. Range query: {"field": [start_value, end_value]} - filters for field >= start_value AND field <= end_value + 2. Exact match: {"field": value} - filters for field == value + """ if not filters: return None conditions = [] for key, value in filters.items(): - if isinstance(value, dict) and ("gte" in value or "lte" in value): + # New syntax: [start, end] represents a range query + if isinstance(value, list) and len(value) == 2: + # Range query: field >= value[0] AND field <= value[1] + # Qdrant's Range only supports numeric values + if isinstance(value[0], (int, float)) and isinstance(value[1], (int, float)): + conditions.append( + FieldCondition( + key=f"metadata.{key}", + range=Range(gte=value[0], lte=value[1]), + ), + ) + else: + # For non-numeric values (e.g., string dates), Qdrant doesn't support range queries + # We need to skip this filter with a warning + logger.warning( + f"Qdrant does not support range queries for non-numeric values. " + f"Skipping range filter for key '{key}' with values {value}. " + f"Consider using numeric timestamps instead." + ) + elif isinstance(value, dict) and ("gte" in value or "lte" in value): range_params = {} + # Check if values are numeric if "gte" in value: - range_params["gte"] = value["gte"] + if isinstance(value["gte"], (int, float)): + range_params["gte"] = value["gte"] + else: + logger.warning( + f"Qdrant range filter for key '{key}' requires numeric gte value, got {type(value['gte']).__name__}. Skipping." + ) + continue if "lte" in value: - range_params["lte"] = value["lte"] - conditions.append( - FieldCondition( - key=f"metadata.{key}", - range=Range(**range_params), - ), - ) - elif isinstance(value, list): - conditions.append( - FieldCondition(key=f"metadata.{key}", match=MatchValue(value=value[0])), - ) + if isinstance(value["lte"], (int, float)): + range_params["lte"] = value["lte"] + else: + logger.warning( + f"Qdrant range filter for key '{key}' requires numeric lte value, got {type(value['lte']).__name__}. Skipping." + ) + continue + + if range_params: # Only add condition if we have valid numeric parameters + conditions.append( + FieldCondition( + key=f"metadata.{key}", + range=Range(**range_params), + ), + ) else: + # Exact match conditions.append( FieldCondition(key=f"metadata.{key}", match=MatchValue(value=value)), ) diff --git a/reme_ai/mem_agent/v3/__init__.py b/reme_ai/mem_agent/v3/__init__.py new file mode 100644 index 00000000..0fd9c86d --- /dev/null +++ b/reme_ai/mem_agent/v3/__init__.py @@ -0,0 +1,9 @@ +from .personal_summarizer_v3 import PersonalSummarizerV3 +from .reme_retriever_v3 import ReMeRetrieverV3 +from .reme_summarizer_v3 import ReMeSummarizerV3 + +__all__ = [ + "PersonalSummarizerV3", + "ReMeRetrieverV3", + "ReMeSummarizerV3", +] diff --git a/reme_ai/mem_agent/v3/personal_summarizer_v3.py b/reme_ai/mem_agent/v3/personal_summarizer_v3.py new file mode 100644 index 00000000..0093884d --- /dev/null +++ b/reme_ai/mem_agent/v3/personal_summarizer_v3.py @@ -0,0 +1,69 @@ +from ..base_memory_agent import BaseMemoryAgent +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message, ToolCall +from ...core.utils import format_messages + + +class PersonalSummarizerV3(BaseMemoryAgent): + memory_type: MemoryType = MemoryType.PERSONAL + + def _build_tool_call(self) -> ToolCall: + return ToolCall( + **{ + "description": self.get_prompt("tool"), + "parameters": { + "type": "object", + "properties": { + "messages": { + "type": "array", + "items": { + "type": "object", + "properties": { + "role": { + "type": "string", + "description": "role", + }, + "content": { + "type": "string", + "description": "content", + }, + }, + "required": ["role", "content"], + }, + }, + }, + "required": ["messages"], + }, + }, + ) + + async def build_messages(self) -> list[Message]: + """Construct messages with context, memory_target, and memory_type information.""" + system_prompt = self.prompt_format( + prompt_name="system_prompt", + context=self.description + "\n" + format_messages(self.get_messages()), + memory_type=self.memory_type.value, + memory_target=self.memory_target, + ) + + messages = [ + Message(role=Role.SYSTEM, content=system_prompt), + Message(role=Role.USER, content=self.get_prompt("user_message")), + ] + return messages + + async def _reasoning_step(self, messages: list[Message], step: int, **kwargs) -> tuple[Message, bool]: + return await super()._reasoning_step(messages, step, **kwargs) + + async def _acting_step(self, assistant_message: Message, step: int, **kwargs) -> list[Message]: + """Execute tool calls with memory_target, memory_type, and author context.""" + messages: list[Message] = await super()._acting_step( + assistant_message, + step, + memory_type=self.memory_type.value, + memory_target=self.memory_target, + ref_memory_id=self.ref_memory_id, + author=self.author, + **kwargs, + ) + return messages diff --git a/reme_ai/mem_agent/v3/personal_summarizer_v3.yaml b/reme_ai/mem_agent/v3/personal_summarizer_v3.yaml new file mode 100644 index 00000000..97d18afc --- /dev/null +++ b/reme_ai/mem_agent/v3/personal_summarizer_v3.yaml @@ -0,0 +1,38 @@ +tool: | + Extract and update personal memories about the user from conversation context. + Analyze dialogues to identify preferences, habits, background, relationships, and key facts. + +system_prompt: | + You are a memory agent managing **{memory_type}** memories about **{memory_target}**. + + ## Latest Conversation: + {context} + + Each message format: `round [] : ` (timestamp: YYYY-MM-DD HH:MM:SS). + + **CRITICAL**: Extract ONLY explicitly stated information. DO NOT infer, assume, or fabricate. + + ## Three-Step Workflow + + ### Step 1: Extract Conversation Memories + Use `AddMemory` to extract key personal facts from the conversation. + - Extract: preferences, habits, status, personal details, decisions, conclusions + - Keep entries concise and distinct (no duplicates, no omissions) + - Record `conversation_time` for each memory (format: 2020-01-01 00:00:00; use 0000-00-00 00:00:00 if unavailable) + + ### Step 2: Read User Profile + Use `ReadUserProfile` to retrieve the current user profile. + - Review existing memories to identify conflicts and duplicates + + ### Step 3: Update User Profile + Use `UpdateUserProfile` to synchronize the profile with new information. + - `profile_ids_to_delete`: Remove outdated or conflicting profiles + - `profiles_to_add`: Add new profiles that are not duplicates + - Use `timestamp` from conversation_time (format: 2020-01-01 00:00:00) + - Keep final profiles concise with no information loss + +user_message: | + Execute the three-step workflow: + 1. Use `AddMemory` to extract personal memories from the conversation + 2. Use `ReadUserProfile` to read existing user profile + 3. Use `UpdateUserProfile` to remove outdated entries and add new profiles diff --git a/reme_ai/mem_agent/v3/reme_retriever_v3.py b/reme_ai/mem_agent/v3/reme_retriever_v3.py new file mode 100644 index 00000000..8f5c62dc --- /dev/null +++ b/reme_ai/mem_agent/v3/reme_retriever_v3.py @@ -0,0 +1,44 @@ +"""ReMe retriever v2 that autonomously retrieves memories from multiple angles.""" + +from typing import List + +from ..base_memory_agent import BaseMemoryAgent +from ...core.enumeration import Role +from ...core.schema import Message +from ...core.utils import format_messages + + +class ReMeRetrieverV3(BaseMemoryAgent): + + def __init__(self, meta_memories: list[dict] | None = None, **kwargs): + super().__init__(**kwargs) + self.meta_memories: list[dict] = meta_memories or [] + + async def _read_meta_memories(self) -> str: + """Fetch all meta-memory entries that define specialized memory agents.""" + from ...mem_tool import ReadMetaMemory + + op = ReadMetaMemory(enable_identity_memory=False) + return op.format_memory_metadata(self.meta_memories) + + async def build_messages(self) -> List[Message]: + """Build messages with system prompt and user message.""" + if self.context.get("query"): + context = self.context.query + elif self.context.get("messages"): + context = format_messages(self.context.messages) + else: + raise ValueError("input must have either `query` or `messages`") + + system_prompt = self.prompt_format( + prompt_name="system_prompt", + meta_memory_info=await self._read_meta_memories(), + context=context, + ) + + messages = [ + Message(role=Role.SYSTEM, content=system_prompt), + Message(role=Role.USER, content=self.get_prompt("user_message")), + ] + + return messages diff --git a/reme_ai/mem_agent/v3/reme_retriever_v3.yaml b/reme_ai/mem_agent/v3/reme_retriever_v3.yaml new file mode 100644 index 00000000..8ff005ee --- /dev/null +++ b/reme_ai/mem_agent/v3/reme_retriever_v3.yaml @@ -0,0 +1,53 @@ +tool: | + Autonomously retrieve relevant memories through a three-step strategy to answer user questions. + Steps: read user profile → vector search with multiple angles → read original conversations. + State "I don't know" if information cannot be found after exhaustive searching. + NEVER hallucinate or fabricate information not present in retrieved memories. + +system_prompt: | + You are a memory retrieval agent. Search for relevant memories to answer the user's question following this strategy: + + ## Available Meta Memories + Format: "- (): " + {meta_memory_info} + + ## User Context + {context} + + ## Three-Step Retrieval Strategy + + **STEP 1: Read User Profile (REQUIRED FIRST)** + - Use `read_user_profile` with memory_type and memory_target from available meta memories + - Check if the user profile directly answers the question + - If sufficient information found, provide the answer and STOP + + **STEP 2: Vector Search (If Step 1 insufficient)** + - Use `retrieve_memory` with memory_type, memory_target, and query + - Try multiple retrieval angles (at least 3 different attempts): + * Direct query with user's question + * Reformulated queries with different phrasing/keywords + * Queries focused on specific entities or concepts + + - **Time Range Filtering** (when applicable): + * Format: [start_date, end_date] in YYYYMMDD format + * Example: [20200101, 20200102] means 20200101 < time < 20200102 + * Single-sided: [0, 20200102] for before, [20200101, 99999999] for after + * If no results, try broader time ranges or remove time constraints + + - If no results after multiple attempts, try different memory_type/memory_target combinations + + **STEP 3: Read Original Conversations (If Step 2 insufficient)** + - Use `read_history` with history_id from retrieved memories + - Prioritize reading: + * Most recent memories with history_id + * Most relevant memories from Step 2 with history_id + - Try multiple history_id entries if needed + + ## Response Rules + - Answer ONLY based on retrieved information - NEVER guess or fabricate + - If nothing found after all three steps: State clearly "I don't know. I cannot find relevant information to answer this question." + - Be persistent: try multiple angles in each step before moving to the next + - Once you find sufficient information, provide a direct answer + +user_message: | + Retrieve relevant memories and answer the question using the three-step strategy. diff --git a/reme_ai/mem_agent/v3/reme_summarizer_v3.py b/reme_ai/mem_agent/v3/reme_summarizer_v3.py new file mode 100644 index 00000000..a0f466b9 --- /dev/null +++ b/reme_ai/mem_agent/v3/reme_summarizer_v3.py @@ -0,0 +1,88 @@ +from loguru import logger + +from ..base_memory_agent import BaseMemoryAgent +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message, MemoryNode, ToolCall +from ...core.utils import format_messages + + +class ReMeSummarizerV3(BaseMemoryAgent): + + def __init__(self, meta_memories: list[dict] | None = None, **kwargs): + """Initialize with meta memories list.""" + super().__init__(**kwargs) + self.meta_memories: list[dict] = meta_memories or [] + + def _build_tool_call(self) -> ToolCall: + return ToolCall( + **{ + "description": self.get_prompt("tool"), + "parameters": { + "type": "object", + "properties": { + "messages": { + "type": "array", + "items": { + "type": "object", + "properties": { + "role": { + "type": "string", + "description": "role", + }, + "content": { + "type": "string", + "description": "content", + }, + }, + "required": ["role", "content"], + }, + }, + }, + "required": ["messages"], + }, + }, + ) + + async def _read_meta_memories(self) -> str: + from ...mem_tool import ReadMetaMemory + + return ReadMetaMemory().format_memory_metadata(self.meta_memories) + + async def build_messages(self) -> list[Message]: + """Construct initial messages with context and meta-memory information.""" + messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages] + self.context["messages_formated"] = self.description + "\n" + format_messages(messages) + self.context["ref_memory_id"] = MemoryNode( + memory_type=MemoryType.HISTORY, + content=self.context["messages_formated"], + ).memory_id + + meta_memory_info = await self._read_meta_memories() + logger.info(f"meta_memory_info={meta_memory_info}") + + system_prompt = self.prompt_format( + prompt_name="system_prompt", + meta_memory_info=meta_memory_info, + context=self.context["messages_formated"], + ) + + user_message = self.get_prompt("user_message") + messages = [ + Message(role=Role.SYSTEM, content=system_prompt), + Message(role=Role.USER, content=user_message), + ] + + return messages + + async def _acting_step(self, assistant_message: Message, step: int, **kwargs) -> list[Message]: + """Execute tool calls with ref_memory_id and author context.""" + return await super()._acting_step( + assistant_message, + step, + messages=self.context.get("messages", []), + description=self.context.get("description"), + ref_memory_id=self.context["ref_memory_id"], + messages_formated=self.context["messages_formated"], + author=self.author, + **kwargs, + ) diff --git a/reme_ai/mem_agent/v3/reme_summarizer_v3.yaml b/reme_ai/mem_agent/v3/reme_summarizer_v3.yaml new file mode 100644 index 00000000..30792a08 --- /dev/null +++ b/reme_ai/mem_agent/v3/reme_summarizer_v3.yaml @@ -0,0 +1,25 @@ +tool: | + Orchestrate the complete memory summarization for the agent. + +system_prompt: | + You are a Memory Agent responsible for performing necessary updates and summaries of the main Agent's memories based on the **context**. + + # Context + {context} + + ## Main Agent's Meta Memory + Each line of meta memory indicates the existence of a specialized Memory Agent dedicated to deep summarization and updating of memories within a specific dimension (memory_type + memory_target). + Format: "- (): " + {meta_memory_info} + + ## Your Task + Use `summary_and_hands_off` tool to: + 1. Create a concise summary in `summary_content` that captures key points, decisions, or important facts from the context. + 2. Identify which memory dimensions need updates and specify them in `memory_tasks` (each with `memory_type` and `memory_target`). + - The `memory_type` and `memory_target` must exactly match existing entries in the "Main Agent's Meta Memory" listed above. + - Multiple tasks can be specified to enable parallel processing by specialized agents. + + Note: If the context contains no memorable information (e.g., simple greetings), output ``. + +user_message: | + Please perform your task based on the context. diff --git a/reme_ai/mem_tool/base_memory_tool.py b/reme_ai/mem_tool/base_memory_tool.py index 6accfe2a..0d92671c 100644 --- a/reme_ai/mem_tool/base_memory_tool.py +++ b/reme_ai/mem_tool/base_memory_tool.py @@ -3,8 +3,6 @@ from abc import ABCMeta from pathlib import Path -from loguru import logger - from ..core.enumeration import MemoryType from ..core.op import BaseOp from ..core.schema import ToolCall, MemoryNode diff --git a/reme_ai/mem_tool/read_local_memories.py b/reme_ai/mem_tool/read_local_memories.py new file mode 100644 index 00000000..98d96187 --- /dev/null +++ b/reme_ai/mem_tool/read_local_memories.py @@ -0,0 +1,54 @@ +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.context import C +from ...core.schema.memory_node import MemoryNode + + +@C.register_op() +class ReadLocalMemories(BaseMemoryTool): + + def __init__(self, **kwargs): + kwargs["enable_multiple"] = False + super().__init__(**kwargs) + + def _build_parameters(self) -> dict: + return { + "type": "object", + "properties": { + "memory_type": { + "type": "string", + "description": self.get_prompt("memory_type"), + }, + "memory_target": { + "type": "string", + "description": self.get_prompt("memory_target"), + }, + }, + "required": ["memory_type", "memory_target"], + } + + async def execute(self): + memory_type = self.context.get("memory_type", "") + memory_target = self.context.get("memory_target", "") + + if not memory_type or not memory_target: + self.output = "memory_type and memory_target are required." + return + + cache_key = f"{memory_type}_{memory_target}" + cached_data = self.meta_memory.load(cache_key, auto_clean=False) + + if not cached_data: + self.output = f"Local memory not found: {memory_type}_{memory_target}" + logger.info(self.output) + return + + memory_nodes = [MemoryNode(**node_data) for node_data in cached_data] + + if not memory_nodes: + self.output = f"No valid memory nodes found in {memory_type}_{memory_target}" + return + + self.output = memory_nodes + logger.info(f"Read {len(memory_nodes)} nodes from cache key: {cache_key}") diff --git a/reme_ai/mem_tool/read_local_memories.yaml b/reme_ai/mem_tool/read_local_memories.yaml new file mode 100644 index 00000000..c8e155f6 --- /dev/null +++ b/reme_ai/mem_tool/read_local_memories.yaml @@ -0,0 +1,8 @@ +tool: | + Read memory nodes from local memory files. + +memory_type: | + The type of local memory to read. + +memory_target: | + The target identifier for the local memory. diff --git a/reme_ai/mem_tool/v3/__init__.py b/reme_ai/mem_tool/v3/__init__.py new file mode 100644 index 00000000..9d06f4c8 --- /dev/null +++ b/reme_ai/mem_tool/v3/__init__.py @@ -0,0 +1,15 @@ +from .add_memory import AddMemory +from .read_history import ReadHistory +from .read_user_profile import ReadUserProfile +from .retrieve_memory import RetrieveMemory +from .summary_and_hands_off import SummaryAndHandsOff +from .update_user_profile import UpdateUserProfile + +__all__ = [ + "AddMemory", + "ReadHistory", + "ReadUserProfile", + "RetrieveMemory", + "SummaryAndHandsOff", + "UpdateUserProfile", +] diff --git a/reme_ai/mem_tool/v3/add_memory.py b/reme_ai/mem_tool/v3/add_memory.py new file mode 100644 index 00000000..a013db42 --- /dev/null +++ b/reme_ai/mem_tool/v3/add_memory.py @@ -0,0 +1,67 @@ +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.schema import MemoryNode + + +class AddMemory(BaseMemoryTool): + + def __init__(self, **kwargs): + kwargs['enable_multiple'] = True + super().__init__(**kwargs) + + def _build_tool_description(self) -> str: + return "Add multiple memories to the vector store for future retrieval." + + def _build_multiple_parameters(self) -> dict: + return { + "type": "object", + "properties": { + "memories": { + "type": "array", + "description": "A list of memory objects to store.", + "items": { + "type": "object", + "properties": { + "memory_content": { + "type": "string", + "description": "memory content", + }, + "conversation_time": { + "type": "object", + "description": "conversation time, e.g. '2020-01-01 00:00:00'", + } + }, + "required": ["memory_content", "conversation_time"], + }, + }, + }, + "required": ["memories"], + } + + async def execute(self): + memories: list[dict] = self.context.get("memories", []) + if not memories: + self.output = "No memories provided for addition." + return + + memory_nodes: list[MemoryNode] = [] + for mem in memories: + memory_content = mem.get("memory_content", "") + conversation_time = mem.get("conversation_time", "") + metadata: dict = {"conversation_time": conversation_time} + try: + metadata["time_int"] = int(conversation_time.split(" ")[0].replace("-", "")) + except Exception: + ... + memory_nodes.append(self._build_memory_node(memory_content, metadata=metadata)) + + vector_nodes = [node.to_vector_node() for node in memory_nodes] + vector_ids: list[str] = [node.vector_id for node in vector_nodes] + + await self.vector_store.delete(vector_ids=vector_ids) + await self.vector_store.insert(nodes=vector_nodes) + self.memory_nodes = memory_nodes + + self.output = f"Successfully added {len(memory_nodes)} memories to vector_store." + logger.info(self.output) diff --git a/reme_ai/mem_tool/v3/read_history.py b/reme_ai/mem_tool/v3/read_history.py new file mode 100644 index 00000000..e9ab2a15 --- /dev/null +++ b/reme_ai/mem_tool/v3/read_history.py @@ -0,0 +1,38 @@ +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.schema import MemoryNode + + +class ReadHistory(BaseMemoryTool): + def __init__(self, **kwargs): + kwargs["enable_multiple"] = False + super().__init__(**kwargs) + + def _build_tool_description(self) -> str: + return "Read original history dialogue." + + def _build_parameters(self) -> dict: + return { + "type": "object", + "properties": { + "history_id": { + "type": "string", + "description": "history_id", + }, + }, + "required": ["history_id"], + } + + async def execute(self): + history_id = self.context.get("history_id", "") + nodes = await self.vector_store.get(vector_ids=[history_id]) + + if not nodes: + self.output = f"No history: {history_id}" + logger.warning(self.output) + return + + memory = MemoryNode.from_vector_node(nodes[0]) + self.output = memory.content + logger.info(f"Successfully read history memory: {history_id}") diff --git a/reme_ai/mem_tool/v3/read_user_profile.py b/reme_ai/mem_tool/v3/read_user_profile.py new file mode 100644 index 00000000..57cff042 --- /dev/null +++ b/reme_ai/mem_tool/v3/read_user_profile.py @@ -0,0 +1,65 @@ +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.schema.memory_node import MemoryNode + + +class ReadUserProfile(BaseMemoryTool): + + def __init__(self, add_memory_type_target: bool = True, **kwargs): + kwargs["enable_multiple"] = False + self.add_memory_type_target = add_memory_type_target + super().__init__(**kwargs) + + def _build_tool_description(self) -> str: + return "Read personal memory profile for the current user." + + def _build_parameters(self) -> dict: + if self.add_memory_type_target: + return { + "type": "object", + "properties": { + "memory_type": { + "type": "string", + "description": "memory_type", + }, + "memory_target": { + "type": "string", + "description": "memory_target", + }, + }, + "required": ["memory_type", "memory_target"], + } + else: + return { + "type": "object", + "properties": {}, + "required": [], + } + + async def execute(self): + cache_key = f"{self.memory_type}_{self.memory_target}" + cached_data = self.meta_memory.load(cache_key, auto_clean=False) + + if not cached_data: + self.output = f"Local memory not found: {self.memory_type}_{self.memory_target}" + logger.info(self.output) + return + + # Convert to MemoryNode objects and sort by conversation_time (oldest first) + memory_nodes = [MemoryNode(**node_data) for node_data in cached_data] + memory_nodes.sort( + key=lambda node: node.metadata.get("conversation_time", "") + ) + + memory_formated = [] + for node in memory_nodes: + node_formated = f"profile_id={node.memory_id} profile_content={node.content}" + if "conversation_time" in node.metadata: + node_formated += f" conversation_time={node.metadata['conversation_time']}" + if node.ref_memory_id: + node_formated += f" history_id={node.ref_memory_id}" + memory_formated.append(node_formated.strip()) + + self.output = "\n".join(memory_formated) + logger.info(f"Read {len(memory_formated)} nodes from cache key: {cache_key}") diff --git a/reme_ai/mem_tool/v3/retrieve_memory.py b/reme_ai/mem_tool/v3/retrieve_memory.py new file mode 100644 index 00000000..32e526d7 --- /dev/null +++ b/reme_ai/mem_tool/v3/retrieve_memory.py @@ -0,0 +1,84 @@ +import json + +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.schema import MemoryNode +from ...core.utils import deduplicate_memories + + +class RetrieveMemory(BaseMemoryTool): + + def __init__(self, top_k: int = 20, **kwargs): + super().__init__(**kwargs) + self.top_k: int = top_k + + def _build_tool_description(self) -> str: + return "Retrieve memories using vector similarity search." + + def _build_multiple_parameters(self) -> dict: + return { + "type": "object", + "properties": { + "query_items": { + "type": "array", + "description": "query_items", + "items": { + "type": "object", + "properties": { + "memory_type": { + "type": "string", + "description": "memory_type", + }, + "memory_target": { + "type": "string", + "description": "memory_target", + }, + "query": { + "type": "string", + "description": "query", + }, + "time_range": { + "type": "string", + "description": "time_range(optional), e.g. [20200101, 20200101]", + }, + }, + "required": ["memory_type", "memory_target", "query"], + }, + }, + }, + "required": ["query_items"], + } + + async def execute(self): + query_items: list[dict] = self.context.get("query_items", []) + memory_nodes: list[MemoryNode] = [] + for query_item in query_items: + memory_type = query_item.get("memory_type") + memory_target = query_item.get("memory_target") + query = query_item.get("query") + time_range = query_item.get("time_range", "") + + filter_dict = { + "memory_type": memory_type, + "memory_target": memory_target, + } + + if time_range: + time_range = json.loads(time_range) + filter_dict["time_range"] = [int(time_range[0]), int(time_range[1])] + + nodes = await self.vector_store.search(query=query, limit=self.top_k, filters=filter_dict) + memory_nodes.extend([MemoryNode.from_vector_node(n) for n in nodes]) + memory_nodes = deduplicate_memories(memory_nodes) + + retrieved_memory_ids = {node.memory_id for node in self.retrieved_nodes if node.memory_id} + new_memory_nodes = [node for node in memory_nodes if node.memory_id not in retrieved_memory_ids] + self.retrieved_nodes.extend(new_memory_nodes) + self.memory_nodes = new_memory_nodes + + if not new_memory_nodes: + self.output = "No new memory_nodes found matching the query (duplicates removed)." + else: + self.output = "\n".join([f"{m.metadata['conversation_time']} {m.content}" for m in new_memory_nodes]) + logger.info(f"Retrieved {len(memory_nodes)} memory_nodes, {len(new_memory_nodes)} new after deduplication") diff --git a/reme_ai/mem_tool/v3/summary_and_hands_off.py b/reme_ai/mem_tool/v3/summary_and_hands_off.py new file mode 100644 index 00000000..19d88744 --- /dev/null +++ b/reme_ai/mem_tool/v3/summary_and_hands_off.py @@ -0,0 +1,140 @@ +import json +from typing import TYPE_CHECKING + +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.enumeration import MemoryType +from ...core.schema import MemoryNode, Message + +if TYPE_CHECKING: + from ...mem_agent import BaseMemoryAgent + + +class SummaryAndHandsOff(BaseMemoryTool): + def __init__(self, memory_agents: list["BaseMemoryAgent"], **kwargs): + kwargs["enable_multiple"] = True + kwargs["sub_ops"] = memory_agents or [] + super().__init__(**kwargs) + from ...mem_agent import BaseMemoryAgent + + self.sub_ops: list[BaseMemoryAgent] = [a for a in self.sub_ops if isinstance(a, BaseMemoryAgent)] + self.messages: list[Message] = [] + + @property + def memory_agent_dict(self) -> dict[MemoryType, "BaseMemoryAgent"]: + return {a.memory_type: a for a in self.sub_ops} + + def _build_tool_description(self) -> str: + return "Summarize and distribute memory tasks to appropriate agents." + + def _build_multiple_parameters(self) -> dict: + return { + "type": "object", + "properties": { + "summary_content": { + "type": "string", + "description": "summary content", + }, + "memory_tasks": { + "type": "array", + "description": "memory_tasks", + "items": { + "type": "object", + "properties": { + "memory_type": { + "type": "string", + "description": "memory_type", + "enum": [k.value for k in self.memory_agent_dict], + }, + "memory_target": { + "type": "string", + "description": "memory_target", + }, + }, + "required": ["memory_type", "memory_target"], + }, + }, + }, + "required": ["summary_content", "memory_tasks"], + } + + @staticmethod + def _parse_memory_type_target(task: dict): + return { + "memory_type": MemoryType(task.get("memory_type", "")), + "memory_target": task.get("memory_target", ""), + } + + def _collect_tasks(self) -> list[dict]: + tasks = [] + for task in self.context.get("memory_tasks", []): + tasks.append(self._parse_memory_type_target(task)) + return tasks + + async def execute(self): + summary_content = self.context.get("summary_content", "") + assert summary_content, "No summary content provided." + + summary_node = MemoryNode( + memory_type=MemoryType.HISTORY, + memory_target="", + when_to_use=summary_content, + content=self.messages_formated, + ref_memory_id="", + author=self.author, + metadata={}, + ) + logger.info(f"Adding summary node: {summary_node.model_dump_json(indent=2, exclude_none=True)}") + self.memory_nodes.append(summary_node) + vector_node = summary_node.to_vector_node() + await self.vector_store.delete(vector_ids=[vector_node.vector_id]) + await self.vector_store.insert([vector_node]) + + tasks = self._collect_tasks() + if not tasks: + self.output = "No valid memory tasks to execute." + return + + agent_list = [] + for i, task in enumerate(tasks): + memory_type: MemoryType = task["memory_type"] + memory_target: str = task["memory_target"] + + if memory_type not in self.memory_agent_dict: + logger.warning(f"No agent found for memory_type={memory_type}") + continue + + agent = self.memory_agent_dict[memory_type].copy() + agent_list.append([agent, memory_type, memory_target]) + + logger.info(f"Task {i}: Submitting {memory_type.value} agent for target={memory_target}") + self.submit_async_task( + agent.call, + query=self.context.get("query", ""), + messages=self.context.get("messages", []), + memory_type=memory_type, + memory_target=memory_target, + description=self.context.get("description"), + ref_memory_id=self.context.get("ref_memory_id", ""), + ) + + await self.join_async_tasks() + + results = [] + for i, (agent, memory_type, memory_target) in enumerate(agent_list): + result_str = str(agent.output) + if agent.memory_nodes: + self.memory_nodes.extend(agent.memory_nodes) + if agent.messages: + self.messages.extend(agent.messages) + + results.append({ + "memory_type": memory_type.value, + "memory_target": memory_target, + "result": result_str[:100] + ("..." if len(result_str) > 100 else ""), + }) + logger.info(f"Task {i}: Completed {memory_type.value} agent for target={memory_target}") + + results_str = json.dumps(results, ensure_ascii=False, indent=2) + self.output = f"Successfully executed summary and {len(results)} hands-off task(s):\n{results_str}" diff --git a/reme_ai/mem_tool/v3/update_user_profile.py b/reme_ai/mem_tool/v3/update_user_profile.py new file mode 100644 index 00000000..56ad4284 --- /dev/null +++ b/reme_ai/mem_tool/v3/update_user_profile.py @@ -0,0 +1,118 @@ +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.context import C +from ...core.schema.memory_node import MemoryNode + + +@C.register_op() +class UpdateUserProfile(BaseMemoryTool): + + def __init__(self, **kwargs): + kwargs["enable_multiple"] = True + super().__init__(**kwargs) + + def _build_multiple_parameters(self) -> dict: + return { + "type": "object", + "properties": { + "profile_ids_to_delete": { + "type": "array", + "description": self.get_prompt("profile_ids_to_delete"), + "items": {"type": "string"}, + }, + "profiles_to_add": { + "type": "array", + "description": self.get_prompt("profiles_to_add"), + "items": { + "type": "object", + "properties": { + "profile_content": { + "type": "string", + "description": self.get_prompt("profile_content"), + }, + "timestamp": { + "type": "string", + "description": self.get_prompt("timestamp"), + }, + }, + "required": ["profile_content", "timestamp"], + }, + }, + }, + "required": ["profile_ids_to_delete", "profiles_to_add"], + } + + async def execute(self): + memory_type = "personal" + memory_target = self.memory_target + assert memory_target, "memory_target is not configured." + + cache_key = f"{memory_type}_{memory_target}" + + profile_ids_to_delete = self.context.get("profile_ids_to_delete", []) + profile_ids_to_delete = [m for m in profile_ids_to_delete if m] + profile_ids_to_delete = list(dict.fromkeys(profile_ids_to_delete)) + + profiles_to_add = self.context.get("profiles_to_add", []) + + if not profile_ids_to_delete and not profiles_to_add: + self.output = "No memories to remove or add. Operation has been done." + return + + cached_data = self.meta_memory.load(cache_key, auto_clean=False) + existing_memory_nodes = [] + if cached_data: + existing_memory_nodes = [MemoryNode(**node_data) for node_data in cached_data] + + removed_count = 0 + added_count = 0 + + if profile_ids_to_delete: + profile_ids_set = set(profile_ids_to_delete) + existing_memory_nodes = [ + node for node in existing_memory_nodes if node.memory_id not in profile_ids_set + ] + removed_count = len(profile_ids_to_delete) + logger.info(f"Removed {removed_count} memories from user profile.") + + new_memory_nodes = [] + if profiles_to_add: + for mem in profiles_to_add: + profile_content = mem.get("profile_content", "") + timestamp = mem.get("timestamp", "") + + if not profile_content: + logger.warning("Skipping memory with empty content") + continue + + memory_node = self._build_memory_node( + memory_content=profile_content, + when_to_use="", + metadata={"timestamp": timestamp} + ) + memory_node.memory_type = MemoryNode.MemoryType.PERSONAL + memory_node.memory_target = memory_target + + new_memory_nodes.append(memory_node) + + added_count = len(new_memory_nodes) + logger.info(f"Added {added_count} new memories to user profile.") + + updated_memory_nodes = existing_memory_nodes + new_memory_nodes + + nodes_data = [node.model_dump(exclude_none=True) for node in updated_memory_nodes] + self.meta_memory.save(cache_key, nodes_data) + + operations = [] + if removed_count > 0: + operations.append(f"removed {removed_count} old memories") + if added_count > 0: + operations.append(f"added {added_count} new memories") + + if operations: + self.output = f"Successfully {' and '.join(operations)} in user profile." + else: + self.output = "Operation has been done." + + logger.info(self.output) diff --git a/reme_ai/mem_tool/write_local_memories.py b/reme_ai/mem_tool/write_local_memories.py new file mode 100644 index 00000000..ac1c5aae --- /dev/null +++ b/reme_ai/mem_tool/write_local_memories.py @@ -0,0 +1,57 @@ +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.context import C +from ...core.schema.memory_node import MemoryNode + + +@C.register_op() +class WriteLocalMemories(BaseMemoryTool): + + def __init__(self, **kwargs): + kwargs["enable_multiple"] = True + super().__init__(**kwargs) + + def _build_multiple_parameters(self) -> dict: + return { + "type": "object", + "properties": { + "memory_nodes": { + "type": "array", + "description": self.get_prompt("memory_nodes"), + "items": { + "type": "object", + "description": "Memory node object", + }, + }, + }, + "required": ["memory_nodes"], + } + + async def execute(self): + memory_nodes = self.context.get("memory_nodes", []) + + if not memory_nodes: + self.output = "No memory nodes provided." + return + + memory_nodes = [MemoryNode(**node) if isinstance(node, dict) else node for node in memory_nodes] + + grouped = {} + for node in memory_nodes: + key = (node.memory_type.value, node.memory_target) + if key not in grouped: + grouped[key] = [] + grouped[key].append(node) + + written_keys = [] + + for (memory_type, memory_target), nodes in grouped.items(): + cache_key = f"{memory_type}_{memory_target}" + nodes_data = [node.model_dump() for node in nodes] + + self.meta_memory.save(cache_key, nodes_data) + written_keys.append(f"{memory_type}_{memory_target}") + logger.info(f"Saved {len(nodes)} nodes to cache key: {cache_key}") + + self.output = f"Successfully written local memories: {', '.join(written_keys)}" diff --git a/reme_ai/mem_tool/write_local_memories.yaml b/reme_ai/mem_tool/write_local_memories.yaml new file mode 100644 index 00000000..81615b1b --- /dev/null +++ b/reme_ai/mem_tool/write_local_memories.yaml @@ -0,0 +1,5 @@ +tool_multiple: | + Write memory nodes to local memory files. + +memory_nodes: | + List of memory nodes to write to local files. diff --git a/reme_ai/reme.py b/reme_ai/reme.py index e3aaf54d..ff9d8822 100644 --- a/reme_ai/reme.py +++ b/reme_ai/reme.py @@ -13,6 +13,11 @@ from .mem_agent.retriever import ReMeRetriever from .mem_agent.retriever_v2 import ReMeRetrieverV2 from .mem_agent.summarizer import ReMeSummarizer, PersonalSummarizer from .mem_agent.summarizer_v2 import ReMeSummarizerV2, PersonalSummarizerV2 +from .mem_agent.v3 import ( + PersonalSummarizerV3, + ReMeRetrieverV3, + ReMeSummarizerV3, +) from .mem_tool import ( HandsOffTool, ReadHistoryMemory, @@ -24,12 +29,19 @@ from .mem_tool import ( ) from .mem_tool.v2 import ( AddMemoryDrafts, - ReadHistory, RetrieveMemories, RetrieveRecentAndSimilarMemories, SummaryAndHandsOff, UpdateMemories, ) +from .mem_tool.v3 import ( + AddMemory as AddMemoryV3, + ReadHistory as ReadHistoryV3, + ReadUserProfile, + RetrieveMemory, + SummaryAndHandsOff as SummaryAndHandsOffV3, + UpdateUserProfile, +) @singleton @@ -314,3 +326,86 @@ class ReMe(Application): else: raise NotImplementedError + + async def summary_v3( + self, + messages: list[dict], + description: str = "", + user_id: str = "", + assistant_id: str = "", + **kwargs, + ): + """Summarizes messages using V3 workflow with user profile management.""" + + if user_id: + meta_memories = [ + { + "memory_type": "personal", + "memory_target": user_id, + }, + ] + messages = self._prepare_messages(messages, user_id, assistant_id) + + personal_summarizer_v3 = PersonalSummarizerV3( + tools=[ + AddMemoryV3(), + ReadUserProfile(add_memory_type_target=False), + UpdateUserProfile(), + ], + ) + + reme_summarizer_v3 = ReMeSummarizerV3( + meta_memories=meta_memories, + tools=[SummaryAndHandsOffV3(memory_agents=[personal_summarizer_v3])], + ) + + # try: + await reme_summarizer_v3.call(messages=messages, description=description, **kwargs) + return reme_summarizer_v3.memory_nodes, reme_summarizer_v3.messages, reme_summarizer_v3.success + # except Exception as e: + # print(f"Warning: reme_summarizer_v3.call failed: {e}") + # return [], [], False + + else: + raise NotImplementedError + + async def retrieve_v3( + self, + query: str = "", + messages: list[dict] | None = None, + description: str = "", + user_id: str = "", + assistant_id: str = "", + top_k: int = 20, + **kwargs, + ): + """Retrieves relevant memories using V3 workflow with user profile support.""" + + if user_id: + messages = self._prepare_messages(messages, user_id, assistant_id) + + meta_memories = [ + { + "memory_type": "personal", + "memory_target": user_id, + }, + ] + + reme_retriever_v3 = ReMeRetrieverV3( + meta_memories=meta_memories, + tools=[ + ReadUserProfile(add_memory_type_target=True), + RetrieveMemory(top_k=top_k), + ReadHistoryV3(), + ], + ) + + # try: + await reme_retriever_v3.call(query=query, messages=messages, description=description, **kwargs) + return reme_retriever_v3.output, reme_retriever_v3.messages, reme_retriever_v3.success + # except Exception as e: + # print(f"Warning: reme_retriever_v3.call failed: {e}") + # return "error, not retrieved", [], False + + else: + raise NotImplementedError diff --git a/tests/test_vector_store.py b/tests/test_vector_store.py index 2f4c1f4d..6ef9ab0e 100644 --- a/tests/test_vector_store.py +++ b/tests/test_vector_store.py @@ -350,35 +350,36 @@ async def test_search_with_single_filter(store: BaseVectorStore, _store_name: st logger.info("✓ Single filter search test passed") -async def test_search_with_list_filter(store: BaseVectorStore, _store_name: str): - """Test vector search with list filter (IN operation).""" - logger.info("=" * 20 + " LIST FILTER SEARCH TEST " + "=" * 20) +async def test_search_with_exact_match_filter(store: BaseVectorStore, _store_name: str): + """Test vector search with exact match filter.""" + logger.info("=" * 20 + " EXACT MATCH FILTER SEARCH TEST " + "=" * 20) - # Test list filter (IN operation) - filters = {"node_type": ["tech", "tech_new"]} + # Test exact match filter + filters = {"node_type": "tech"} results = await store.search( query="What is artificial intelligence?", limit=5, filters=filters, ) - logger.info(f"Filtered search (node_type IN [tech, tech_new]) returned {len(results)} results") + logger.info(f"Filtered search (node_type=tech) returned {len(results)} results") for i, r in enumerate(results, 1): node_type = r.metadata.get("node_type") logger.info(f" Result {i}: type={node_type}, content={r.content[:50]}...") - assert node_type in ["tech", "tech_new"], "Result should have node_type in [tech, tech_new]" + assert node_type == "tech", "Result should have node_type='tech'" - logger.info("✓ List filter search test passed") + logger.info("✓ Exact match filter search test passed") async def test_search_with_multiple_filters(store: BaseVectorStore, _store_name: str): """Test vector search with multiple metadata filters (AND operation).""" logger.info("=" * 20 + " MULTIPLE FILTERS SEARCH TEST " + "=" * 20) - # Test multiple filters (AND operation) + # Test multiple exact match filters (AND operation) filters = { - "node_type": ["tech", "tech_new"], + "node_type": "tech", "source": "research", + "priority": "high", } results = await store.search( query="What is artificial intelligence?", @@ -387,14 +388,16 @@ async def test_search_with_multiple_filters(store: BaseVectorStore, _store_name: ) logger.info( - f"Multi-filter search (node_type IN [tech, tech_new] AND source=research) " f"returned {len(results)} results", + f"Multi-filter search (node_type=tech AND source=research AND priority=high) " f"returned {len(results)} results", ) for i, r in enumerate(results, 1): node_type = r.metadata.get("node_type") source = r.metadata.get("source") - logger.info(f" Result {i}: type={node_type}, source={source}, content={r.content[:40]}...") - assert node_type in ["tech", "tech_new"], "Result should have node_type in [tech, tech_new]" + priority = r.metadata.get("priority") + logger.info(f" Result {i}: type={node_type}, source={source}, priority={priority}") + assert node_type == "tech", "Result should have node_type='tech'" assert source == "research", "Result should have source='research'" + assert priority == "high", "Result should have priority='high'" logger.info("✓ Multiple filters search test passed") @@ -789,10 +792,9 @@ async def test_complex_metadata_queries(store: BaseVectorStore, _store_name: str await store.insert(complex_nodes) logger.info(f"✓ Inserted {len(complex_nodes)} nodes with complex metadata") - # Test 1: Multiple field filters with list values + # Test 1: Multiple exact match filters filters_1 = { "domain": "AI", - "year": ["2023", "2024"], "impact_factor": "high", } results_1 = await store.search( @@ -800,26 +802,25 @@ async def test_complex_metadata_queries(store: BaseVectorStore, _store_name: str limit=10, filters=filters_1, ) - logger.info(f"Test 1 - AI + high impact + recent years: {len(results_1)} results") + logger.info(f"Test 1 - AI + high impact: {len(results_1)} results") for r in results_1: assert r.metadata.get("domain") == "AI" assert r.metadata.get("impact_factor") == "high" - assert r.metadata.get("year") in ["2023", "2024"] - # Test 2: List filter with multiple subdomains + # Test 2: Single exact match filter filters_2 = { - "subdomain": ["nlp", "computer_vision"], + "subdomain": "nlp", } results_2 = await store.search( query="deep learning applications", limit=10, filters=filters_2, ) - logger.info(f"Test 2 - NLP or Computer Vision: {len(results_2)} results") + logger.info(f"Test 2 - NLP subdomain: {len(results_2)} results") for r in results_2: - assert r.metadata.get("subdomain") in ["nlp", "computer_vision"] + assert r.metadata.get("subdomain") == "nlp" - # Test 3: Year-based filtering + # Test 3: Year-based exact match filtering filters_3 = { "year": "2024", } @@ -1119,65 +1120,47 @@ async def test_filter_combinations(store: BaseVectorStore, _store_name: str): results_1 = await store.search(query="technology", filters={}, limit=10) logger.info(f"Test 1 - Empty filter: {len(results_1)} results") - # Test 2: Single value filter + # Test 2: Single exact match filter results_2 = await store.search( query="technology", filters={"node_type": "tech"}, limit=10, ) - logger.info(f"Test 2 - Single value filter: {len(results_2)} results") + logger.info(f"Test 2 - Single exact match filter: {len(results_2)} results") for r in results_2: assert r.metadata.get("node_type") == "tech" - # Test 3: List filter with single item + # Test 3: Multiple exact match filters (AND operation) results_3 = await store.search( - query="technology", - filters={"node_type": ["tech"]}, - limit=10, - ) - logger.info(f"Test 3 - List filter (single item): {len(results_3)} results") - - # Test 4: List filter with multiple items - results_4 = await store.search( - query="technology", - filters={"category": ["AI", "ML", "DL"]}, - limit=10, - ) - logger.info(f"Test 4 - List filter (multiple items): {len(results_4)} results") - for r in results_4: - assert r.metadata.get("category") in ["AI", "ML", "DL"] - - # Test 5: Multiple filters (AND operation) - results_5 = await store.search( query="technology", filters={ - "node_type": ["tech", "tech_new"], + "node_type": "tech", "source": "research", "priority": "high", }, limit=10, ) - logger.info(f"Test 5 - Multiple filters (AND): {len(results_5)} results") - for r in results_5: - assert r.metadata.get("node_type") in ["tech", "tech_new"] + logger.info(f"Test 3 - Multiple exact match filters (AND): {len(results_3)} results") + for r in results_3: + assert r.metadata.get("node_type") == "tech" assert r.metadata.get("source") == "research" assert r.metadata.get("priority") == "high" - # Test 6: Filter with non-existent value - results_6 = await store.search( + # Test 4: Filter with non-existent value + results_4 = await store.search( query="technology", filters={"category": "NON_EXISTENT_CATEGORY"}, limit=10, ) - logger.info(f"Test 6 - Non-existent filter value: {len(results_6)} results") - assert len(results_6) == 0, "Should return no results for non-existent filter value" + logger.info(f"Test 4 - Non-existent filter value: {len(results_4)} results") + assert len(results_4) == 0, "Should return no results for non-existent filter value" - # Test 7: List operation with filters + # Test 5: List operation with multiple exact match filters list_results = await store.list( filters={"node_type": "tech", "priority": "high"}, limit=20, ) - logger.info(f"Test 7 - List with filters: {len(list_results)} results") + logger.info(f"Test 5 - List with multiple filters: {len(list_results)} results") for r in list_results: assert r.metadata.get("node_type") == "tech" assert r.metadata.get("priority") == "high" @@ -1185,6 +1168,329 @@ async def test_filter_combinations(store: BaseVectorStore, _store_name: str): logger.info("✓ Filter combinations test passed") +async def test_range_query_filters(store: BaseVectorStore, _store_name: str): + """Test range query filters using the new [start, end] syntax.""" + logger.info("=" * 20 + " RANGE QUERY FILTERS TEST " + "=" * 20) + + # Clean up any existing test data first + try: + existing_nodes = await store.list(filters={"test_type": "range_query_test"}) + if existing_nodes: + await store.delete([node.vector_id for node in existing_nodes]) + logger.info(f"Cleaned up {len(existing_nodes)} existing test nodes") + except Exception as e: + logger.warning(f"Failed to clean up existing nodes: {e}") + + # Create test nodes with numeric metadata for range queries + import time + + base_timestamp = int(time.time()) + test_nodes = [] + + for i in range(20): + node = VectorNode( + vector_id=f"range_node_{i}", + content=f"Test content for range query node {i}", + metadata={ + "test_type": "range_query_test", + "timestamp": base_timestamp + i * 1000, # Each node is 1000 seconds apart + "rating": 50 + i * 2, # Ratings from 50 to 88 + "priority": i % 3, # 0, 1, or 2 + "category": ["tech", "science", "business"][i % 3], + }, + ) + test_nodes.append(node) + + # Insert test nodes + await store.insert(test_nodes) + logger.info(f"Inserted {len(test_nodes)} test nodes with numeric metadata") + + # Test 1: Range query on timestamp field + start_time = base_timestamp + 5000 + end_time = base_timestamp + 15000 + results_1 = await store.search( + query="test content", + limit=20, + filters={ + "timestamp": [start_time, end_time], # Range query: >= start_time AND <= end_time + }, + ) + logger.info(f"Test 1 - Timestamp range [{start_time}, {end_time}]: {len(results_1)} results") + + # Verify all results are within range + for r in results_1: + ts = r.metadata.get("timestamp") + assert ts >= start_time, f"Timestamp {ts} should be >= {start_time}" + assert ts <= end_time, f"Timestamp {ts} should be <= {end_time}" + logger.debug(f" Node {r.vector_id}: timestamp={ts}") + + # Expected nodes: range_node_5 to range_node_15 (11 nodes) + assert len(results_1) >= 10, f"Expected at least 10 results, got {len(results_1)}" + logger.info("✓ Timestamp range query validated") + + # Test 2: Range query on rating field + results_2 = await store.search( + query="test content", + limit=20, + filters={ + "rating": [60, 80], # Range query: rating >= 60 AND rating <= 80 + }, + ) + logger.info(f"Test 2 - Rating range [60, 80]: {len(results_2)} results") + + # Verify all results are within rating range + for r in results_2: + rating = r.metadata.get("rating") + assert rating >= 60, f"Rating {rating} should be >= 60" + assert rating <= 80, f"Rating {rating} should be <= 80" + logger.debug(f" Node {r.vector_id}: rating={rating}") + + # Expected: ratings from 60 to 80 (nodes 5-15) + assert len(results_2) >= 10, f"Expected at least 10 results, got {len(results_2)}" + logger.info("✓ Rating range query validated") + + # Test 3: Combine range query with exact match filter + results_3 = await store.search( + query="test content", + limit=20, + filters={ + "timestamp": [start_time, end_time], + "category": "tech", # Exact match + }, + ) + logger.info( + f"Test 3 - Timestamp range + exact match (category=tech): {len(results_3)} results", + ) + + # Verify filters + for r in results_3: + ts = r.metadata.get("timestamp") + category = r.metadata.get("category") + assert ts >= start_time and ts <= end_time, "Timestamp should be in range" + assert category == "tech", f"Category should be 'tech', got '{category}'" + logger.debug(f" Node {r.vector_id}: timestamp={ts}, category={category}") + + # Expected: nodes within range AND category=tech + assert len(results_3) >= 3, f"Expected at least 3 results, got {len(results_3)}" + logger.info("✓ Combined range + exact match query validated") + + # Test 4: Multiple range queries + results_4 = await store.search( + query="test content", + limit=20, + filters={ + "timestamp": [base_timestamp + 8000, base_timestamp + 12000], + "rating": [65, 75], + }, + ) + logger.info(f"Test 4 - Multiple range queries: {len(results_4)} results") + + # Verify both ranges + for r in results_4: + ts = r.metadata.get("timestamp") + rating = r.metadata.get("rating") + assert ts >= base_timestamp + 8000 and ts <= base_timestamp + 12000, "Timestamp out of range" + assert rating >= 65 and rating <= 75, f"Rating {rating} out of range [65, 75]" + logger.debug(f" Node {r.vector_id}: timestamp={ts}, rating={rating}") + + # Expected: nodes 8-12 (5 nodes) with overlapping ranges + assert len(results_4) >= 3, f"Expected at least 3 results, got {len(results_4)}" + logger.info("✓ Multiple range queries validated") + + # Test 5: Range query with list operation + results_5 = await store.list( + filters={ + "rating": [60, 70], + "test_type": "range_query_test", + }, + limit=20, + ) + logger.info(f"Test 5 - Range query in list operation: {len(results_5)} results") + + # Verify rating range in list results + for r in results_5: + rating = r.metadata.get("rating") + assert rating >= 60 and rating <= 70, f"Rating {rating} should be in range [60, 70]" + + logger.info("✓ Range query in list operation validated") + + # Test 6: Edge case - exact boundary values + results_6 = await store.list( + filters={ + "rating": [60, 60], # Exact match using range syntax + "test_type": "range_query_test", + }, + limit=20, + ) + logger.info(f"Test 6 - Exact value using range syntax [60, 60]: {len(results_6)} results") + + # Should return exactly one node (range_node_5 with rating=60) + for r in results_6: + rating = r.metadata.get("rating") + assert rating == 60, f"Rating should be exactly 60, got {rating}" + + logger.info("✓ Boundary value range query validated") + + # Test 7: Range query with sorting + results_7 = await store.list( + filters={ + "rating": [60, 80], + "test_type": "range_query_test", + }, + sort_key="rating", + reverse=True, + limit=5, + ) + logger.info(f"Test 7 - Range query with sorting: {len(results_7)} results") + + # Verify results are sorted and within range + for i in range(len(results_7) - 1): + rating1 = results_7[i].metadata.get("rating") + rating2 = results_7[i + 1].metadata.get("rating") + assert rating1 >= rating2, f"Results not sorted: {rating1} < {rating2}" + assert rating1 >= 60 and rating1 <= 80, "Rating out of range" + + logger.info("✓ Range query with sorting validated") + + # Clean up test data + await store.delete([node.vector_id for node in test_nodes]) + logger.info("Cleaned up test nodes") + + logger.info("✓ Range query filters test passed") + + +async def test_string_range_queries(store: BaseVectorStore, store_name: str): + """Test range queries with string values (e.g., date strings, timestamps).""" + logger.info("=" * 20 + " STRING RANGE QUERIES TEST " + "=" * 20) + + # Skip this test for stores that don't support string range queries properly + # Qdrant and ChromaDB only support numeric range queries, not string range queries + if store_name not in ["PGVectorStore", "LocalVectorStore", "ESVectorStore"]: + logger.info(f"Skipping string range query test for {store_name}") + return + + # Clean up any existing test data first + try: + existing_nodes = await store.list(filters={"test_type": "string_range_test"}) + if existing_nodes: + await store.delete([node.vector_id for node in existing_nodes]) + logger.info(f"Cleaned up {len(existing_nodes)} existing test nodes") + except Exception as e: + logger.warning(f"Failed to clean up existing nodes: {e}") + + # Create test nodes with string date metadata + test_nodes = [] + dates = [ + "2024-01-01", + "2024-01-15", + "2024-02-01", + "2024-02-15", + "2024-03-01", + "2024-03-15", + "2024-04-01", + ] + + for i, date in enumerate(dates): + node = VectorNode( + vector_id=f"string_range_node_{i}", + content=f"Test content for date {date}", + metadata={ + "test_type": "string_range_test", + "date": date, + "index": i, + }, + ) + test_nodes.append(node) + + # Insert test nodes + await store.insert(test_nodes) + logger.info(f"Inserted {len(test_nodes)} test nodes with string dates") + + # Test 1: String range query on date field + try: + results = await store.search( + query="test content", + limit=20, + filters={ + "date": ["2024-02-01", "2024-03-15"], # Range query on string dates + }, + ) + logger.info(f"Test 1 - String date range ['2024-02-01', '2024-03-15']: {len(results)} results") + + # Verify all results are within range + expected_dates = ["2024-02-01", "2024-02-15", "2024-03-01", "2024-03-15"] + for r in results: + date = r.metadata.get("date") + assert date >= "2024-02-01", f"Date {date} should be >= '2024-02-01'" + assert date <= "2024-03-15", f"Date {date} should be <= '2024-03-15'" + logger.debug(f" Node {r.vector_id}: date={date}") + + assert len(results) >= 3, f"Expected at least 3 results, got {len(results)}" + logger.info("✓ String range query validated") + except Exception as e: + # For PGVector, this might fail on older implementations + if "PGVector" in store_name: + logger.warning(f"String range query failed for PGVector (expected if not updated): {e}") + else: + raise + + # Clean up test data + await store.delete([node.vector_id for node in test_nodes]) + logger.info("Cleaned up test nodes") + + logger.info("✓ String range queries test passed") + + +async def test_sql_injection_protection(store: BaseVectorStore, store_name: str): + """Test SQL injection protection in filter keys and collection names.""" + logger.info("=" * 20 + " SQL INJECTION PROTECTION TEST " + "=" * 20) + + # This test is only relevant for SQL-based stores + if store_name not in ["PGVectorStore"]: + logger.info(f"Skipping SQL injection test for {store_name}") + return + + # Test 1: Invalid collection name (SQL injection attempt) + try: + from reme_ai.core.vector_store import PGVectorStore + from reme_ai.core.embedding import OpenAIEmbeddingModel + + embedding_model = OpenAIEmbeddingModel() + + # This should raise ValueError due to invalid table name + try: + invalid_store = PGVectorStore( + collection_name="test'; DROP TABLE users; --", + embedding_model=embedding_model, + ) + logger.error("❌ FAILED: Invalid collection name was accepted (SQL injection risk!)") + assert False, "Should have raised ValueError for invalid collection name" + except ValueError as e: + logger.info(f"✓ Invalid collection name rejected: {e}") + + # Test 2: Invalid metadata key in filters + try: + results = await store.search( + query="test", + filters={ + "normal_key": "value", + "bad'; DROP TABLE users; --": "value", + }, + ) + logger.error("❌ FAILED: Invalid metadata key was accepted (SQL injection risk!)") + assert False, "Should have raised ValueError for invalid metadata key" + except ValueError as e: + logger.info(f"✓ Invalid metadata key rejected: {e}") + + logger.info("✓ SQL injection protection validated") + + except Exception as e: + logger.error(f"SQL injection protection test failed: {e}") + raise + + logger.info("✓ SQL injection protection test passed") + + async def test_list_with_sorting(store: BaseVectorStore, _store_name: str): """Test list operation with sorting by timestamp to get most recent top 10 items.""" logger.info("=" * 20 + " LIST WITH SORTING TEST " + "=" * 20) @@ -1353,7 +1659,7 @@ async def run_all_tests_for_store(store_type: str, store_name: str): await test_insert(store, store_name) await test_search(store, store_name) await test_search_with_single_filter(store, store_name) - await test_search_with_list_filter(store, store_name) + await test_search_with_exact_match_filter(store, store_name) await test_search_with_multiple_filters(store, store_name) await test_get_by_id(store, store_name) await test_list_all(store, store_name) @@ -1374,6 +1680,9 @@ async def run_all_tests_for_store(store_type: str, store_name: str): await test_metadata_statistics(store, store_name) await test_update_metadata_only(store, store_name) await test_filter_combinations(store, store_name) + await test_range_query_filters(store, store_name) + await test_string_range_queries(store, store_name) + await test_sql_injection_protection(store, store_name) await test_list_with_sorting(store, store_name) # ========== Collection Management Tests ========== From 08b771b6c45bcbfd22b307ea3ef34d2d1d4ebbb5 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sun, 18 Jan 2026 15:50:46 +0800 Subject: [PATCH 02/19] feat(mem-agent): introduce version 4 memory agents and tools --- bench/halumem/compute_qa_stats_v4.py | 327 +++++++++ bench/halumem/compute_stats_from_tmp.py | 2 +- bench/halumem/eval_reme_simple_v3.py | 8 + bench/halumem/eval_reme_simple_v4.py | 671 ++++++++++++++++++ bench/halumem/halumem.yaml | 40 +- reme_ai/core/llm/base_llm.py | 4 +- reme_ai/core/schema/tool_call.py | 37 + reme_ai/mem_agent/base_memory_agent.py | 7 +- .../mem_agent/v3/personal_summarizer_v3.yaml | 4 + reme_ai/mem_agent/v3/reme_retriever_v3.yaml | 11 +- reme_ai/mem_agent/v4/__init__.py | 11 + reme_ai/mem_agent/v4/personal_retriever_v4.py | 46 ++ .../mem_agent/v4/personal_retriever_v4.yaml | 37 + .../mem_agent/v4/personal_summarizer_v4.py | 127 ++++ .../mem_agent/v4/personal_summarizer_v4.yaml | 45 ++ reme_ai/mem_agent/v4/reme_retriever_v4.py | 53 ++ reme_ai/mem_agent/v4/reme_retriever_v4.yaml | 25 + reme_ai/mem_agent/v4/reme_summarizer_v4.py | 63 ++ reme_ai/mem_agent/v4/reme_summarizer_v4.yaml | 26 + reme_ai/mem_tool/base_memory_tool.py | 14 +- reme_ai/mem_tool/v3/add_memory.py | 2 +- reme_ai/mem_tool/v3/read_user_profile.py | 2 +- reme_ai/mem_tool/v3/update_user_profile.py | 51 +- reme_ai/mem_tool/v4/__init__.py | 15 + reme_ai/mem_tool/v4/add_summary_memory.py | 63 ++ reme_ai/mem_tool/v4/hands_off.py | 112 +++ reme_ai/mem_tool/v4/read_history.py | 38 + reme_ai/mem_tool/v4/read_user_profile.py | 62 ++ reme_ai/mem_tool/v4/retrieve_memory.py | 103 +++ reme_ai/mem_tool/v4/update_user_profile.py | 105 +++ reme_ai/reme.py | 106 ++- 31 files changed, 2143 insertions(+), 74 deletions(-) create mode 100644 bench/halumem/compute_qa_stats_v4.py create mode 100644 bench/halumem/eval_reme_simple_v4.py create mode 100644 reme_ai/mem_agent/v4/__init__.py create mode 100644 reme_ai/mem_agent/v4/personal_retriever_v4.py create mode 100644 reme_ai/mem_agent/v4/personal_retriever_v4.yaml create mode 100644 reme_ai/mem_agent/v4/personal_summarizer_v4.py create mode 100644 reme_ai/mem_agent/v4/personal_summarizer_v4.yaml create mode 100644 reme_ai/mem_agent/v4/reme_retriever_v4.py create mode 100644 reme_ai/mem_agent/v4/reme_retriever_v4.yaml create mode 100644 reme_ai/mem_agent/v4/reme_summarizer_v4.py create mode 100644 reme_ai/mem_agent/v4/reme_summarizer_v4.yaml create mode 100644 reme_ai/mem_tool/v4/__init__.py create mode 100644 reme_ai/mem_tool/v4/add_summary_memory.py create mode 100644 reme_ai/mem_tool/v4/hands_off.py create mode 100644 reme_ai/mem_tool/v4/read_history.py create mode 100644 reme_ai/mem_tool/v4/read_user_profile.py create mode 100644 reme_ai/mem_tool/v4/retrieve_memory.py create mode 100644 reme_ai/mem_tool/v4/update_user_profile.py diff --git a/bench/halumem/compute_qa_stats_v4.py b/bench/halumem/compute_qa_stats_v4.py new file mode 100644 index 00000000..7cdee50a --- /dev/null +++ b/bench/halumem/compute_qa_stats_v4.py @@ -0,0 +1,327 @@ +""" +Compute Question Answering statistics from eval_reme_simple_v4.py results. + +This script processes the output from eval_reme_simple_v4.py and computes +comprehensive QA metrics. + +Usage: + python bench/halumem/compute_qa_stats_v4.py --results_file bench_results/reme_simple_v4/eval_results.jsonl +""" + +import json +import os +from pathlib import Path +from typing import Any + +from loguru import logger + + +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 + + +def compute_time_metrics(results_file: str) -> dict[str, float]: + """Compute timing metrics from evaluation results.""" + add_duration = 0 + 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) + + # 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 + } + + +def load_from_tmp_dir(tmp_dir: str) -> tuple[str, list[dict]]: + """Load data from tmp directory and generate eval_results.jsonl file.""" + tmp_path = Path(tmp_dir) + parent_dir = tmp_path.parent + eval_results_file = parent_dir / "eval_results.jsonl" + + print(f"\n📁 Loading data from tmp directory: {tmp_dir}") + print(f"📝 Will generate: {eval_results_file}") + + # Collect all user directories + user_dirs = [d for d in tmp_path.iterdir() if d.is_dir()] + print(f" Found {len(user_dirs)} user directories") + + users_data = [] + + for user_dir in user_dirs: + user_name = user_dir.name + + # Load all session files for this user (sorted by session number) + 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]) # Sort by session number + ) + + if not session_files: + print(f" ⚠️ No session files found for user: {user_name}") + continue + + # Load first session to get user metadata + 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": [] + } + + # Load all sessions + for session_file in session_files: + with open(session_file, "r", encoding="utf-8") as f: + session_data = json.load(f) + # Remove redundant user metadata + session_data.pop("uuid", None) + session_data.pop("user_name", None) + user_data["sessions"].append(session_data) + + users_data.append(user_data) + print(f" ✓ Loaded user {user_name}: {len(session_files)} sessions") + + # Write to eval_results.jsonl + 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), users_data + + +def main(input_path: str): + """Main function to compute statistics from eval results.""" + + if not os.path.exists(input_path): + logger.error(f"Input path not found: {input_path}") + return + + print("\n" + "=" * 80) + print("COMPUTING QUESTION ANSWERING STATISTICS - REME V4") + print("=" * 80) + + # Determine if input is a directory (tmp) or file (eval_results.jsonl) + if os.path.isdir(input_path): + results_file, users_data = load_from_tmp_dir(input_path) + else: + results_file = input_path + users_data = None + print(f"\n📁 Using existing results file: {results_file}") + + # Collect all QA records with metadata + qa_records = [] + qa_records_with_metadata = [] # Store records with user/session/question info + user_count = 0 + 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 # Track the index of valid (non-skipped) sessions + 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", {}) + session_qa_records = eval_results.get("question_answering_records", []) + + # Add records with metadata + for qa_idx, qa in enumerate(session_qa_records): + qa_records.append(qa) + qa_records_with_metadata.append({ + "user_name": user_name, + "session_idx": valid_session_idx, + "original_session_idx": original_idx, + "question_idx": qa_idx, + "qa_record": qa + }) + + valid_session_idx += 1 + + print(f"\n📊 Data loaded:") + print(f" Users: {user_count}") + print(f" Sessions: {session_count}") + print(f" QA Records: {len(qa_records)}") + + # Compute metrics + print("\n🔄 Computing metrics...") + qa_metrics = compute_qa_metrics(qa_records) + time_metrics = 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 + output_dir = Path(results_file).parent + report_file = output_dir / "reme_eval_stat_result.json" + + with open(report_file, "w", encoding="utf-8") as f: + json.dump(final_results, f, ensure_ascii=False, indent=4) + + print(f"\n✅ Statistics saved to: {report_file}") + + # Print summary + print("\n" + "=" * 80) + print("EVALUATION SUMMARY - REME V4") + print("=" * 80) + + print("\n📊 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 non-Correct QA records + print("\n" + "=" * 80) + print("NON-CORRECT QA RECORDS") + print("=" * 80) + + non_correct_records = [ + record for record in qa_records_with_metadata + if record["qa_record"].get("result_type") not in ["Correct", ""] + ] + + if non_correct_records: + print(f"\nFound {len(non_correct_records)} non-correct records:\n") + for record in non_correct_records: + user_name = record["user_name"] + session_idx = record["session_idx"] + original_idx = record["original_session_idx"] + question_idx = record["question_idx"] + qa = record["qa_record"] + result_type = qa.get("result_type", "Unknown") + question = qa.get("question", "N/A") + answer = qa.get("answer", "N/A") + + print(f"👤 User: {user_name}") + print(f"📅 Session: {original_idx} (valid session index: {session_idx})") + print(f"❓ Question #{question_idx}") + print(f"🏷️ Result Type: {result_type}") + print(f"💬 Question: {question}") + print(f"💡 Answer: {answer}") + print("-" * 80) + else: + print("\n✅ All QA records are Correct!") + + 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, + required=False, + help="Path to eval_results.jsonl file (e.g., bench_results/reme_simple_v4/eval_results.jsonl)" + ) + parser.add_argument( + "--tmp_dir", + type=str, + required=False, + help="Path to tmp directory (e.g., bench_results/reme_simple_v4/tmp)" + ) + + args = parser.parse_args() + + # Determine input path + if args.tmp_dir: + input_path = args.tmp_dir + elif args.results_file: + input_path = args.results_file + else: + parser.error("Either --results_file or --tmp_dir must be provided") + + main(input_path=input_path) diff --git a/bench/halumem/compute_stats_from_tmp.py b/bench/halumem/compute_stats_from_tmp.py index 3891c21a..5ab401ba 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 diff --git a/bench/halumem/eval_reme_simple_v3.py b/bench/halumem/eval_reme_simple_v3.py index 62b1fe64..7c0dfd8e 100644 --- a/bench/halumem/eval_reme_simple_v3.py +++ b/bench/halumem/eval_reme_simple_v3.py @@ -17,6 +17,7 @@ import asyncio import json import os import re +import shutil import time from dataclasses import dataclass from datetime import datetime, timezone @@ -497,6 +498,13 @@ class HaluMemEvaluatorV3: # 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] diff --git a/bench/halumem/eval_reme_simple_v4.py b/bench/halumem/eval_reme_simple_v4.py new file mode 100644 index 00000000..c17622a6 --- /dev/null +++ b/bench/halumem/eval_reme_simple_v4.py @@ -0,0 +1,671 @@ +""" +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 +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[str, list, float]: + """ + Search memory using ReMe and return response. + + Returns: + tuple: (response, agent_messages, duration_ms) + """ + start = time.time() + + response, agent_messages, success = await self.reme.retrieve_v4( + 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: + 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 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/halumem.yaml b/bench/halumem/halumem.yaml index 60f7cbda..cd7a2e6f 100644 --- a/bench/halumem/halumem.yaml +++ b/bench/halumem/halumem.yaml @@ -424,10 +424,7 @@ EVALUATION_PROMPT_FOR_QUESTION: | 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 +432,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 +493,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/reme_ai/core/llm/base_llm.py b/reme_ai/core/llm/base_llm.py index f83a39a7..58db4b3f 100644 --- a/reme_ai/core/llm/base_llm.py +++ b/reme_ai/core/llm/base_llm.py @@ -68,7 +68,9 @@ class BaseLLM(ABC): if tool.name not in tool_dict: continue - if not tool.check_argument(): + # First try sanitizing arguments + if not tool.sanitize_and_check_argument(): + logger.error(f"Tool call {tool.name} has invalid JSON arguments after sanitization attempt: {tool.arguments}") raise ValueError(f"Tool call {tool.name} has invalid JSON arguments: {tool.arguments}") validated_tools.append(tool.simple_output_dump()) diff --git a/reme_ai/core/schema/tool_call.py b/reme_ai/core/schema/tool_call.py index 035c63f1..21dde495 100644 --- a/reme_ai/core/schema/tool_call.py +++ b/reme_ai/core/schema/tool_call.py @@ -177,6 +177,43 @@ class ToolCall(BaseModel): 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 def simple_output_dump(self) -> dict: """Convert ToolCall to output format dictionary for API responses.""" diff --git a/reme_ai/mem_agent/base_memory_agent.py b/reme_ai/mem_agent/base_memory_agent.py index fb52dc0d..15c6dc04 100644 --- a/reme_ai/mem_agent/base_memory_agent.py +++ b/reme_ai/mem_agent/base_memory_agent.py @@ -152,7 +152,7 @@ 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") + logger.info(f"[{self.__class__.__name__}] step{step + 1}.{j} join tool_result={tool_result[:2000]}...\n\n") return tool_result_messages async def react(self, messages: list[Message]): @@ -209,3 +209,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/v3/personal_summarizer_v3.yaml b/reme_ai/mem_agent/v3/personal_summarizer_v3.yaml index 97d18afc..95907612 100644 --- a/reme_ai/mem_agent/v3/personal_summarizer_v3.yaml +++ b/reme_ai/mem_agent/v3/personal_summarizer_v3.yaml @@ -17,6 +17,10 @@ system_prompt: | ### 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) diff --git a/reme_ai/mem_agent/v3/reme_retriever_v3.yaml b/reme_ai/mem_agent/v3/reme_retriever_v3.yaml index 8ff005ee..7a3e575f 100644 --- a/reme_ai/mem_agent/v3/reme_retriever_v3.yaml +++ b/reme_ai/mem_agent/v3/reme_retriever_v3.yaml @@ -5,13 +5,13 @@ tool: | NEVER hallucinate or fabricate information not present in retrieved memories. system_prompt: | - You are a memory retrieval agent. Search for relevant memories to answer the user's question following this strategy: + 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 Context + ## User's Question {context} ## Three-Step Retrieval Strategy @@ -19,7 +19,6 @@ system_prompt: | **STEP 1: Read User Profile (REQUIRED FIRST)** - Use `read_user_profile` with memory_type and memory_target from available meta memories - Check if the user profile directly answers the question - - If sufficient information found, provide the answer and STOP **STEP 2: Vector Search (If Step 1 insufficient)** - Use `retrieve_memory` with memory_type, memory_target, and query @@ -44,10 +43,8 @@ system_prompt: | - Try multiple history_id entries if needed ## Response Rules - - Answer ONLY based on retrieved information - NEVER guess or fabricate - - If nothing found after all three steps: State clearly "I don't know. I cannot find relevant information to answer this question." + - 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 - - Once you find sufficient information, provide a direct answer user_message: | - Retrieve relevant memories and answer the question using the three-step strategy. + Answer the question using the three-step strategy. 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..2b33897e --- /dev/null +++ b/reme_ai/mem_agent/v4/personal_retriever_v4.py @@ -0,0 +1,46 @@ +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]: + 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`") + + read_profile_tool = ReadUserProfile() + await read_profile_tool.call(memory_type=self.memory_type.value, memory_target=self.memory_target) + + messages = [ + Message( + role=Role.SYSTEM, + content=self.prompt_format( + prompt_name="system_prompt", + memory_type=self.memory_type.value, + memory_target=self.memory_target, + user_profile=read_profile_tool.output, + context=context, + )), + 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, + memory_type=self.memory_type.value, + memory_target=self.memory_target, + **kwargs, + ) 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..cc5dcc34 --- /dev/null +++ b/reme_ai/mem_agent/v4/personal_retriever_v4.yaml @@ -0,0 +1,37 @@ +tool: | + Retrieve relevant personal memories to answer user questions through vector search and history reading. + +system_prompt: | + You are a memory agent managing **{memory_type}** memories about **{memory_target}**. + + ## User Profile + {user_profile} + + ## Question + {context} + + ## Retrieval Strategy + + **Tool 1: Vector Search (`retrieve_memory`) + - Try at least 3-5 different queries before moving to next tool: + * Direct question + * Reformulated phrasings + * Entity-focused queries + * Different keyword combinations + - If no results: retry with different time ranges or remove time constraints: [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 this ONLY after completing multiple retrieve_memory attempts + - Use history_id from retrieved memories to read original conversations + - Prioritize most relevant or recent entries + - Read multiple if needed for complete context + + **Response** + - **CRITICAL: Answer ONLY based on retrieved memories. Do NOT hallucinate or infer information not present in the search results.** + - Try multiple angles before giving up + - State "nothing found after thorough search" if nothing found after thorough search + +user_message: | + Answer the question using the retrieval strategy. 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..4b867bed --- /dev/null +++ b/reme_ai/mem_agent/v4/personal_summarizer_v4.py @@ -0,0 +1,127 @@ +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.SYSTEM, + content=self.prompt_format( + prompt_name="system_prompt_phase1", + context=history_node.content, + memory_type=self.memory_type.value, + memory_target=self.memory_target, + )), + Message( + role=Role.USER, + content=self.get_prompt("user_message_phase1"), + ), + ] + 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.SYSTEM, + content=self.prompt_format( + prompt_name="system_prompt_phase2", + context=history_node.content, + memory_type=self.memory_type.value, + memory_target=self.memory_target, + user_profile=user_profile, + )), + Message( + role=Role.USER, + content=self.get_prompt("user_message_phase2"), + ), + ] + return messages + + 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, + 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__}] 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__}] phase1.step0.{i} {message.role} " + f"{message.simple_dump(enable_json_dump=True)}", + ) + + messages_phase1, success_phase1 = await self.react(messages_phase1) + if not success_phase1: + logger.warning(f"[{self.__class__.__name__}] Phase 1 did not complete successfully") + + # Phase 2: Read user profile and UpdateUserProfile + logger.info(f"[{self.__class__.__name__}] 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 + logger.info(f"[{self.__class__.__name__}] Loading user profile with ReadUserProfile") + await read_profile_tool.call(memory_type=self.memory_type.value, memory_target=self.memory_target) + user_profile = str(read_profile_tool.output) + logger.info(f"[{self.__class__.__name__}] User profile loaded: {user_profile}...") + else: + logger.warning(f"[{self.__class__.__name__}] 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__}] phase2.step0.{i} {message.role} " + f"{message.simple_dump(enable_json_dump=True)}", + ) + + messages_phase2, success_phase2 = await self.react(messages_phase2) + + # 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..41d3ac61 --- /dev/null +++ b/reme_ai/mem_agent/v4/personal_summarizer_v4.yaml @@ -0,0 +1,45 @@ +tool: | + Extract and update personal memories about the user from conversation context. + Identify preferences, habits, background, relationships, and key facts. + +system_prompt_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) + +user_message_phase1: | + Extract personal memories from the conversation using `AddSummaryMemory`. + +# capturing complete contexts with preconditions, causes, and consequences +system_prompt_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/memories with new information from the conversation, including **{memory_target}**' current status: + - `profile_ids_to_delete`: Remove outdated, conflicting, or redundant entries. + - `profiles_to_add`: Add new profiles/memories with `conversation_time`, e.g. `YYYY-MM-DD HH:MM:SS`, {memory_target} did something. + - Maintain profiles that are concise, mutually exclusive, and collectively comprehensive with no information loss. + +user_message_phase2: | + 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..b8d5fbbe --- /dev/null +++ b/reme_ai/mem_agent/v4/reme_retriever_v4.py @@ -0,0 +1,53 @@ +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 [] + + 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]: + return await super()._acting_step( + assistant_message, + step, + query=self.context.get("query", ""), + messages=self.context.get("messages", []), + **kwargs, + ) 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..a6d322ef --- /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_tool/base_memory_tool.py b/reme_ai/mem_tool/base_memory_tool.py index 0d92671c..8b124496 100644 --- a/reme_ai/mem_tool/base_memory_tool.py +++ b/reme_ai/mem_tool/base_memory_tool.py @@ -80,11 +80,21 @@ class BaseMemoryTool(BaseOp, metaclass=ABCMeta): """Get the reference memory ID from context.""" return self.context.get("ref_memory_id", "") + @property + def description(self) -> str: + """Get the description from context.""" + return self.context.get("description", "") + @property def messages_formated(self) -> str: """Get the formated messages from context.""" return self.context.get("messages_formated", "") + @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.""" @@ -115,8 +125,4 @@ class BaseMemoryTool(BaseOp, metaclass=ABCMeta): 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/v3/add_memory.py b/reme_ai/mem_tool/v3/add_memory.py index a013db42..ee488639 100644 --- a/reme_ai/mem_tool/v3/add_memory.py +++ b/reme_ai/mem_tool/v3/add_memory.py @@ -28,7 +28,7 @@ class AddMemory(BaseMemoryTool): "description": "memory content", }, "conversation_time": { - "type": "object", + "type": "string", "description": "conversation time, e.g. '2020-01-01 00:00:00'", } }, diff --git a/reme_ai/mem_tool/v3/read_user_profile.py b/reme_ai/mem_tool/v3/read_user_profile.py index 57cff042..3dba2bcf 100644 --- a/reme_ai/mem_tool/v3/read_user_profile.py +++ b/reme_ai/mem_tool/v3/read_user_profile.py @@ -49,7 +49,7 @@ class ReadUserProfile(BaseMemoryTool): # Convert to MemoryNode objects and sort by conversation_time (oldest first) memory_nodes = [MemoryNode(**node_data) for node_data in cached_data] memory_nodes.sort( - key=lambda node: node.metadata.get("conversation_time", "") + key=lambda n: n.metadata.get("conversation_time", "") ) memory_formated = [] diff --git a/reme_ai/mem_tool/v3/update_user_profile.py b/reme_ai/mem_tool/v3/update_user_profile.py index 56ad4284..46879e2b 100644 --- a/reme_ai/mem_tool/v3/update_user_profile.py +++ b/reme_ai/mem_tool/v3/update_user_profile.py @@ -1,42 +1,43 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C from ...core.schema.memory_node import MemoryNode -@C.register_op() class UpdateUserProfile(BaseMemoryTool): def __init__(self, **kwargs): kwargs["enable_multiple"] = True super().__init__(**kwargs) + def _build_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": self.get_prompt("profile_ids_to_delete"), + "description": "profile_ids_to_delete", "items": {"type": "string"}, }, "profiles_to_add": { "type": "array", - "description": self.get_prompt("profiles_to_add"), + "description": "profiles_to_add", "items": { "type": "object", "properties": { "profile_content": { "type": "string", - "description": self.get_prompt("profile_content"), + "description": "profile_content", }, - "timestamp": { + "conversation_time": { "type": "string", - "description": self.get_prompt("timestamp"), + "description": "conversation_time, e.g. '2020-01-01 00:00:00'", }, }, - "required": ["profile_content", "timestamp"], + "required": ["profile_content", "conversation_time"], }, }, }, @@ -44,22 +45,16 @@ class UpdateUserProfile(BaseMemoryTool): } async def execute(self): - memory_type = "personal" - memory_target = self.memory_target - assert memory_target, "memory_target is not configured." - - cache_key = f"{memory_type}_{memory_target}" - profile_ids_to_delete = self.context.get("profile_ids_to_delete", []) profile_ids_to_delete = [m for m in profile_ids_to_delete if m] profile_ids_to_delete = list(dict.fromkeys(profile_ids_to_delete)) - profiles_to_add = self.context.get("profiles_to_add", []) if not profile_ids_to_delete and not profiles_to_add: self.output = "No memories to remove or add. Operation has been done." return + 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: @@ -70,9 +65,7 @@ class UpdateUserProfile(BaseMemoryTool): 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 - ] + 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.") @@ -80,24 +73,13 @@ class UpdateUserProfile(BaseMemoryTool): if profiles_to_add: for mem in profiles_to_add: profile_content = mem.get("profile_content", "") - timestamp = mem.get("timestamp", "") - - if not profile_content: - logger.warning("Skipping memory with empty content") - continue - - memory_node = self._build_memory_node( + conversation_time = mem.get("conversation_time", "") + new_memory_nodes.append(self._build_memory_node( memory_content=profile_content, when_to_use="", - metadata={"timestamp": timestamp} - ) - memory_node.memory_type = MemoryNode.MemoryType.PERSONAL - memory_node.memory_target = memory_target - - new_memory_nodes.append(memory_node) - - added_count = len(new_memory_nodes) - logger.info(f"Added {added_count} new memories to user profile.") + 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 @@ -114,5 +96,4 @@ class UpdateUserProfile(BaseMemoryTool): 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..54b61f6d --- /dev/null +++ b/reme_ai/mem_tool/v4/hands_off.py @@ -0,0 +1,112 @@ +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] = [] + + @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) + + 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/read_history.py b/reme_ai/mem_tool/v4/read_history.py new file mode 100644 index 00000000..e9ab2a15 --- /dev/null +++ b/reme_ai/mem_tool/v4/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/v4/read_user_profile.py b/reme_ai/mem_tool/v4/read_user_profile.py new file mode 100644 index 00000000..b4aae557 --- /dev/null +++ b/reme_ai/mem_tool/v4/read_user_profile.py @@ -0,0 +1,62 @@ +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 = False, **kwargs): + kwargs["enable_multiple"] = False + super().__init__(**kwargs) + self.add_memory_type_target = add_memory_type_target + + def _build_tool_description(self) -> str: + return "Read user profile." + + 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}".replace(" ", "_").lower() + cached_data = self.meta_memory.load(cache_key, auto_clean=False) + + if not cached_data: + self.output = "" + logger.info(f"empty cached_data={cache_key}") + return + + 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 and node.metadata["conversation_time"]: + 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/v4/retrieve_memory.py b/reme_ai/mem_tool/v4/retrieve_memory.py new file mode 100644 index 00000000..e7f77f96 --- /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 = "\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/v4/update_user_profile.py b/reme_ai/mem_tool/v4/update_user_profile.py new file mode 100644 index 00000000..085daf80 --- /dev/null +++ b/reme_ai/mem_tool/v4/update_user_profile.py @@ -0,0 +1,105 @@ +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.schema.memory_node import MemoryNode +from ...core.utils import deduplicate_memories + + +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": { + "conversation_time": { + "type": "string", + "description": "conversation_time, e.g. '2020-01-01 00:00:00'", + }, + "profile_content": { + "type": "string", + "description": "profile_content", + }, + }, + "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 profiles to remove or add. Operation has been done." + return + + cache_key = f"{self.memory_type}_{self.memory_target}".replace(" ", "_").lower() + cached_data = self.meta_memory.load(cache_key, auto_clean=False) + if cached_data: + existing_memory_nodes = [MemoryNode(**node_data) for node_data in cached_data] + else: + existing_memory_nodes = [] + + removed_count = 0 + if profile_ids_to_delete: + original_count = len(existing_memory_nodes) + existing_memory_nodes = [n for n in existing_memory_nodes if n.memory_id not in profile_ids_to_delete] + removed_count = original_count - len(existing_memory_nodes) + logger.info(f"Removed {removed_count} profiles.") + + added_count = 0 + new_memory_nodes = [] + if profiles_to_add: + for mem in profiles_to_add: + memory_node = MemoryNode( + memory_type=self.memory_type, + memory_target=self.memory_target, + when_to_use="", + content=mem.get("profile_content", ""), + ref_memory_id=self.ref_memory_id, + author=self.author, + metadata={"conversation_time": mem.get("conversation_time", "")}, + ) + new_memory_nodes.append(memory_node) + added_count = len(new_memory_nodes) + logger.info(f"Added {added_count} new profiles.") + + updated_memory_nodes = deduplicate_memories(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 profiles") + if added_count > 0: + operations.append(f"added {added_count} new profiles") + + 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/reme.py b/reme_ai/reme.py index ff9d8822..02485159 100644 --- a/reme_ai/reme.py +++ b/reme_ai/reme.py @@ -18,6 +18,12 @@ from .mem_agent.v3 import ( ReMeRetrieverV3, ReMeSummarizerV3, ) +from .mem_agent.v4 import ( + PersonalSummarizerV4, + PersonalRetrieverV4, + ReMeRetrieverV4, + ReMeSummarizerV4, +) from .mem_tool import ( HandsOffTool, ReadHistoryMemory, @@ -42,6 +48,14 @@ from .mem_tool.v3 import ( 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 @@ -348,9 +362,9 @@ class ReMe(Application): personal_summarizer_v3 = PersonalSummarizerV3( tools=[ - AddMemoryV3(), - ReadUserProfile(add_memory_type_target=False), - UpdateUserProfile(), + AddMemoryV3(enable_thinking_params=True), + ReadUserProfile(enable_thinking_params=True, add_memory_type_target=False), + UpdateUserProfile(enable_thinking_params=True), ], ) @@ -394,9 +408,9 @@ class ReMe(Application): reme_retriever_v3 = ReMeRetrieverV3( meta_memories=meta_memories, tools=[ - ReadUserProfile(add_memory_type_target=True), - RetrieveMemory(top_k=top_k), - ReadHistoryV3(), + ReadUserProfile(enable_thinking_params=True, add_memory_type_target=True), + RetrieveMemory(enable_thinking_params=True, top_k=top_k), + ReadHistoryV3(enable_thinking_params=True), ], ) @@ -409,3 +423,83 @@ class ReMe(Application): 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.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.messages, reme_retriever_v4.success + + else: + raise NotImplementedError From 74c1386a697bc2e9f767cdb8b1bedfc8fab9ba3c Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Mon, 19 Jan 2026 00:46:52 +0800 Subject: [PATCH 03/19] refactor(mem_agent): optimize agent execution and enhance evaluation pipeline --- bench/halumem/compute_qa_stats_v4.py | 220 +++++++----------- bench/halumem/eval_reme_simple_v4.py | 37 ++- bench/halumem/eval_tools.py | 37 +++ bench/halumem/halumem.yaml | 24 ++ reme_ai/core/config/default.yaml | 1 + reme_ai/mem_agent/base_memory_agent.py | 34 ++- reme_ai/mem_agent/v4/personal_retriever_v4.py | 44 ++-- .../mem_agent/v4/personal_retriever_v4.yaml | 21 +- .../mem_agent/v4/personal_summarizer_v4.py | 47 ++-- .../mem_agent/v4/personal_summarizer_v4.yaml | 20 +- reme_ai/mem_agent/v4/reme_retriever_v4.py | 78 ++++++- reme_ai/mem_tool/v4/hands_off.py | 3 + reme_ai/mem_tool/v4/read_history.py | 2 +- reme_ai/mem_tool/v4/read_user_profile.py | 35 ++- reme_ai/mem_tool/v4/retrieve_memory.py | 2 +- reme_ai/mem_tool/v4/update_user_profile.py | 2 +- reme_ai/reme.py | 4 +- 17 files changed, 374 insertions(+), 237 deletions(-) diff --git a/bench/halumem/compute_qa_stats_v4.py b/bench/halumem/compute_qa_stats_v4.py index 7cdee50a..c7bdfab2 100644 --- a/bench/halumem/compute_qa_stats_v4.py +++ b/bench/halumem/compute_qa_stats_v4.py @@ -1,11 +1,9 @@ """ Compute Question Answering statistics from eval_reme_simple_v4.py results. -This script processes the output from eval_reme_simple_v4.py and computes -comprehensive QA metrics. - 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 @@ -13,8 +11,6 @@ import os from pathlib import Path from typing import Any -from loguru import logger - def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: """Compute question answering metrics.""" @@ -31,51 +27,37 @@ def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: "qa_num": 0 } - correct = 0 - hallucination = 0 - omission = 0 - valid = 0 + correct = hallucination = omission = valid = 0 for qa in qa_records: result_type = qa.get("result_type", "") - - if result_type in ["Correct", "Hallucination", "Omission"]: + if result_type == "Correct": + correct += 1 + valid += 1 + elif result_type == "Hallucination": + hallucination += 1 + valid += 1 + elif result_type == "Omission": + omission += 1 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, + "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 } - 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 def compute_time_metrics(results_file: str) -> dict[str, float]: """Compute timing metrics from evaluation results.""" - add_duration = 0 - search_duration = 0 + add_duration = search_duration = 0 with open(results_file, "r", encoding="utf-8") as f: for line in f: @@ -85,12 +67,10 @@ def compute_time_metrics(results_file: str) -> dict[str, float]: 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) - # Convert to minutes return { "add_dialogue_duration_time": add_duration / 1000 / 60, "search_memory_duration_time": search_duration / 1000 / 60, @@ -98,36 +78,27 @@ def compute_time_metrics(results_file: str) -> dict[str, float]: } -def load_from_tmp_dir(tmp_dir: str) -> tuple[str, list[dict]]: +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) - parent_dir = tmp_path.parent - eval_results_file = parent_dir / "eval_results.jsonl" + eval_results_file = tmp_path.parent / "eval_results.jsonl" - print(f"\n📁 Loading data from tmp directory: {tmp_dir}") - print(f"📝 Will generate: {eval_results_file}") + print(f"\n📁 Loading from: {tmp_dir}") + print(f"📝 Generating: {eval_results_file}") - # Collect all user directories user_dirs = [d for d in tmp_path.iterdir() if d.is_dir()] - print(f" Found {len(user_dirs)} user directories") + print(f" Found {len(user_dirs)} users") users_data = [] - for user_dir in user_dirs: - user_name = user_dir.name - - # Load all session files for this user (sorted by session number) 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]) # Sort by session number + [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: - print(f" ⚠️ No session files found for user: {user_name}") continue - # Load first session to get user metadata with open(session_files[0], "r", encoding="utf-8") as f: first_session = json.load(f) @@ -137,52 +108,46 @@ def load_from_tmp_dir(tmp_dir: str) -> tuple[str, list[dict]]: "sessions": [] } - # Load all sessions for session_file in session_files: with open(session_file, "r", encoding="utf-8") as f: session_data = json.load(f) - # Remove redundant user metadata session_data.pop("uuid", None) session_data.pop("user_name", None) user_data["sessions"].append(session_data) users_data.append(user_data) - print(f" ✓ Loaded user {user_name}: {len(session_files)} sessions") + print(f" ✓ {user_dir.name}: {len(session_files)} sessions") - # Write to eval_results.jsonl 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), users_data + 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): - logger.error(f"Input path not found: {input_path}") + print(f"❌ Error: Path not found: {input_path}") return print("\n" + "=" * 80) - print("COMPUTING QUESTION ANSWERING STATISTICS - REME V4") + print("REME V4 - QUESTION ANSWERING STATISTICS") print("=" * 80) - # Determine if input is a directory (tmp) or file (eval_results.jsonl) + # Load or generate eval_results.jsonl if os.path.isdir(input_path): - results_file, users_data = load_from_tmp_dir(input_path) + results_file = load_from_tmp_dir(input_path) else: results_file = input_path - users_data = None - print(f"\n📁 Using existing results file: {results_file}") + print(f"\n📁 Using: {results_file}") - # Collect all QA records with metadata + # Collect QA records with metadata qa_records = [] - qa_records_with_metadata = [] # Store records with user/session/question info - user_count = 0 - session_count = 0 + qa_with_metadata = [] + user_count = session_count = 0 with open(results_file, "r", encoding="utf-8") as f: for line in f: @@ -192,38 +157,38 @@ def main(input_path: str): user_count += 1 user_name = user_data.get("user_name", "Unknown") - valid_session_idx = 0 # Track the index of valid (non-skipped) sessions + 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", {}) - session_qa_records = eval_results.get("question_answering_records", []) - # Add records with metadata - for qa_idx, qa in enumerate(session_qa_records): + for qa_idx, qa in enumerate(eval_results.get("question_answering_records", [])): qa_records.append(qa) - qa_records_with_metadata.append({ + qa_with_metadata.append({ "user_name": user_name, "session_idx": valid_session_idx, - "original_session_idx": original_idx, "question_idx": qa_idx, "qa_record": qa }) valid_session_idx += 1 - print(f"\n📊 Data loaded:") + print(f"\n📊 Data Summary:") print(f" Users: {user_count}") print(f" Sessions: {session_count}") print(f" QA Records: {len(qa_records)}") # Compute metrics - print("\n🔄 Computing 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, @@ -232,22 +197,16 @@ def main(input_path: str): "question_answering_records": qa_records } - # Save final report - output_dir = Path(results_file).parent - report_file = output_dir / "reme_eval_stat_result.json" - with open(report_file, "w", encoding="utf-8") as f: json.dump(final_results, f, ensure_ascii=False, indent=4) - print(f"\n✅ Statistics saved to: {report_file}") + print(f"\n✅ Results saved to: {report_file}") - # Print summary + # Print metrics print("\n" + "=" * 80) - print("EVALUATION SUMMARY - REME V4") + print("📊 QUESTION ANSWERING METRICS") print("=" * 80) - - print("\n📊 Question Answering:") - print(f" Correct (all): {qa_metrics['correct_qa_ratio(all)']:.4f}") + 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}") @@ -255,42 +214,56 @@ def main(input_path: str): 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"\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 non-Correct QA records + # Print error records print("\n" + "=" * 80) - print("NON-CORRECT QA RECORDS") + print("❌ ERROR RECORDS (Non-Correct)") print("=" * 80) - non_correct_records = [ - record for record in qa_records_with_metadata - if record["qa_record"].get("result_type") not in ["Correct", ""] - ] + error_records = [r for r in qa_with_metadata if r["qa_record"].get("result_type") not in ["Correct", ""]] - if non_correct_records: - print(f"\nFound {len(non_correct_records)} non-correct records:\n") - for record in non_correct_records: - user_name = record["user_name"] - session_idx = record["session_idx"] - original_idx = record["original_session_idx"] - question_idx = record["question_idx"] - qa = record["qa_record"] - result_type = qa.get("result_type", "Unknown") - question = qa.get("question", "N/A") - answer = qa.get("answer", "N/A") - - print(f"👤 User: {user_name}") - print(f"📅 Session: {original_idx} (valid session index: {session_idx})") - print(f"❓ Question #{question_idx}") - print(f"🏷️ Result Type: {result_type}") - print(f"💬 Question: {question}") - print(f"💡 Answer: {answer}") - print("-" * 80) + if not error_records: + print("\n✅ All QA records are correct!") else: - print("\n✅ All QA records are Correct!") + 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) @@ -298,30 +271,15 @@ def main(input_path: str): 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, - required=False, - help="Path to eval_results.jsonl file (e.g., bench_results/reme_simple_v4/eval_results.jsonl)" - ) - parser.add_argument( - "--tmp_dir", - type=str, - required=False, - help="Path to tmp directory (e.g., bench_results/reme_simple_v4/tmp)" - ) + 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() - # Determine input path if args.tmp_dir: - input_path = args.tmp_dir + main(input_path=args.tmp_dir) elif args.results_file: - input_path = args.results_file + main(input_path=args.results_file) else: parser.error("Either --results_file or --tmp_dir must be provided") - - main(input_path=input_path) diff --git a/bench/halumem/eval_reme_simple_v4.py b/bench/halumem/eval_reme_simple_v4.py index c17622a6..825697cb 100644 --- a/bench/halumem/eval_reme_simple_v4.py +++ b/bench/halumem/eval_reme_simple_v4.py @@ -26,7 +26,7 @@ from typing import Any from loguru import logger -from eval_tools import evaluation_for_question2 +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 @@ -239,23 +239,35 @@ class MemoryProcessor: query: str, user_id: str, top_k: int = 20 - ) -> tuple[str, list, float]: + ) -> tuple[dict, list, float]: """ - Search memory using ReMe and return response. + Search memory using ReMe and return structured answer with reasoning. Returns: - tuple: (response, agent_messages, duration_ms) + tuple: (answer_dict, agent_messages, duration_ms) + answer_dict contains: {"reasoning": str, "answer": str, "memories": str} """ start = time.time() - response, agent_messages, success = await self.reme.retrieve_v4( + # 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 response, agent_messages, duration_ms + return answer_result, agent_messages, duration_ms # ==================== Evaluation ==================== @@ -279,19 +291,24 @@ class QuestionAnsweringEvaluator: results = [] for qa in questions: - response, agent_messages, duration_ms = await self.memory_processor.search_memory( + 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, - response, + system_answer, formatted_dialogue ) @@ -300,7 +317,9 @@ class QuestionAnsweringEvaluator: **qa, "uuid": uuid, "session_id": session_id, - "system_response": response, + "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"), diff --git a/bench/halumem/eval_tools.py b/bench/halumem/eval_tools.py index bf68ce2d..8377e87d 100644 --- a/bench/halumem/eval_tools.py +++ b/bench/halumem/eval_tools.py @@ -131,3 +131,40 @@ async def evaluation_for_question2( result = await llm_request_for_json(prompt) 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 cd7a2e6f..08c77ecc 100644 --- a/bench/halumem/halumem.yaml +++ b/bench/halumem/halumem.yaml @@ -2,6 +2,30 @@ 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: | You are an intelligent memory assistant tasked with retrieving accurate information from conversation memories. diff --git a/reme_ai/core/config/default.yaml b/reme_ai/core/config/default.yaml index 5891d180..c175e6a9 100644 --- a/reme_ai/core/config/default.yaml +++ b/reme_ai/core/config/default.yaml @@ -17,6 +17,7 @@ llm: backend: openai model_name: qwen3-30b-a3b-instruct-2507 max_concurrency: 20 + temperature: 0.0001 qwen3_max_instruct: backend: openai diff --git a/reme_ai/mem_agent/base_memory_agent.py b/reme_ai/mem_agent/base_memory_agent.py index 15c6dc04..d53cc8d9 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[:2000]}...\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 diff --git a/reme_ai/mem_agent/v4/personal_retriever_v4.py b/reme_ai/mem_agent/v4/personal_retriever_v4.py index 2b33897e..5135acd9 100644 --- a/reme_ai/mem_agent/v4/personal_retriever_v4.py +++ b/reme_ai/mem_agent/v4/personal_retriever_v4.py @@ -9,32 +9,25 @@ class PersonalRetrieverV4(BaseMemoryAgent): memory_type: MemoryType = MemoryType.PERSONAL async def build_messages(self) -> list[Message]: - if self.context.get("query"): - context = self.context.query - elif self.context.get("messages"): - context = format_messages(self.context.messages) - else: + 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() + 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 - messages = [ - Message( - role=Role.SYSTEM, - content=self.prompt_format( - prompt_name="system_prompt", - memory_type=self.memory_type.value, - memory_target=self.memory_target, - user_profile=read_profile_tool.output, - context=context, - )), + return [ Message( role=Role.USER, - content=self.get_prompt("user_message"), - ), + 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, + )) ] - return messages async def _acting_step(self, assistant_message: Message, step: int, **kwargs) -> list[Message]: return await super()._acting_step( @@ -44,3 +37,16 @@ class PersonalRetrieverV4(BaseMemoryAgent): 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 index cc5dcc34..85f7e013 100644 --- a/reme_ai/mem_agent/v4/personal_retriever_v4.yaml +++ b/reme_ai/mem_agent/v4/personal_retriever_v4.yaml @@ -1,7 +1,7 @@ tool: | Retrieve relevant personal memories to answer user questions through vector search and history reading. -system_prompt: | +user_message: | You are a memory agent managing **{memory_type}** memories about **{memory_target}**. ## User Profile @@ -10,28 +10,23 @@ system_prompt: | ## Question {context} - ## Retrieval Strategy + ## Task + Search for relevant memories to answer the question above. - **Tool 1: Vector Search (`retrieve_memory`) - - Try at least 3-5 different queries before moving to next tool: + **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 or remove time constraints: [start, end] in YYYYMMDD format + - 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 this ONLY after completing multiple retrieve_memory attempts - Use history_id from retrieved memories to read original conversations - - Prioritize most relevant or recent entries - Read multiple if needed for complete context **Response** - - **CRITICAL: Answer ONLY based on retrieved memories. Do NOT hallucinate or infer information not present in the search results.** - - Try multiple angles before giving up - - State "nothing found after thorough search" if nothing found after thorough search - -user_message: | - Answer the question using the retrieval strategy. + - 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 index 4b867bed..c0e1d4f1 100644 --- a/reme_ai/mem_agent/v4/personal_summarizer_v4.py +++ b/reme_ai/mem_agent/v4/personal_summarizer_v4.py @@ -13,17 +13,13 @@ class PersonalSummarizerV4(BaseMemoryAgent): history_node: MemoryNode = self.context.history_node messages = [ Message( - role=Role.SYSTEM, + role=Role.USER, content=self.prompt_format( - prompt_name="system_prompt_phase1", + prompt_name="user_message_phase1", context=history_node.content, memory_type=self.memory_type.value, memory_target=self.memory_target, )), - Message( - role=Role.USER, - content=self.get_prompt("user_message_phase1"), - ), ] return messages @@ -32,25 +28,22 @@ class PersonalSummarizerV4(BaseMemoryAgent): history_node: MemoryNode = self.context.history_node messages = [ Message( - role=Role.SYSTEM, + role=Role.USER, content=self.prompt_format( - prompt_name="system_prompt_phase2", + prompt_name="user_message_phase2", context=history_node.content, memory_type=self.memory_type.value, memory_target=self.memory_target, user_profile=user_profile, )), - Message( - role=Role.USER, - content=self.get_prompt("user_message_phase2"), - ), ] return messages - 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]: 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, @@ -68,7 +61,7 @@ class PersonalSummarizerV4(BaseMemoryAgent): ) # Phase 1: AddSummaryMemory - logger.info(f"[{self.__class__.__name__}] Starting 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() @@ -77,16 +70,16 @@ class PersonalSummarizerV4(BaseMemoryAgent): messages_phase1 = await self.build_messages_phase1() for i, message in enumerate(messages_phase1): logger.info( - f"[{self.__class__.__name__}] phase1.step0.{i} {message.role} " + 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) + messages_phase1, success_phase1 = await self.react(messages_phase1, stage="S1") if not success_phase1: - logger.warning(f"[{self.__class__.__name__}] Phase 1 did not complete successfully") + 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__}] Starting Phase 2: UpdateUserProfile") + logger.info(f"[{self.__class__.__name__}-S2] Starting Phase 2: UpdateUserProfile") # Restore original tools and get ReadUserProfile tool self.tools = original_tools @@ -94,13 +87,17 @@ class PersonalSummarizerV4(BaseMemoryAgent): user_profile = "" if read_profile_tool: - # Call ReadUserProfile to load current profile - logger.info(f"[{self.__class__.__name__}] Loading user profile with ReadUserProfile") - await read_profile_tool.call(memory_type=self.memory_type.value, memory_target=self.memory_target) + # 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__}] User profile loaded: {user_profile}...") + logger.info(f"[{self.__class__.__name__}-S2] User profile loaded: {user_profile}...") else: - logger.warning(f"[{self.__class__.__name__}] ReadUserProfile tool not found") + 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"] @@ -108,11 +105,11 @@ class PersonalSummarizerV4(BaseMemoryAgent): messages_phase2 = await self.build_messages_phase2(user_profile) for i, message in enumerate(messages_phase2): logger.info( - f"[{self.__class__.__name__}] phase2.step0.{i} {message.role} " + 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) + messages_phase2, success_phase2 = await self.react(messages_phase2, stage="S2") # Restore original tools self.tools = original_tools diff --git a/reme_ai/mem_agent/v4/personal_summarizer_v4.yaml b/reme_ai/mem_agent/v4/personal_summarizer_v4.yaml index 41d3ac61..5e7d12b7 100644 --- a/reme_ai/mem_agent/v4/personal_summarizer_v4.yaml +++ b/reme_ai/mem_agent/v4/personal_summarizer_v4.yaml @@ -2,7 +2,7 @@ tool: | Extract and update personal memories about the user from conversation context. Identify preferences, habits, background, relationships, and key facts. -system_prompt_phase1: | +user_message_phase1: | You are a memory agent managing **{memory_type}** memories about **{memory_target}**. ## Latest Conversation: @@ -17,11 +17,10 @@ system_prompt_phase1: | 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) -user_message_phase1: | Extract personal memories from the conversation using `AddSummaryMemory`. # capturing complete contexts with preconditions, causes, and consequences -system_prompt_phase2: | +user_message_phase2: | You are a memory agent managing **{memory_type}** memories about **{memory_target}**. ## Latest Conversation: @@ -36,10 +35,15 @@ system_prompt_phase2: | ## Task: Update Profile with `UpdateUserProfile` - Synchronize profile/memories with new information from the conversation, including **{memory_target}**' current status: - - `profile_ids_to_delete`: Remove outdated, conflicting, or redundant entries. - - `profiles_to_add`: Add new profiles/memories with `conversation_time`, e.g. `YYYY-MM-DD HH:MM:SS`, {memory_target} did something. - - Maintain profiles that are concise, mutually exclusive, and collectively comprehensive with no information loss. + 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 -user_message_phase2: | 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 index b8d5fbbe..7dbf2b8a 100644 --- a/reme_ai/mem_agent/v4/reme_retriever_v4.py +++ b/reme_ai/mem_agent/v4/reme_retriever_v4.py @@ -11,6 +11,7 @@ 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 @@ -44,10 +45,73 @@ class ReMeRetrieverV4(BaseMemoryAgent): return messages async def _acting_step(self, assistant_message: Message, step: int, **kwargs) -> list[Message]: - return await super()._acting_step( - assistant_message, - step, - query=self.context.get("query", ""), - messages=self.context.get("messages", []), - **kwargs, - ) + 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_tool/v4/hands_off.py b/reme_ai/mem_tool/v4/hands_off.py index 54b61f6d..17a33429 100644 --- a/reme_ai/mem_tool/v4/hands_off.py +++ b/reme_ai/mem_tool/v4/hands_off.py @@ -20,6 +20,7 @@ class HandsOff(BaseMemoryTool): 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"]: @@ -105,6 +106,8 @@ class HandsOff(BaseMemoryTool): 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}") diff --git a/reme_ai/mem_tool/v4/read_history.py b/reme_ai/mem_tool/v4/read_history.py index e9ab2a15..78a90eb4 100644 --- a/reme_ai/mem_tool/v4/read_history.py +++ b/reme_ai/mem_tool/v4/read_history.py @@ -34,5 +34,5 @@ class ReadHistory(BaseMemoryTool): return memory = MemoryNode.from_vector_node(nodes[0]) - self.output = memory.content + self.output = f"### Historical Dialogue\n{memory.content}" logger.info(f"Successfully read history memory: {history_id}") diff --git a/reme_ai/mem_tool/v4/read_user_profile.py b/reme_ai/mem_tool/v4/read_user_profile.py index b4aae557..aa8a57f8 100644 --- a/reme_ai/mem_tool/v4/read_user_profile.py +++ b/reme_ai/mem_tool/v4/read_user_profile.py @@ -1,3 +1,4 @@ +from typing import Literal from loguru import logger from ..base_memory_tool import BaseMemoryTool @@ -6,10 +7,11 @@ from ...core.schema.memory_node import MemoryNode class ReadUserProfile(BaseMemoryTool): - def __init__(self, add_memory_type_target: bool = False, **kwargs): + def __init__(self, add_memory_type_target: bool = False, show_ids: Literal["both", "profile", "history", "none"] = "both", **kwargs): kwargs["enable_multiple"] = False super().__init__(**kwargs) self.add_memory_type_target = add_memory_type_target + self.show_ids = show_ids def _build_tool_description(self) -> str: return "Read user profile." @@ -37,12 +39,16 @@ class ReadUserProfile(BaseMemoryTool): "required": [], } - async def execute(self): + async def execute(self): + # Determine which IDs to show + show_profile_id = self.show_ids in ("both", "profile") + show_history_id = self.show_ids in ("both", "history") + cache_key = f"{self.memory_type}_{self.memory_target}".replace(" ", "_").lower() cached_data = self.meta_memory.load(cache_key, auto_clean=False) if not cached_data: - self.output = "" + self.output = "### User Profile\nNo user profile found." logger.info(f"empty cached_data={cache_key}") return @@ -51,12 +57,25 @@ class ReadUserProfile(BaseMemoryTool): memory_formated = [] for node in memory_nodes: - node_formated = f"profile_id={node.memory_id} profile_content={node.content}" + node_formated_parts = [] + + # Add profile_id if enabled + if show_profile_id: + node_formated_parts.append(f"profile_id={node.memory_id}") + + # Always add profile_content + node_formated_parts.append(f"profile_content={node.content}") + + # Add conversation_time if available if "conversation_time" in node.metadata and node.metadata["conversation_time"]: - node_formated += f" conversation_time={node.metadata['conversation_time']}" - if node.ref_memory_id: - node_formated += f" history_id={node.ref_memory_id}" + node_formated_parts.append(f"conversation_time={node.metadata['conversation_time']}") + + # Add history_id if enabled and available + if show_history_id and node.ref_memory_id: + node_formated_parts.append(f"history_id={node.ref_memory_id}") + + node_formated = " ".join(node_formated_parts) memory_formated.append(node_formated.strip()) - self.output = "\n".join(memory_formated) + self.output = "### User Profile\n" + "\n".join(memory_formated) logger.info(f"Read {len(memory_formated)} nodes from cache key: {cache_key}") diff --git a/reme_ai/mem_tool/v4/retrieve_memory.py b/reme_ai/mem_tool/v4/retrieve_memory.py index e7f77f96..a8138bdc 100644 --- a/reme_ai/mem_tool/v4/retrieve_memory.py +++ b/reme_ai/mem_tool/v4/retrieve_memory.py @@ -98,6 +98,6 @@ class RetrieveMemory(BaseMemoryTool): if node.ref_memory_id: line += f"history_id={node.ref_memory_id} " output.append(line.strip()) - self.output = "\n".join(output) + 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/v4/update_user_profile.py b/reme_ai/mem_tool/v4/update_user_profile.py index 085daf80..1b9c5c97 100644 --- a/reme_ai/mem_tool/v4/update_user_profile.py +++ b/reme_ai/mem_tool/v4/update_user_profile.py @@ -80,7 +80,7 @@ class UpdateUserProfile(BaseMemoryTool): memory_target=self.memory_target, when_to_use="", content=mem.get("profile_content", ""), - ref_memory_id=self.ref_memory_id, + ref_memory_id=self.history_node.memory_id, author=self.author, metadata={"conversation_time": mem.get("conversation_time", "")}, ) diff --git a/reme_ai/reme.py b/reme_ai/reme.py index 02485159..d0a71e6f 100644 --- a/reme_ai/reme.py +++ b/reme_ai/reme.py @@ -458,7 +458,7 @@ class ReMe(Application): ) await reme_summarizer_v4.call(messages=messages, description=description, **kwargs) - return reme_summarizer_v4.memory_nodes, reme_summarizer_v4.messages, reme_summarizer_v4.success + return reme_summarizer_v4.memory_nodes, reme_summarizer_v4.tool_messages, reme_summarizer_v4.success else: raise NotImplementedError @@ -499,7 +499,7 @@ class ReMe(Application): ) await reme_retriever_v4.call(query=query, messages=messages, description=description, **kwargs) - return reme_retriever_v4.output, reme_retriever_v4.messages, reme_retriever_v4.success + return reme_retriever_v4.output, reme_retriever_v4.tool_messages, reme_retriever_v4.success else: raise NotImplementedError From e7a36067ebcdb208cb2d10776dbab304530e3880 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Mon, 19 Jan 2026 01:02:55 +0800 Subject: [PATCH 04/19] refactor(llm): replace concurrency control with request rate limiting --- reme_ai/core/config/default.yaml | 4 +-- reme_ai/core/llm/base_llm.py | 48 +++++++++++++++++++------------- 2 files changed, 31 insertions(+), 21 deletions(-) diff --git a/reme_ai/core/config/default.yaml b/reme_ai/core/config/default.yaml index c175e6a9..2c8488a8 100644 --- a/reme_ai/core/config/default.yaml +++ b/reme_ai/core/config/default.yaml @@ -16,14 +16,14 @@ llm: default: backend: openai model_name: qwen3-30b-a3b-instruct-2507 - max_concurrency: 20 + request_interval: 1 temperature: 0.0001 qwen3_max_instruct: backend: openai model_name: qwen3-max # temperature: 0.6 - max_concurrency: 20 + request_interval: 1 embedding_model: default: diff --git a/reme_ai/core/llm/base_llm.py b/reme_ai/core/llm/base_llm.py index 58db4b3f..a4a9d50a 100644 --- a/reme_ai/core/llm/base_llm.py +++ b/reme_ai/core/llm/base_llm.py @@ -17,24 +17,25 @@ from ..schema import ToolCall class BaseLLM(ABC): """Abstract base class defining the standard interface for LLM interactions.""" - def __init__(self, model_name: str, max_retries: int = 10, raise_exception: bool = False, max_concurrency: int | None = None, **kwargs): + def __init__(self, model_name: str, max_retries: int = 10, raise_exception: bool = False, request_interval: float = 0.0, **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. + request_interval: Minimum time interval (in seconds) between consecutive requests. Default is 0.0 (no interval). **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 + # Request rate control for async operations + 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]): @@ -128,14 +129,18 @@ class BaseLLM(ABC): 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 + # Apply request rate limiting if configured + 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: + sleep_time = self.request_interval - elapsed + await asyncio.sleep(sleep_time) + 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, @@ -356,12 +361,17 @@ class BaseLLM(ABC): 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) + # Apply request rate limiting if configured + 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: + sleep_time = self.request_interval - elapsed + await asyncio.sleep(sleep_time) + 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, From d145a67843dc62bf62f355e91db1dd9097faa290 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Mon, 19 Jan 2026 16:52:35 +0800 Subject: [PATCH 05/19] feat(config): update default LLM model configuration - Changed default model from qwen3-30b-a3b-instruct-2507 to qwen-flash - Updated evaluation tools to use EVALUATION_PROMPT_FOR_QUESTION instead of QUESTION2 - Added new PROMPT_MEMZERO_JSON2 configuration with context priority rules - Modified llm_request_for_json to use qwen-flash as default model - Updated halumem evaluation to specify qwen3-max model explicitly for certain requests --- bench/halumem/eval_tools.py | 5 +++-- bench/halumem/halumem.yaml | 24 ++++++++++++++++++++++++ bench/halumem/llms.py | 3 ++- reme_ai/core/config/default.yaml | 3 ++- 4 files changed, 31 insertions(+), 4 deletions(-) diff --git a/bench/halumem/eval_tools.py b/bench/halumem/eval_tools.py index 8377e87d..29868ca8 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,7 +129,7 @@ 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 diff --git a/bench/halumem/halumem.yaml b/bench/halumem/halumem.yaml index 08c77ecc..51419c37 100644 --- a/bench/halumem/halumem.yaml +++ b/bench/halumem/halumem.yaml @@ -26,6 +26,30 @@ PROMPT_MEMZERO_JSON: | }} ``` +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. diff --git a/bench/halumem/llms.py b/bench/halumem/llms.py index a6a06218..24edc703 100644 --- a/bench/halumem/llms.py +++ b/bench/halumem/llms.py @@ -58,7 +58,8 @@ 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: diff --git a/reme_ai/core/config/default.yaml b/reme_ai/core/config/default.yaml index 2c8488a8..6ae47d6a 100644 --- a/reme_ai/core/config/default.yaml +++ b/reme_ai/core/config/default.yaml @@ -15,7 +15,8 @@ http: llm: default: backend: openai - model_name: qwen3-30b-a3b-instruct-2507 +# model_name: qwen3-30b-a3b-instruct-2507 + model_name: qwen-flash request_interval: 1 temperature: 0.0001 From db510dd96ef16e793a30b77e8f9f76375b0ef079 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Tue, 20 Jan 2026 20:25:26 +0800 Subject: [PATCH 06/19] feat(bench): add human-in-the-loop evaluation framework for AI memory systems --- .gitignore | 4 +- bench/human_in_the_loop/__init__.py | 0 bench/human_in_the_loop/compute_qa_stats.py | 237 ++++++++ bench/human_in_the_loop/eval.yaml | 145 +++++ bench/human_in_the_loop/reevaluate_qa.py | 550 +++++++++++++++++++ bench/human_in_the_loop2/__init__.py | 0 bench/human_in_the_loop2/compute_qa_stats.py | 237 ++++++++ bench/human_in_the_loop2/eval.yaml | 145 +++++ bench/human_in_the_loop2/reevaluate_qa.py | 550 +++++++++++++++++++ reme_ai/core/config/default.yaml | 2 +- reme_ai/core/llm/lite_llm.py | 2 +- reme_ai/core/llm/lite_llm_sync.py | 2 +- reme_ai/core/llm/openai_llm.py | 2 +- reme_ai/core/llm/openai_llm_sync.py | 2 +- 14 files changed, 1872 insertions(+), 6 deletions(-) create mode 100644 bench/human_in_the_loop/__init__.py create mode 100644 bench/human_in_the_loop/compute_qa_stats.py create mode 100644 bench/human_in_the_loop/eval.yaml create mode 100644 bench/human_in_the_loop/reevaluate_qa.py create mode 100644 bench/human_in_the_loop2/__init__.py create mode 100644 bench/human_in_the_loop2/compute_qa_stats.py create mode 100644 bench/human_in_the_loop2/eval.yaml create mode 100644 bench/human_in_the_loop2/reevaluate_qa.py 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/bench/human_in_the_loop/__init__.py b/bench/human_in_the_loop/__init__.py new file mode 100644 index 00000000..e69de29b 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..65ac6dc8 --- /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..41d36943 --- /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..98b61b26 --- /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/bench/human_in_the_loop2/__init__.py b/bench/human_in_the_loop2/__init__.py new file mode 100644 index 00000000..e69de29b 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..65ac6dc8 --- /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..41d36943 --- /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..a65c4d1c --- /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/reme_ai/core/config/default.yaml b/reme_ai/core/config/default.yaml index 6ae47d6a..00ef4062 100644 --- a/reme_ai/core/config/default.yaml +++ b/reme_ai/core/config/default.yaml @@ -24,7 +24,7 @@ llm: backend: openai model_name: qwen3-max # temperature: 0.6 - request_interval: 1 + request_interval: 2 embedding_model: default: diff --git a/reme_ai/core/llm/lite_llm.py b/reme_ai/core/llm/lite_llm.py index 88177184..934fa52d 100644 --- a/reme_ai/core/llm/lite_llm.py +++ b/reme_ai/core/llm/lite_llm.py @@ -98,7 +98,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_ai/core/llm/lite_llm_sync.py index c3a925d2..11fbc7af 100644 --- a/reme_ai/core/llm/lite_llm_sync.py +++ b/reme_ai/core/llm/lite_llm_sync.py @@ -31,7 +31,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_ai/core/llm/openai_llm.py index 9e9ffc4a..7b645036 100644 --- a/reme_ai/core/llm/openai_llm.py +++ b/reme_ai/core/llm/openai_llm.py @@ -93,7 +93,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_ai/core/llm/openai_llm_sync.py index 51da29f4..a2bcfee3 100644 --- a/reme_ai/core/llm/openai_llm_sync.py +++ b/reme_ai/core/llm/openai_llm_sync.py @@ -35,7 +35,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 From ff7d25e3131faedcd721047b0af03bb36be4e62e Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Tue, 20 Jan 2026 20:35:20 +0800 Subject: [PATCH 07/19] chore(mem_tool): reorder required fields in user profile update schema --- reme_ai/mem_tool/v4/update_user_profile.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/reme_ai/mem_tool/v4/update_user_profile.py b/reme_ai/mem_tool/v4/update_user_profile.py index 1b9c5c97..a8fa04f5 100644 --- a/reme_ai/mem_tool/v4/update_user_profile.py +++ b/reme_ai/mem_tool/v4/update_user_profile.py @@ -40,7 +40,7 @@ class UpdateUserProfile(BaseMemoryTool): "description": "profile_content", }, }, - "required": ["profile_content", "conversation_time"], + "required": ["conversation_time", "profile_content"], }, }, }, From c3c7b5a4a072b0699438091c3fa6f89c4700591a Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Tue, 20 Jan 2026 20:47:24 +0800 Subject: [PATCH 08/19] fix(llm): enhance rate limit error detection in base LLM implementation --- reme_ai/core/llm/base_llm.py | 12 ++++++++++-- 1 file changed, 10 insertions(+), 2 deletions(-) diff --git a/reme_ai/core/llm/base_llm.py b/reme_ai/core/llm/base_llm.py index a4a9d50a..8f743ce6 100644 --- a/reme_ai/core/llm/base_llm.py +++ b/reme_ai/core/llm/base_llm.py @@ -402,7 +402,11 @@ class BaseLLM(ABC): # 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") @@ -475,7 +479,11 @@ class BaseLLM(ABC): # 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") From c7fc8255b168de22784cc43fba407791a86691d2 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 21 Jan 2026 16:34:09 +0800 Subject: [PATCH 09/19] refactor(core): move core module to core_old and update import paths --- .pre-commit-config.yaml | 12 +- bench/eval_reme_old.py | 4 +- bench/halumem/analyze_dataset_stats.py | 174 ++++---- bench/halumem/analyze_results.py | 68 +-- bench/halumem/compute_qa_stats_v4.py | 86 ++-- bench/halumem/compute_stats_from_tmp.py | 4 +- bench/halumem/eval_baseline_simple.py | 168 ++++---- bench/halumem/eval_reme.py | 12 +- bench/halumem/eval_reme_simple.py | 206 ++++----- bench/halumem/eval_reme_simple_v3.py | 214 +++++----- bench/halumem/eval_reme_simple_v4.py | 216 +++++----- bench/halumem/eval_tools.py | 8 +- bench/halumem/halumem.yaml | 6 +- bench/halumem/llms.py | 14 +- bench/human_in_the_loop/compute_qa_stats.py | 76 ++-- bench/human_in_the_loop/eval.yaml | 2 +- bench/human_in_the_loop/reevaluate_qa.py | 148 +++---- bench/human_in_the_loop2/compute_qa_stats.py | 76 ++-- bench/human_in_the_loop2/eval.yaml | 2 +- bench/human_in_the_loop2/reevaluate_qa.py | 148 +++---- docs/todo.md | 3 + reme_ai/core/__init__.py | 17 - reme_ai/core/context/prompt_handler.py | 382 ++++++++++++++--- reme_ai/core/context/registry.py | 144 ++++++- reme_ai/core/enumeration/json_schema_enum.py | 26 +- reme_ai/core/enumeration/memory_type.py | 30 +- reme_ai/core_old/__init__.py | 17 + reme_ai/{core => core_old}/application.py | 0 reme_ai/{core => core_old}/config/__init__.py | 0 .../{core => core_old}/config/default.yaml | 0 .../config/reme_config_parser.py | 0 reme_ai/core_old/context/__init__.py | 16 + reme_ai/core_old/context/base_context.py | 41 ++ reme_ai/core_old/context/prompt_handler.py | 95 +++++ reme_ai/core_old/context/registry.py | 46 ++ reme_ai/core_old/context/runtime_context.py | 79 ++++ .../context/service_context.py | 0 .../{core => core_old}/embedding/__init__.py | 0 .../embedding/base_embedding_model.py | 0 .../embedding/openai_embedding_model.py | 0 .../embedding/openai_embedding_model_sync.py | 0 reme_ai/core_old/enumeration/__init__.py | 17 + reme_ai/core_old/enumeration/chunk_enum.py | 25 ++ reme_ai/core_old/enumeration/http_enum.py | 22 + .../core_old/enumeration/json_schema_enum.py | 18 + reme_ai/core_old/enumeration/memory_type.py | 25 ++ reme_ai/core_old/enumeration/registry_enum.py | 28 ++ reme_ai/core_old/enumeration/role.py | 19 + reme_ai/{core => core_old}/flow/__init__.py | 0 reme_ai/{core => core_old}/flow/base_flow.py | 0 reme_ai/{core => core_old}/flow/cmd_flow.py | 0 .../flow/expression_flow.py | 0 .../{core => core_old}/flow/simple_flow.py | 0 reme_ai/{core => core_old}/llm/__init__.py | 0 reme_ai/{core => core_old}/llm/base_llm.py | 42 +- reme_ai/{core => core_old}/llm/lite_llm.py | 4 +- .../{core => core_old}/llm/lite_llm_sync.py | 0 reme_ai/{core => core_old}/llm/openai_llm.py | 4 +- .../{core => core_old}/llm/openai_llm_sync.py | 0 reme_ai/{core => core_old}/main.py | 0 reme_ai/{core => core_old}/op/__init__.py | 0 reme_ai/{core => core_old}/op/base_op.py | 0 reme_ai/{core => core_old}/op/base_ray_op.py | 0 reme_ai/{core => core_old}/op/mcp_tool.py | 0 reme_ai/{core => core_old}/op/parallel_op.py | 0 .../{core => core_old}/op/sequential_op.py | 0 reme_ai/{ => core_old}/reme.py | 108 +++-- reme_ai/{core => core_old}/schema/__init__.py | 0 .../{core => core_old}/schema/memory_node.py | 0 reme_ai/{core => core_old}/schema/message.py | 0 reme_ai/{core => core_old}/schema/request.py | 0 reme_ai/{core => core_old}/schema/response.py | 0 .../schema/service_config.py | 0 .../{core => core_old}/schema/stream_chunk.py | 0 .../{core => core_old}/schema/tool_call.py | 10 +- .../{core => core_old}/schema/vector_node.py | 0 .../{core => core_old}/service/__init__.py | 0 .../service/base_service.py | 0 .../{core => core_old}/service/cmd_service.py | 0 .../service/http_service.py | 0 .../{core => core_old}/service/mcp_service.py | 0 .../token_counter/__init__.py | 0 .../token_counter/base_token_counter.py | 0 .../token_counter/hf_token_counter.py | 0 .../token_counter/openai_token_counter.py | 0 reme_ai/{core => core_old}/utils/__init__.py | 0 .../{core => core_old}/utils/cache_handler.py | 0 .../utils/case_converter.py | 0 .../{core => core_old}/utils/common_utils.py | 0 reme_ai/{core => core_old}/utils/env_utils.py | 0 .../{core => core_old}/utils/execute_tuils.py | 0 .../{core => core_old}/utils/http_client.py | 0 reme_ai/{core => core_old}/utils/llm_utils.py | 0 .../{core => core_old}/utils/logger_utils.py | 0 .../{core => core_old}/utils/logo_utils.py | 0 .../{core => core_old}/utils/mcp_client.py | 0 .../utils/pydantic_config_parser.py | 0 .../utils/pydantic_utils.py | 0 reme_ai/{core => core_old}/utils/singleton.py | 0 reme_ai/{core => core_old}/utils/time.py | 0 .../vector_store/__init__.py | 0 .../vector_store/base_vector_store.py | 6 +- .../vector_store/chroma_vector_store.py | 4 +- .../vector_store/es_vector_store.py | 0 .../vector_store/local_vector_store.py | 2 +- .../vector_store/pgvector_store.py | 8 +- .../vector_store/qdrant_vector_store.py | 4 +- reme_ai/mem_agent/base_memory_agent.py | 10 +- reme_ai/mem_agent/chat/remy_agent.py | 8 +- reme_ai/mem_agent/chat/simple_chat.py | 8 +- reme_ai/mem_agent/chat/stream_chat.py | 8 +- reme_ai/mem_agent/retriever/reme_retriever.py | 8 +- .../retriever_v2/reme_retriever_v2.py | 14 +- .../retriever_v2/reme_retriever_v2.yaml | 34 +- .../reme_retriever_v2_simple.yaml | 30 +- .../summarizer/identity_summarizer.py | 8 +- .../summarizer/personal_summarizer.py | 8 +- .../summarizer/procedural_summarizer.py | 8 +- .../mem_agent/summarizer/reme_summarizer.py | 14 +- .../mem_agent/summarizer/tool_summarizer.py | 8 +- .../summarizer_v2/personal_summarizer_v2.py | 10 +- .../personal_summarizer_v2_simple.yaml | 2 +- .../summarizer_v2/reme_summarizer_v2.py | 8 +- .../summarizer_v2/reme_summarizer_v2.yaml | 2 +- .../mem_agent/v3/personal_summarizer_v3.py | 6 +- reme_ai/mem_agent/v3/reme_retriever_v3.py | 6 +- reme_ai/mem_agent/v3/reme_retriever_v3.yaml | 4 +- reme_ai/mem_agent/v3/reme_summarizer_v3.py | 6 +- reme_ai/mem_agent/v3/reme_summarizer_v3.yaml | 2 +- reme_ai/mem_agent/v4/personal_retriever_v4.py | 8 +- .../mem_agent/v4/personal_summarizer_v4.py | 4 +- .../mem_agent/v4/personal_summarizer_v4.yaml | 6 +- reme_ai/mem_agent/v4/reme_retriever_v4.py | 16 +- reme_ai/mem_agent/v4/reme_summarizer_v4.py | 6 +- reme_ai/mem_agent/v4/reme_summarizer_v4.yaml | 2 +- .../mem_agent/wk/personal_summarizer_wk.py | 6 +- reme_ai/mem_agent/wk/reme_retriever_wk.py | 6 +- reme_ai/mem_agent/wk/reme_retriever_wk.yaml | 34 +- reme_ai/mem_agent/wk/reme_summarizer_wk.py | 6 +- reme_ai/mem_agent/wk/reme_summarizer_wk.yaml | 2 +- reme_ai/mem_tool/base_memory_tool.py | 8 +- reme_ai/mem_tool/hands_off_tool.py | 4 +- .../mem_tool/history/add_history_memory.py | 8 +- .../mem_tool/history/read_history_memory.py | 4 +- .../mem_tool/identity/read_identity_memory.py | 2 +- .../identity/update_identity_memory.py | 2 +- reme_ai/mem_tool/meta/add_meta_memory.py | 4 +- reme_ai/mem_tool/meta/read_meta_memory.py | 4 +- reme_ai/mem_tool/read_local_memories.py | 2 +- reme_ai/mem_tool/think_tool.py | 4 +- reme_ai/mem_tool/v2/add_memory_drafts.py | 2 +- reme_ai/mem_tool/v2/read_history.py | 8 +- reme_ai/mem_tool/v2/retrieve_memories.py | 6 +- reme_ai/mem_tool/v2/retrieve_memories.yaml | 2 +- .../retrieve_recent_and_similar_memories.py | 6 +- .../retrieve_recent_and_similar_memories.yaml | 8 +- reme_ai/mem_tool/v2/summary_and_hands_off.py | 6 +- reme_ai/mem_tool/v2/update_memories.py | 4 +- reme_ai/mem_tool/v3/add_memory.py | 2 +- reme_ai/mem_tool/v3/read_history.py | 2 +- reme_ai/mem_tool/v3/read_user_profile.py | 2 +- reme_ai/mem_tool/v3/retrieve_memory.py | 4 +- reme_ai/mem_tool/v3/summary_and_hands_off.py | 4 +- reme_ai/mem_tool/v3/update_user_profile.py | 2 +- reme_ai/mem_tool/v4/add_summary_memory.py | 2 +- reme_ai/mem_tool/v4/hands_off.py | 8 +- reme_ai/mem_tool/v4/read_history.py | 2 +- reme_ai/mem_tool/v4/read_user_profile.py | 16 +- reme_ai/mem_tool/v4/retrieve_memory.py | 6 +- reme_ai/mem_tool/v4/update_user_profile.py | 4 +- reme_ai/mem_tool/vector_store/add_memory.py | 4 +- .../vector_store/add_summary_memory.py | 6 +- .../mem_tool/vector_store/delete_memory.py | 2 +- .../vector_store/retrieve_recent_memory.py | 6 +- .../mem_tool/vector_store/update_memory.py | 4 +- .../vector_store/vector_retrieve_memory.py | 8 +- reme_ai/mem_tool/wk/add_memory.py | 2 +- reme_ai/mem_tool/wk/read_history.py | 2 +- reme_ai/mem_tool/wk/summary_and_hands_off.py | 4 +- reme_ai/mem_tool/wk/update_memory.py | 2 +- reme_ai/mem_tool/wk/vector_retrieve_memory.py | 6 +- reme_ai/mem_tool/write_local_memories.py | 8 +- reme_ai/tool/execute/execute_code.py | 8 +- reme_ai/tool/execute/execute_shell.py | 8 +- reme_ai/tool/search/dashscope_search.py | 6 +- reme_ai/tool/search/mock_search.py | 10 +- reme_ai/tool/search/tavily_search.py | 6 +- test/http_client_test.py | 165 -------- test/mcp_client_test.py | 45 -- {tests => test}/mcp_servers_demo.json | 0 test/record_audio.py | 153 ------- test/test1.py | 12 - test/test2.py | 393 ------------------ test/test3.py | 92 ---- test/test4.py | 67 --- test/test5.py | 7 - test/test6.py | 15 - {tests => test}/test_base_context.py | 2 +- {tests => test}/test_cache_handler.py | 2 +- {tests => test}/test_embedding.py | 6 +- {tests => test}/test_embedding_sync.py | 6 +- {tests => test}/test_llm.py | 8 +- {tests => test}/test_llm_sync.py | 8 +- {tests => test}/test_logo.py | 4 +- {tests => test}/test_mcp_client.py | 2 +- {tests => test}/test_mcp_server.py | 4 +- .../test_memory_vector_conversion.py | 0 {tests => test}/test_message.py | 4 +- {tests => test}/test_op_composition.py | 4 +- {tests => test}/test_reme.py | 2 +- {tests => test}/test_timer.py | 2 +- {tests => test}/test_token_counter.py | 6 +- {tests => test}/test_tool.py | 4 +- {tests => test}/test_tool_call.py | 2 +- test/test_update_insight_op.py | 128 ------ {tests => test}/test_vector_store.py | 23 +- 216 files changed, 2190 insertions(+), 2400 deletions(-) create mode 100644 docs/todo.md create mode 100644 reme_ai/core_old/__init__.py rename reme_ai/{core => core_old}/application.py (100%) rename reme_ai/{core => core_old}/config/__init__.py (100%) rename reme_ai/{core => core_old}/config/default.yaml (100%) rename reme_ai/{core => core_old}/config/reme_config_parser.py (100%) create mode 100644 reme_ai/core_old/context/__init__.py create mode 100644 reme_ai/core_old/context/base_context.py create mode 100644 reme_ai/core_old/context/prompt_handler.py create mode 100644 reme_ai/core_old/context/registry.py create mode 100644 reme_ai/core_old/context/runtime_context.py rename reme_ai/{core => core_old}/context/service_context.py (100%) rename reme_ai/{core => core_old}/embedding/__init__.py (100%) rename reme_ai/{core => core_old}/embedding/base_embedding_model.py (100%) rename reme_ai/{core => core_old}/embedding/openai_embedding_model.py (100%) rename reme_ai/{core => core_old}/embedding/openai_embedding_model_sync.py (100%) create mode 100644 reme_ai/core_old/enumeration/__init__.py create mode 100644 reme_ai/core_old/enumeration/chunk_enum.py create mode 100644 reme_ai/core_old/enumeration/http_enum.py create mode 100644 reme_ai/core_old/enumeration/json_schema_enum.py create mode 100644 reme_ai/core_old/enumeration/memory_type.py create mode 100644 reme_ai/core_old/enumeration/registry_enum.py create mode 100644 reme_ai/core_old/enumeration/role.py rename reme_ai/{core => core_old}/flow/__init__.py (100%) rename reme_ai/{core => core_old}/flow/base_flow.py (100%) rename reme_ai/{core => core_old}/flow/cmd_flow.py (100%) rename reme_ai/{core => core_old}/flow/expression_flow.py (100%) rename reme_ai/{core => core_old}/flow/simple_flow.py (100%) rename reme_ai/{core => core_old}/llm/__init__.py (100%) rename reme_ai/{core => core_old}/llm/base_llm.py (98%) rename reme_ai/{core => core_old}/llm/lite_llm.py (99%) rename reme_ai/{core => core_old}/llm/lite_llm_sync.py (100%) rename reme_ai/{core => core_old}/llm/openai_llm.py (99%) rename reme_ai/{core => core_old}/llm/openai_llm_sync.py (100%) rename reme_ai/{core => core_old}/main.py (100%) rename reme_ai/{core => core_old}/op/__init__.py (100%) rename reme_ai/{core => core_old}/op/base_op.py (100%) rename reme_ai/{core => core_old}/op/base_ray_op.py (100%) rename reme_ai/{core => core_old}/op/mcp_tool.py (100%) rename reme_ai/{core => core_old}/op/parallel_op.py (100%) rename reme_ai/{core => core_old}/op/sequential_op.py (100%) rename reme_ai/{ => core_old}/reme.py (90%) rename reme_ai/{core => core_old}/schema/__init__.py (100%) rename reme_ai/{core => core_old}/schema/memory_node.py (100%) rename reme_ai/{core => core_old}/schema/message.py (100%) rename reme_ai/{core => core_old}/schema/request.py (100%) rename reme_ai/{core => core_old}/schema/response.py (100%) rename reme_ai/{core => core_old}/schema/service_config.py (100%) rename reme_ai/{core => core_old}/schema/stream_chunk.py (100%) rename reme_ai/{core => core_old}/schema/tool_call.py (99%) rename reme_ai/{core => core_old}/schema/vector_node.py (100%) rename reme_ai/{core => core_old}/service/__init__.py (100%) rename reme_ai/{core => core_old}/service/base_service.py (100%) rename reme_ai/{core => core_old}/service/cmd_service.py (100%) rename reme_ai/{core => core_old}/service/http_service.py (100%) rename reme_ai/{core => core_old}/service/mcp_service.py (100%) rename reme_ai/{core => core_old}/token_counter/__init__.py (100%) rename reme_ai/{core => core_old}/token_counter/base_token_counter.py (100%) rename reme_ai/{core => core_old}/token_counter/hf_token_counter.py (100%) rename reme_ai/{core => core_old}/token_counter/openai_token_counter.py (100%) rename reme_ai/{core => core_old}/utils/__init__.py (100%) rename reme_ai/{core => core_old}/utils/cache_handler.py (100%) rename reme_ai/{core => core_old}/utils/case_converter.py (100%) rename reme_ai/{core => core_old}/utils/common_utils.py (100%) rename reme_ai/{core => core_old}/utils/env_utils.py (100%) rename reme_ai/{core => core_old}/utils/execute_tuils.py (100%) rename reme_ai/{core => core_old}/utils/http_client.py (100%) rename reme_ai/{core => core_old}/utils/llm_utils.py (100%) rename reme_ai/{core => core_old}/utils/logger_utils.py (100%) rename reme_ai/{core => core_old}/utils/logo_utils.py (100%) rename reme_ai/{core => core_old}/utils/mcp_client.py (100%) rename reme_ai/{core => core_old}/utils/pydantic_config_parser.py (100%) rename reme_ai/{core => core_old}/utils/pydantic_utils.py (100%) rename reme_ai/{core => core_old}/utils/singleton.py (100%) rename reme_ai/{core => core_old}/utils/time.py (100%) rename reme_ai/{core => core_old}/vector_store/__init__.py (100%) rename reme_ai/{core => core_old}/vector_store/base_vector_store.py (97%) rename reme_ai/{core => core_old}/vector_store/chroma_vector_store.py (99%) rename reme_ai/{core => core_old}/vector_store/es_vector_store.py (100%) rename reme_ai/{core => core_old}/vector_store/local_vector_store.py (99%) rename reme_ai/{core => core_old}/vector_store/pgvector_store.py (99%) rename reme_ai/{core => core_old}/vector_store/qdrant_vector_store.py (99%) delete mode 100644 test/http_client_test.py delete mode 100644 test/mcp_client_test.py rename {tests => test}/mcp_servers_demo.json (100%) delete mode 100644 test/record_audio.py delete mode 100644 test/test1.py delete mode 100644 test/test2.py delete mode 100644 test/test3.py delete mode 100644 test/test4.py delete mode 100644 test/test5.py delete mode 100644 test/test6.py rename {tests => test}/test_base_context.py (97%) rename {tests => test}/test_cache_handler.py (97%) rename {tests => test}/test_embedding.py (98%) rename {tests => test}/test_embedding_sync.py (98%) rename {tests => test}/test_llm.py (98%) rename {tests => test}/test_llm_sync.py (98%) rename {tests => test}/test_logo.py (59%) rename {tests => test}/test_mcp_client.py (99%) rename {tests => test}/test_mcp_server.py (97%) rename {tests => test}/test_memory_vector_conversion.py (100%) rename {tests => test}/test_message.py (98%) rename {tests => test}/test_op_composition.py (99%) rename {tests => test}/test_reme.py (98%) rename {tests => test}/test_timer.py (97%) rename {tests => test}/test_token_counter.py (99%) rename {tests => test}/test_tool.py (98%) rename {tests => test}/test_tool_call.py (99%) delete mode 100644 test/test_update_insight_op.py rename {tests => test}/test_vector_store.py (99%) diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 1795bc17..3d6fcf96 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/core_old/|reme_ai/mem_agent/|reme_ai/mem_tool/|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/core_old/|reme_ai/mem_agent/|reme_ai/mem_tool/|bench) - repo: https://github.com/psf/black rev: 25.9.0 hooks: - id: black - exclude: ^(test/|cookbook/) + exclude: ^(test/|cookbook/|reme_ai/core_old/|reme_ai/mem_agent/|reme_ai/mem_tool/|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/core_old/|reme_ai/mem_agent/|reme_ai/mem_tool/|bench) args: [ "--extend-ignore=E203", "--max-line-length=120" @@ -44,6 +44,10 @@ repos: | \.demo$ | \.md$ | \.html$ + | reme_ai/core_old/ + | reme_ai/mem_agent/ + | reme_ai/mem_tool/ + | bench ) args: [ --disable=W0511, diff --git a/bench/eval_reme_old.py b/bench/eval_reme_old.py index 51265da7..41fde3b1 100644 --- a/bench/eval_reme_old.py +++ b/bench/eval_reme_old.py @@ -10,8 +10,8 @@ from datetime import datetime, timezone from tqdm import tqdm -from reme_ai.core.enumeration import Role -from reme_ai.core.schema import Message, MemoryNode +from reme_ai.core_old.enumeration import Role +from reme_ai.core_old.schema import Message, MemoryNode from reme_ai.reme import ReMe TEMPLATE_REME = """Memories for user {user_id}: diff --git a/bench/halumem/analyze_dataset_stats.py b/bench/halumem/analyze_dataset_stats.py index d5a213b4..061f96ae 100644 --- a/bench/halumem/analyze_dataset_stats.py +++ b/bench/halumem/analyze_dataset_stats.py @@ -38,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 数量 @@ -69,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 中提取用户名""" @@ -84,7 +84,7 @@ class DatasetAnalyzer: if not match: return "Unknown" return match.group(1).strip() - + @staticmethod def calculate_dialogue_length(dialogue: list[dict]) -> int: """计算对话的总长度(字符数)""" @@ -93,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: """ @@ -102,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) @@ -137,72 +137,72 @@ 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, @@ -212,25 +212,25 @@ class DatasetAnalyzer: 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 = ( @@ -240,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, @@ -300,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 = [] @@ -318,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}") @@ -355,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) @@ -392,50 +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 = ( @@ -444,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 = { @@ -512,10 +512,10 @@ class DatasetAnalyzer: 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}") @@ -525,27 +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) @@ -557,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" ) @@ -578,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 index c7bdfab2..afd07e4b 100644 --- a/bench/halumem/compute_qa_stats_v4.py +++ b/bench/halumem/compute_qa_stats_v4.py @@ -26,9 +26,9 @@ def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: "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": @@ -40,7 +40,7 @@ def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: elif result_type == "Omission": omission += 1 valid += 1 - + metrics = { "correct_qa_ratio(all)": correct / total, "hallucination_qa_ratio(all)": hallucination / total, @@ -51,26 +51,26 @@ def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: "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, @@ -82,73 +82,73 @@ 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(): @@ -156,15 +156,15 @@ def main(input_path: str): 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({ @@ -173,22 +173,22 @@ def main(input_path: str): "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, @@ -196,12 +196,12 @@ def main(input_path: str): }, "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") @@ -213,27 +213,27 @@ def main(input_path: str): 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}") @@ -264,19 +264,19 @@ def main(input_path: str): 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: diff --git a/bench/halumem/compute_stats_from_tmp.py b/bench/halumem/compute_stats_from_tmp.py index 5ab401ba..2d0201ec 100644 --- a/bench/halumem/compute_stats_from_tmp.py +++ b/bench/halumem/compute_stats_from_tmp.py @@ -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..3254c332 100644 --- a/bench/halumem/eval_reme.py +++ b/bench/halumem/eval_reme.py @@ -35,8 +35,8 @@ from eval_tools import ( evaluation_for_update_memory, ) from llms import llm_request -from reme_ai.core.enumeration import MemoryType -from reme_ai.core.schema import MemoryNode +from reme_ai.core_old.enumeration import MemoryType +from reme_ai.core_old.schema import MemoryNode from reme_ai.reme import ReMe # Template for formatting memories (from shared YAML config) @@ -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..4635c556 100644 --- a/bench/halumem/eval_reme_simple.py +++ b/bench/halumem/eval_reme_simple.py @@ -26,8 +26,8 @@ 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.core_old.enumeration import MemoryType +from reme_ai.core_old.schema import MemoryNode from reme_ai.reme import ReMe @@ -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 index 7c0dfd8e..43cee5d9 100644 --- a/bench/halumem/eval_reme_simple_v3.py +++ b/bench/halumem/eval_reme_simple_v3.py @@ -27,8 +27,8 @@ 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.core_old.enumeration import MemoryType +from reme_ai.core_old.schema import MemoryNode from reme_ai.reme import ReMe @@ -49,13 +49,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.""" @@ -63,7 +63,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 with conversation_time (user messages only).""" @@ -80,7 +80,7 @@ class DataLoader: 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.""" @@ -89,10 +89,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" @@ -103,29 +103,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) @@ -133,38 +133,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: @@ -173,7 +173,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") @@ -181,19 +181,19 @@ class FileManager: class MemoryProcessor: """Handles ReMe V3 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 = 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) """ @@ -201,24 +201,24 @@ class MemoryProcessor: 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, + 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: @@ -234,28 +234,28 @@ 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 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, + query=query, + user_id=user_id, top_k=top_k ) - + duration_ms = (time.time() - start) * 1000 return response, agent_messages, duration_ms @@ -264,11 +264,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], @@ -279,7 +279,7 @@ class QuestionAnsweringEvaluator: ) -> 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( @@ -287,7 +287,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( @@ -297,7 +297,7 @@ class QuestionAnsweringEvaluator: response, formatted_dialogue ) - + # Build result record qa_result = { **qa, @@ -310,13 +310,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.""" @@ -332,15 +332,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": @@ -349,7 +349,7 @@ class MetricsAggregator: hallucination += 1 elif result_type == "Omission": omission += 1 - + metrics = { "correct_qa_ratio(all)": correct / total, "hallucination_qa_ratio(all)": hallucination / total, @@ -357,7 +357,7 @@ class MetricsAggregator: "qa_valid_num": valid, "qa_num": total } - + if valid > 0: metrics.update({ "correct_qa_ratio(valid)": correct / valid, @@ -370,28 +370,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, @@ -404,18 +404,18 @@ class MetricsAggregator: 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, + self.memory_processor, config.top_k ) self.data_loader = DataLoader() - + async def process_session( self, session: dict, @@ -430,29 +430,29 @@ class HaluMemEvaluatorV3: "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) @@ -463,97 +463,97 @@ class HaluMemEvaluatorV3: 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) + 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: @@ -561,20 +561,20 @@ class HaluMemEvaluatorV3: 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, @@ -582,23 +582,23 @@ class HaluMemEvaluatorV3: }, "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}") @@ -607,7 +607,7 @@ class HaluMemEvaluatorV3: 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") @@ -630,14 +630,14 @@ def main( 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)" ) @@ -665,9 +665,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_v4.py b/bench/halumem/eval_reme_simple_v4.py index 825697cb..63cb4f37 100644 --- a/bench/halumem/eval_reme_simple_v4.py +++ b/bench/halumem/eval_reme_simple_v4.py @@ -27,8 +27,8 @@ 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.core_old.enumeration import MemoryType +from reme_ai.core_old.schema import MemoryNode from reme_ai.reme import ReMe @@ -49,13 +49,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.""" @@ -63,7 +63,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 with conversation_time (user messages only).""" @@ -80,7 +80,7 @@ class DataLoader: 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.""" @@ -89,10 +89,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" @@ -103,29 +103,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) @@ -133,38 +133,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: @@ -173,7 +173,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") @@ -181,19 +181,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 = 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) """ @@ -201,23 +201,23 @@ class MemoryProcessor: 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, + 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: @@ -233,39 +233,39 @@ 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[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, + 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 @@ -274,11 +274,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], @@ -289,19 +289,19 @@ class QuestionAnsweringEvaluator: ) -> 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( @@ -311,7 +311,7 @@ class QuestionAnsweringEvaluator: system_answer, formatted_dialogue ) - + # Build result record qa_result = { **qa, @@ -326,13 +326,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.""" @@ -348,15 +348,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": @@ -365,7 +365,7 @@ class MetricsAggregator: hallucination += 1 elif result_type == "Omission": omission += 1 - + metrics = { "correct_qa_ratio(all)": correct / total, "hallucination_qa_ratio(all)": hallucination / total, @@ -373,7 +373,7 @@ class MetricsAggregator: "qa_valid_num": valid, "qa_num": total } - + if valid > 0: metrics.update({ "correct_qa_ratio(valid)": correct / valid, @@ -386,28 +386,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, @@ -426,11 +426,11 @@ class HaluMemEvaluatorV4: 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, @@ -445,7 +445,7 @@ class HaluMemEvaluatorV4: "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 @@ -453,20 +453,20 @@ class HaluMemEvaluatorV4: 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) @@ -477,97 +477,97 @@ class HaluMemEvaluatorV4: 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) + 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: @@ -575,20 +575,20 @@ class HaluMemEvaluatorV4: 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, @@ -596,23 +596,23 @@ class HaluMemEvaluatorV4: }, "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}") @@ -621,7 +621,7 @@ class HaluMemEvaluatorV4: 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") @@ -644,14 +644,14 @@ def main( 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)" ) @@ -679,9 +679,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_tools.py b/bench/halumem/eval_tools.py index 29868ca8..e0fae45d 100644 --- a/bench/halumem/eval_tools.py +++ b/bench/halumem/eval_tools.py @@ -141,12 +141,12 @@ async def answer_question_with_memories( ): """ 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 """ @@ -158,13 +158,13 @@ async def answer_question_with_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") diff --git a/bench/halumem/halumem.yaml b/bench/halumem/halumem.yaml index 51419c37..56dde2db 100644 --- a/bench/halumem/halumem.yaml +++ b/bench/halumem/halumem.yaml @@ -11,7 +11,7 @@ PROMPT_MEMZERO_JSON: | 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} @@ -35,7 +35,7 @@ PROMPT_MEMZERO_JSON2: | 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} @@ -467,7 +467,7 @@ EVALUATION_PROMPT_FOR_QUESTION: | "evaluation_result": "Correct | Hallucination | Omission" }} ``` - + EVALUATION_PROMPT_FOR_QUESTION2: | You are an **evaluation expert for AI memory system question answering**. diff --git a/bench/halumem/llms.py b/bench/halumem/llms.py index 24edc703..5a59b2b2 100644 --- a/bench/halumem/llms.py +++ b/bench/halumem/llms.py @@ -5,8 +5,8 @@ import re from tenacity import retry, stop_after_attempt, wait_random_exponential, before_sleep_log -from reme_ai.core.schema import Message -from reme_ai.core.utils import load_env +from reme_ai.core_old.schema import Message +from reme_ai.core_old.utils import load_env from reme_ai.reme import ReMe logger = logging.getLogger(__name__) @@ -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 """ @@ -61,15 +61,15 @@ async def llm_request(prompt, model_name: str = "qwen3-max", **kwargs) -> str: 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/bench/human_in_the_loop/compute_qa_stats.py b/bench/human_in_the_loop/compute_qa_stats.py index 65ac6dc8..69d65e55 100644 --- a/bench/human_in_the_loop/compute_qa_stats.py +++ b/bench/human_in_the_loop/compute_qa_stats.py @@ -22,9 +22,9 @@ def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: "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": @@ -36,7 +36,7 @@ def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: elif result_type == "Omission": omission += 1 valid += 1 - + metrics = { "correct_qa_ratio(all)": correct / total, "hallucination_qa_ratio(all)": hallucination / total, @@ -47,21 +47,21 @@ def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: "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, @@ -72,21 +72,21 @@ def compute_time_metrics(users_data: list[dict]) -> dict[str, float]: 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(): @@ -103,71 +103,71 @@ def load_from_tmp_dir(tmp_dir: str) -> list[dict]: 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({ @@ -176,17 +176,17 @@ def main(tmp_dir: str): "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, @@ -194,10 +194,10 @@ def main(tmp_dir: str): }, "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:") @@ -205,12 +205,12 @@ def main(tmp_dir: str): 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"] @@ -218,13 +218,13 @@ def main(tmp_dir: str): 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", @@ -232,6 +232,6 @@ if __name__ == "__main__": 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 index 41d36943..3e685bc7 100644 --- a/bench/human_in_the_loop/eval.yaml +++ b/bench/human_in_the_loop/eval.yaml @@ -64,7 +64,7 @@ EVALUATION_PROMPT_FOR_QUESTION: | "evaluation_result": "Correct | Hallucination | Omission" }} ``` - + EVALUATION_PROMPT_FOR_QUESTION2: | You are an **evaluation expert for AI memory system question answering**. diff --git a/bench/human_in_the_loop/reevaluate_qa.py b/bench/human_in_the_loop/reevaluate_qa.py index 98b61b26..a013616c 100644 --- a/bench/human_in_the_loop/reevaluate_qa.py +++ b/bench/human_in_the_loop/reevaluate_qa.py @@ -16,8 +16,8 @@ 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.core_old.schema import Message +from reme_ai.core_old.utils import load_env from reme_ai.reme import ReMe from tenacity import retry, stop_after_attempt, wait_random_exponential @@ -82,7 +82,7 @@ async def evaluate_qa_record( 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 @@ -90,9 +90,9 @@ async def evaluate_qa_record( response: System response to evaluate dialogue: Dialogue context (optional) model_name: LLM model name - prompt_version: "v1" for EVALUATION_PROMPT_FOR_QUESTION, + prompt_version: "v1" for EVALUATION_PROMPT_FOR_QUESTION, "v2" for EVALUATION_PROMPT_FOR_QUESTION2 - + Returns: dict with evaluation_result and reasoning """ @@ -101,7 +101,7 @@ async def evaluate_qa_record( prompt_template = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION2"] else: prompt_template = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION"] - + # Format prompt prompt = prompt_template.format( question=question, @@ -110,7 +110,7 @@ async def evaluate_qa_record( response=response, dialogue=dialogue or "N/A" ) - + result = await llm_request_for_json(prompt, model_name=model_name) return result @@ -118,14 +118,14 @@ async def evaluate_qa_record( 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) @@ -135,18 +135,18 @@ def load_from_data_dir(data_dir: str) -> list[dict]: "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, + sessions, key=lambda s: s["data"].get("session_idx", 0) ) users_data.extend(sessions_sorted) - + return users_data - + return [] @@ -155,7 +155,7 @@ def format_dialogue_context(session_data: dict) -> str: dialogue = session_data.get("session", {}).get("dialogue", []) if not dialogue: return "N/A" - + formatted_turns = [] for turn in dialogue: role = turn.get("role", "unknown") @@ -175,34 +175,34 @@ async def reevaluate_session( 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; + 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, @@ -210,19 +210,19 @@ async def reevaluate_session( 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( @@ -234,28 +234,28 @@ async def reevaluate_session( 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)...") @@ -268,18 +268,18 @@ async def reevaluate_session( 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: @@ -287,7 +287,7 @@ async def reevaluate_session( 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] @@ -296,21 +296,21 @@ async def reevaluate_session( 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 @@ -328,9 +328,9 @@ def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: "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": @@ -342,7 +342,7 @@ def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: elif result_type == "Omission": omission += 1 valid += 1 - + metrics = { "correct_qa_ratio(all)": correct / total, "hallucination_qa_ratio(all)": hallucination / total, @@ -353,7 +353,7 @@ def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: "qa_valid_num": valid, "qa_num": total } - + return metrics @@ -364,28 +364,28 @@ async def main( 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)}") @@ -393,27 +393,27 @@ async def main( 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, @@ -421,25 +421,25 @@ async def main( 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: @@ -452,10 +452,10 @@ async def main( "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}") @@ -463,17 +463,17 @@ async def main( 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: @@ -485,28 +485,28 @@ async def main( "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." @@ -539,9 +539,9 @@ if __name__ == "__main__": 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, diff --git a/bench/human_in_the_loop2/compute_qa_stats.py b/bench/human_in_the_loop2/compute_qa_stats.py index 65ac6dc8..69d65e55 100644 --- a/bench/human_in_the_loop2/compute_qa_stats.py +++ b/bench/human_in_the_loop2/compute_qa_stats.py @@ -22,9 +22,9 @@ def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: "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": @@ -36,7 +36,7 @@ def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: elif result_type == "Omission": omission += 1 valid += 1 - + metrics = { "correct_qa_ratio(all)": correct / total, "hallucination_qa_ratio(all)": hallucination / total, @@ -47,21 +47,21 @@ def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: "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, @@ -72,21 +72,21 @@ def compute_time_metrics(users_data: list[dict]) -> dict[str, float]: 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(): @@ -103,71 +103,71 @@ def load_from_tmp_dir(tmp_dir: str) -> list[dict]: 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({ @@ -176,17 +176,17 @@ def main(tmp_dir: str): "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, @@ -194,10 +194,10 @@ def main(tmp_dir: str): }, "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:") @@ -205,12 +205,12 @@ def main(tmp_dir: str): 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"] @@ -218,13 +218,13 @@ def main(tmp_dir: str): 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", @@ -232,6 +232,6 @@ if __name__ == "__main__": 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 index 41d36943..3e685bc7 100644 --- a/bench/human_in_the_loop2/eval.yaml +++ b/bench/human_in_the_loop2/eval.yaml @@ -64,7 +64,7 @@ EVALUATION_PROMPT_FOR_QUESTION: | "evaluation_result": "Correct | Hallucination | Omission" }} ``` - + EVALUATION_PROMPT_FOR_QUESTION2: | You are an **evaluation expert for AI memory system question answering**. diff --git a/bench/human_in_the_loop2/reevaluate_qa.py b/bench/human_in_the_loop2/reevaluate_qa.py index a65c4d1c..a5e6d0fd 100644 --- a/bench/human_in_the_loop2/reevaluate_qa.py +++ b/bench/human_in_the_loop2/reevaluate_qa.py @@ -16,8 +16,8 @@ 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.core_old.schema import Message +from reme_ai.core_old.utils import load_env from reme_ai.reme import ReMe from tenacity import retry, stop_after_attempt, wait_random_exponential @@ -82,7 +82,7 @@ async def evaluate_qa_record( 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 @@ -90,9 +90,9 @@ async def evaluate_qa_record( response: System response to evaluate dialogue: Dialogue context (optional) model_name: LLM model name - prompt_version: "v1" for EVALUATION_PROMPT_FOR_QUESTION, + prompt_version: "v1" for EVALUATION_PROMPT_FOR_QUESTION, "v2" for EVALUATION_PROMPT_FOR_QUESTION2 - + Returns: dict with evaluation_result and reasoning """ @@ -101,7 +101,7 @@ async def evaluate_qa_record( prompt_template = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION2"] else: prompt_template = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION"] - + # Format prompt prompt = prompt_template.format( question=question, @@ -110,7 +110,7 @@ async def evaluate_qa_record( response=response, dialogue=dialogue or "N/A" ) - + result = await llm_request_for_json(prompt, model_name=model_name) return result @@ -118,14 +118,14 @@ async def evaluate_qa_record( 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) @@ -135,18 +135,18 @@ def load_from_data_dir(data_dir: str) -> list[dict]: "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, + sessions, key=lambda s: s["data"].get("session_idx", 0) ) users_data.extend(sessions_sorted) - + return users_data - + return [] @@ -155,7 +155,7 @@ def format_dialogue_context(session_data: dict) -> str: dialogue = session_data.get("session", {}).get("dialogue", []) if not dialogue: return "N/A" - + formatted_turns = [] for turn in dialogue: role = turn.get("role", "unknown") @@ -175,34 +175,34 @@ async def reevaluate_session( 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; + 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, @@ -210,19 +210,19 @@ async def reevaluate_session( 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( @@ -234,28 +234,28 @@ async def reevaluate_session( 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)...") @@ -268,18 +268,18 @@ async def reevaluate_session( 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: @@ -287,7 +287,7 @@ async def reevaluate_session( 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] @@ -296,21 +296,21 @@ async def reevaluate_session( 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 @@ -328,9 +328,9 @@ def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: "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": @@ -342,7 +342,7 @@ def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: elif result_type == "Omission": omission += 1 valid += 1 - + metrics = { "correct_qa_ratio(all)": correct / total, "hallucination_qa_ratio(all)": hallucination / total, @@ -353,7 +353,7 @@ def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: "qa_valid_num": valid, "qa_num": total } - + return metrics @@ -364,28 +364,28 @@ async def main( 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)}") @@ -393,27 +393,27 @@ async def main( 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, @@ -421,25 +421,25 @@ async def main( 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: @@ -452,10 +452,10 @@ async def main( "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}") @@ -463,17 +463,17 @@ async def main( 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: @@ -485,28 +485,28 @@ async def main( "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." @@ -539,9 +539,9 @@ if __name__ == "__main__": 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, 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/reme_ai/core/__init__.py b/reme_ai/core/__init__.py index 8eab5792..e69de29b 100644 --- a/reme_ai/core/__init__.py +++ b/reme_ai/core/__init__.py @@ -1,17 +0,0 @@ -"""Core module for ReMe AI framework.""" - -# pylint: disable=wrong-import-position -# flake8: noqa: F401 - -from . import config -from . import context -from . import embedding -from . import enumeration -from . import flow -from . import llm -from . import op -from . import schema -from . import service -from . import token_counter -from . import utils -from . import vector_store diff --git a/reme_ai/core/context/prompt_handler.py b/reme_ai/core/context/prompt_handler.py index b428b163..9639cbf8 100644 --- a/reme_ai/core/context/prompt_handler.py +++ b/reme_ai/core/context/prompt_handler.py @@ -1,24 +1,100 @@ -"""Module for managing and formatting prompt templates from files or dictionaries.""" +"""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 -from .service_context import C + + +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.""" + pass class PromptHandler(BaseContext): - """A context-aware handler for loading, retrieving, and formatting prompt templates.""" + """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 handler with a specific language and optional context data.""" + """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 or C.language + self.language: str = language.strip() - def load_prompt_by_file(self, prompt_file_path: Path | str = None): - """Load prompt configurations from a YAML file into the context.""" + 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 @@ -26,70 +102,272 @@ class PromptHandler(BaseContext): 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 - 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) + 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: dict = None): - """Merge a dictionary of prompt strings into the current context.""" + 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 isinstance(value, str): - if key in self: - logger.warning(f"Overwriting prompt key={key}, old_value={self[key]}, new_value={value}") + 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"Adding new prompt key={key}, value={value}") + 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): - """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() + 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() - assert key in self, f"prompt_name={key} not found." - return self[key].strip() + # Try base name + if prompt_name in self: + return self[prompt_name].strip() - 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) + # 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() - # 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)} + # Not found, raise error with helpful message + available = list(self.keys()) + raise PromptNotFoundError(prompt_name, available) - 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 + 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 - hit = True - hit_flag = flag - # Remove the flag prefix from the line - line = line.strip(f"[{key}]") + 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 - # 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) + 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 - prompt = "\n".join(split_prompt) + return "\n".join(filtered_lines) - if other_kwargs: - # Apply standard Python string formatting - prompt = prompt.format(**other_kwargs) + 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) - return prompt + # 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_ai/core/context/registry.py b/reme_ai/core/context/registry.py index f403037d..9fa35d30 100644 --- a/reme_ai/core/context/registry.py +++ b/reme_ai/core/context/registry.py @@ -1,19 +1,143 @@ """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 ..enumeration import RegistryEnum +from ...core_old.utils import singleton + +T = TypeVar("T") +@singleton class Registry(BaseContext): - """A registry container that uses decorators to map and store class references.""" + """A singleton registry manager that maintains separate registries for different component types. - def register(self, name: str = "", add_cls: bool = True): - """Return a decorator that registers a class under a specific name in the registry.""" + This class serves as the central registry hub for the entire ReMe application, providing: + - Component registration for different types (LLMs, embeddings, vector stores, etc.) + - Convenient access methods for retrieving registered classes + - Decorator-based registration API - 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 + The singleton pattern ensures only one instance exists throughout the application lifecycle, + accessible via the global `R` variable exported at the bottom of this module. + """ - return decorator + def __init__(self, **kwargs): + """Initialize the registry manager with separate registries for each component type.""" + super().__init__(**kwargs) + + # Registry system: stores class definitions for different component types + self.registry_dict: dict[RegistryEnum, dict] = { + v: {} for v in RegistryEnum.__members__.values() + } + + def register(self, name: str | type = "", register_type: RegistryEnum = None) -> Callable[[type[T]], type[T]] | type[T]: + """Return a decorator to register a component within a specific registry category. + + Can be used in multiple ways: + - @R.register_op() # with empty parentheses, uses class name + - @R.register_op # without parentheses, uses class name + - @R.register_op("custom_name") # with custom name + + Args: + name: Either a string name for the class, or the class itself when used without parentheses + register_type: The type of registry (LLM, EMBEDDING_MODEL, VECTOR_STORE, etc.) + + Returns: + Either a decorator function or the registered class itself + + Example: + @R.register("my_llm", RegistryEnum.LLM) + class MyLLM(BaseLLM): + pass + """ + if inspect.isclass(name): + # Used without parentheses: @R.register_op + self.registry_dict[register_type][name.__name__] = name + return name + else: + # Used with parentheses: @R.register_op() or @R.register_op("name") + def decorator(cls): + key = name if isinstance(name, str) and name else cls.__name__ + self.registry_dict[register_type][key] = cls + return cls + + return decorator + + 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) + + +# Export a global singleton instance for easy access across the application +# This is the primary way to access the registry throughout the codebase +R = Registry() diff --git a/reme_ai/core/enumeration/json_schema_enum.py b/reme_ai/core/enumeration/json_schema_enum.py index 507645f4..d66882e2 100644 --- a/reme_ai/core/enumeration/json_schema_enum.py +++ b/reme_ai/core/enumeration/json_schema_enum.py @@ -1,18 +1,38 @@ -"""Defines the standard data types supported by JSON Schema.""" +"""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.""" + """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: - """Returns the string representation of the enum value.""" + """Return the lowercase JSON Schema type name for this enum member.""" return self.name.lower() diff --git a/reme_ai/core/enumeration/memory_type.py b/reme_ai/core/enumeration/memory_type.py index 22d35481..b9f5ed29 100644 --- a/reme_ai/core/enumeration/memory_type.py +++ b/reme_ai/core/enumeration/memory_type.py @@ -1,25 +1,33 @@ -"""Memory type enumeration for the three-layer memory architecture.""" +"""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): - """ - Three-layer memory architecture for agent memory management. + """Enumeration of memory categories used by the memory subsystem. - 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 + 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_old/__init__.py b/reme_ai/core_old/__init__.py new file mode 100644 index 00000000..8eab5792 --- /dev/null +++ b/reme_ai/core_old/__init__.py @@ -0,0 +1,17 @@ +"""Core module for ReMe AI framework.""" + +# pylint: disable=wrong-import-position +# flake8: noqa: F401 + +from . import config +from . import context +from . import embedding +from . import enumeration +from . import flow +from . import llm +from . import op +from . import schema +from . import service +from . import token_counter +from . import utils +from . import vector_store diff --git a/reme_ai/core/application.py b/reme_ai/core_old/application.py similarity index 100% rename from reme_ai/core/application.py rename to reme_ai/core_old/application.py diff --git a/reme_ai/core/config/__init__.py b/reme_ai/core_old/config/__init__.py similarity index 100% rename from reme_ai/core/config/__init__.py rename to reme_ai/core_old/config/__init__.py diff --git a/reme_ai/core/config/default.yaml b/reme_ai/core_old/config/default.yaml similarity index 100% rename from reme_ai/core/config/default.yaml rename to reme_ai/core_old/config/default.yaml diff --git a/reme_ai/core/config/reme_config_parser.py b/reme_ai/core_old/config/reme_config_parser.py similarity index 100% rename from reme_ai/core/config/reme_config_parser.py rename to reme_ai/core_old/config/reme_config_parser.py diff --git a/reme_ai/core_old/context/__init__.py b/reme_ai/core_old/context/__init__.py new file mode 100644 index 00000000..7f26d600 --- /dev/null +++ b/reme_ai/core_old/context/__init__.py @@ -0,0 +1,16 @@ +"""context""" + +from .base_context import BaseContext +from .prompt_handler import PromptHandler +from .registry import Registry +from .runtime_context import RuntimeContext +from .service_context import ServiceContext, C + +__all__ = [ + "BaseContext", + "PromptHandler", + "Registry", + "RuntimeContext", + "ServiceContext", + "C", +] diff --git a/reme_ai/core_old/context/base_context.py b/reme_ai/core_old/context/base_context.py new file mode 100644 index 00000000..dabd8cdb --- /dev/null +++ b/reme_ai/core_old/context/base_context.py @@ -0,0 +1,41 @@ +"""Module providing a dictionary subclass with attribute-style access and pickling support.""" + +from typing import Generic, TypeVar + +_KT = TypeVar("_KT") +_VT = TypeVar("_VT") + + +class BaseContext(dict, Generic[_KT, _VT]): + """A dictionary subclass that enables accessing and modifying keys as attributes.""" + + def __getattr__(self, name: str) -> _VT: + """Retrieve a dictionary item as an attribute.""" + try: + return self[name] + except KeyError as e: + raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") from e + + def __setattr__(self, name: str, value: _VT) -> None: + """Assign a value to a dictionary item using attribute syntax.""" + self[name] = value + + def __delattr__(self, name: str) -> None: + """Remove a dictionary item using attribute syntax.""" + try: + # Delete item from dict via key + del self[name] + except KeyError as e: + raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") from e + + def __getstate__(self) -> dict: + """Return the dictionary representation for pickling.""" + return dict(self) + + def __setstate__(self, state: dict) -> None: + """Restore the dictionary state from a pickled object.""" + self.update(state) + + def __reduce__(self): + """Define the reconstruction logic for pickling processes.""" + return self.__class__, (), self.__getstate__() diff --git a/reme_ai/core_old/context/prompt_handler.py b/reme_ai/core_old/context/prompt_handler.py new file mode 100644 index 00000000..b428b163 --- /dev/null +++ b/reme_ai/core_old/context/prompt_handler.py @@ -0,0 +1,95 @@ +"""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].strip() + + 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_old/context/registry.py b/reme_ai/core_old/context/registry.py new file mode 100644 index 00000000..08ea2271 --- /dev/null +++ b/reme_ai/core_old/context/registry.py @@ -0,0 +1,46 @@ +"""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 + +T = TypeVar('T') + + +class Registry(BaseContext): + """A registry container that uses decorators to map and store class references.""" + + def register(self, name: str | type = "", add_cls: bool = True) -> Callable[[type[T]], type[T]] | type[T]: + """Return a decorator that registers a class under a specific name in the registry. + + Can be used in three ways: + - @C.register_op() # with empty parentheses, uses class name + - @C.register_op # without parentheses, uses class name + - @C.register_op("custom_name") # with custom name + + Args: + name: Either a string name for the class, or the class itself when used without parentheses + add_cls: Whether to actually add the class to the registry + + Returns: + Either a decorator function or the registered class itself + """ + + def decorator(cls): + if add_cls: + # Use provided name or default to the class name as the key + key = name if isinstance(name, str) and name else cls.__name__ + self[key] = cls + return cls + + # If used without parentheses: @C.register_op + if inspect.isclass(name): + cls = name + # Register with class name as key + if add_cls: + self[cls.__name__] = cls + return cls + + # If used with parentheses: @C.register_op() or @C.register_op("name") + return decorator diff --git a/reme_ai/core_old/context/runtime_context.py b/reme_ai/core_old/context/runtime_context.py new file mode 100644 index 00000000..d7112e1c --- /dev/null +++ b/reme_ai/core_old/context/runtime_context.py @@ -0,0 +1,79 @@ +"""Runtime context for managing response states and asynchronous data streaming.""" + +import asyncio + +from .base_context import BaseContext +from ..enumeration import ChunkEnum +from ..schema import Response, StreamChunk + + +class RuntimeContext(BaseContext): + """Context for execution state, response metadata, and stream queues.""" + + def __init__( + self, + response: Response | None = None, + stream_queue: asyncio.Queue | None = None, + **kwargs, + ): + """Initialize the context with optional response and queue.""" + super().__init__(**kwargs) + self.response = response or Response() + self.stream_queue = stream_queue + + @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 + + async def _enqueue(self, chunk: StreamChunk) -> None: + """Internal helper to put a chunk into the queue if it exists.""" + if self.stream_queue: + await self.stream_queue.put(chunk) + + async def add_stream_string(self, chunk: str, chunk_type: ChunkEnum) -> "RuntimeContext": + """Enqueue a stream chunk from a raw string and type.""" + await self._enqueue(StreamChunk(chunk_type=chunk_type, chunk=chunk)) + return self + + async def add_stream_chunk(self, stream_chunk: StreamChunk) -> "RuntimeContext": + """Enqueue an existing stream chunk.""" + await self._enqueue(stream_chunk) + return self + + async def add_stream_done(self) -> "RuntimeContext": + """Enqueue a termination chunk to signal the end of the stream.""" + await self._enqueue(StreamChunk(chunk_type=ChunkEnum.DONE, chunk="", done=True)) + return self + + def add_response_error(self, e: Exception) -> "RuntimeContext": + """Record an exception into the response object.""" + self.response.success = False + self.response.answer = str(e) + return self + + def apply_mapping(self, mapping: dict[str, str]) -> "RuntimeContext": + """Copy internal values based on a source-to-target key map.""" + if not mapping: + return self + + for source, target in mapping.items(): + if source in self: + self[target] = self[source] + return self + + def validate_required_keys(self, required_keys: dict[str, bool], context_name: str = "context") -> "RuntimeContext": + """Ensure all required keys are present in the context. + + Args: + required_keys: Dictionary mapping key names to boolean indicating if required + context_name: Name of the context for error messages (e.g., operator name) + """ + for key, is_required in required_keys.items(): + if is_required and key not in self: + raise ValueError(f"{context_name}: missing required input '{key}'") + return self diff --git a/reme_ai/core/context/service_context.py b/reme_ai/core_old/context/service_context.py similarity index 100% rename from reme_ai/core/context/service_context.py rename to reme_ai/core_old/context/service_context.py diff --git a/reme_ai/core/embedding/__init__.py b/reme_ai/core_old/embedding/__init__.py similarity index 100% rename from reme_ai/core/embedding/__init__.py rename to reme_ai/core_old/embedding/__init__.py diff --git a/reme_ai/core/embedding/base_embedding_model.py b/reme_ai/core_old/embedding/base_embedding_model.py similarity index 100% rename from reme_ai/core/embedding/base_embedding_model.py rename to reme_ai/core_old/embedding/base_embedding_model.py diff --git a/reme_ai/core/embedding/openai_embedding_model.py b/reme_ai/core_old/embedding/openai_embedding_model.py similarity index 100% rename from reme_ai/core/embedding/openai_embedding_model.py rename to reme_ai/core_old/embedding/openai_embedding_model.py diff --git a/reme_ai/core/embedding/openai_embedding_model_sync.py b/reme_ai/core_old/embedding/openai_embedding_model_sync.py similarity index 100% rename from reme_ai/core/embedding/openai_embedding_model_sync.py rename to reme_ai/core_old/embedding/openai_embedding_model_sync.py diff --git a/reme_ai/core_old/enumeration/__init__.py b/reme_ai/core_old/enumeration/__init__.py new file mode 100644 index 00000000..7202a949 --- /dev/null +++ b/reme_ai/core_old/enumeration/__init__.py @@ -0,0 +1,17 @@ +"""enumeration""" + +from .chunk_enum import ChunkEnum +from .http_enum import HttpEnum +from .json_schema_enum import JsonSchemaEnum +from .memory_type import MemoryType +from .registry_enum import RegistryEnum +from .role import Role + +__all__ = [ + "ChunkEnum", + "HttpEnum", + "JsonSchemaEnum", + "MemoryType", + "RegistryEnum", + "Role", +] diff --git a/reme_ai/core_old/enumeration/chunk_enum.py b/reme_ai/core_old/enumeration/chunk_enum.py new file mode 100644 index 00000000..dbe37106 --- /dev/null +++ b/reme_ai/core_old/enumeration/chunk_enum.py @@ -0,0 +1,25 @@ +"""Defines the types of data chunks used in streaming responses.""" + +from enum import Enum + + +class ChunkEnum(str, Enum): + """Enumeration of possible chunk categories for stream processing.""" + + # Internal reasoning or chain-of-thought process + THINK = "think" + + # The final generated response content + ANSWER = "answer" + + # Metadata or calls related to external tools + TOOL = "tool" + + # Resource consumption and token usage statistics + USAGE = "usage" + + # Error messages or exception details + ERROR = "error" + + # Final signal indicating the completion of the stream + DONE = "done" diff --git a/reme_ai/core_old/enumeration/http_enum.py b/reme_ai/core_old/enumeration/http_enum.py new file mode 100644 index 00000000..19622242 --- /dev/null +++ b/reme_ai/core_old/enumeration/http_enum.py @@ -0,0 +1,22 @@ +"""Provides a collection of standard HTTP request methods.""" + +from enum import Enum + + +class HttpEnum(str, Enum): + """Enumeration of supported HTTP methods for network requests.""" + + # Retrieves data from a specified resource + GET = "get" + + # Submits data to be processed to a specified resource + POST = "post" + + # Identical to GET but only retrieves the response headers + HEAD = "head" + + # Uploads or replaces the representation of a target resource + PUT = "put" + + # Deletes the specified resource from the server + DELETE = "delete" diff --git a/reme_ai/core_old/enumeration/json_schema_enum.py b/reme_ai/core_old/enumeration/json_schema_enum.py new file mode 100644 index 00000000..507645f4 --- /dev/null +++ b/reme_ai/core_old/enumeration/json_schema_enum.py @@ -0,0 +1,18 @@ +"""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_old/enumeration/memory_type.py b/reme_ai/core_old/enumeration/memory_type.py new file mode 100644 index 00000000..22d35481 --- /dev/null +++ b/reme_ai/core_old/enumeration/memory_type.py @@ -0,0 +1,25 @@ +"""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_old/enumeration/registry_enum.py b/reme_ai/core_old/enumeration/registry_enum.py new file mode 100644 index 00000000..876c06b8 --- /dev/null +++ b/reme_ai/core_old/enumeration/registry_enum.py @@ -0,0 +1,28 @@ +"""Defines the registry categories for core components of the system.""" + +from enum import Enum + + +class RegistryEnum(str, Enum): + """Enumeration of component types registered within the application lifecycle.""" + + # Large Language Model interfaces + LLM = "llm" + + # Models used for generating vector embeddings + EMBEDDING_MODEL = "embedding_model" + + # Databases or storage systems for vector search + VECTOR_STORE = "vector_store" + + # Atomic operations or functional units + OP = "op" + + # Orchestrated sequences of operations or workflows + FLOW = "flow" + + # External APIs or shared internal services + SERVICE = "service" + + # Utilities for tracking and limiting token consumption + TOKEN_COUNTER = "token_counter" diff --git a/reme_ai/core_old/enumeration/role.py b/reme_ai/core_old/enumeration/role.py new file mode 100644 index 00000000..4acad7e5 --- /dev/null +++ b/reme_ai/core_old/enumeration/role.py @@ -0,0 +1,19 @@ +"""Defines the participant roles in a chat completion sequence.""" + +from enum import Enum + + +class Role(str, Enum): + """Enumeration of standard personas involved in a conversation flow.""" + + # High-level instructions to guide the model's behavior + SYSTEM = "system" + + # Input or queries provided by the human user + USER = "user" + + # Responses or messages generated by the AI model + ASSISTANT = "assistant" + + # Output or results returned from external tool executions + TOOL = "tool" diff --git a/reme_ai/core/flow/__init__.py b/reme_ai/core_old/flow/__init__.py similarity index 100% rename from reme_ai/core/flow/__init__.py rename to reme_ai/core_old/flow/__init__.py diff --git a/reme_ai/core/flow/base_flow.py b/reme_ai/core_old/flow/base_flow.py similarity index 100% rename from reme_ai/core/flow/base_flow.py rename to reme_ai/core_old/flow/base_flow.py diff --git a/reme_ai/core/flow/cmd_flow.py b/reme_ai/core_old/flow/cmd_flow.py similarity index 100% rename from reme_ai/core/flow/cmd_flow.py rename to reme_ai/core_old/flow/cmd_flow.py diff --git a/reme_ai/core/flow/expression_flow.py b/reme_ai/core_old/flow/expression_flow.py similarity index 100% rename from reme_ai/core/flow/expression_flow.py rename to reme_ai/core_old/flow/expression_flow.py diff --git a/reme_ai/core/flow/simple_flow.py b/reme_ai/core_old/flow/simple_flow.py similarity index 100% rename from reme_ai/core/flow/simple_flow.py rename to reme_ai/core_old/flow/simple_flow.py diff --git a/reme_ai/core/llm/__init__.py b/reme_ai/core_old/llm/__init__.py similarity index 100% rename from reme_ai/core/llm/__init__.py rename to reme_ai/core_old/llm/__init__.py diff --git a/reme_ai/core/llm/base_llm.py b/reme_ai/core_old/llm/base_llm.py similarity index 98% rename from reme_ai/core/llm/base_llm.py rename to reme_ai/core_old/llm/base_llm.py index 8f743ce6..fcd543b6 100644 --- a/reme_ai/core/llm/base_llm.py +++ b/reme_ai/core_old/llm/base_llm.py @@ -19,7 +19,7 @@ class BaseLLM(ABC): def __init__(self, model_name: str, max_retries: int = 10, raise_exception: bool = False, request_interval: float = 0.0, **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 @@ -32,7 +32,7 @@ class BaseLLM(ABC): self.raise_exception: bool = raise_exception self.request_interval: float = request_interval self.kwargs: dict = kwargs - + # Request rate control for async operations self._last_request_time: float = 0.0 self._request_lock: asyncio.Lock = asyncio.Lock() @@ -87,7 +87,7 @@ class BaseLLM(ABC): **kwargs, ) -> dict: """Construct provider-specific parameters for streaming API requests. - + Args: messages: List of conversation messages tools: Optional list of tool calls @@ -122,7 +122,7 @@ class BaseLLM(ABC): **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 @@ -138,10 +138,10 @@ class BaseLLM(ABC): sleep_time = self.request_interval - elapsed await asyncio.sleep(sleep_time) 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], @@ -178,7 +178,7 @@ class BaseLLM(ABC): **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 @@ -213,7 +213,7 @@ class BaseLLM(ABC): **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 @@ -281,7 +281,7 @@ class BaseLLM(ABC): **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 @@ -351,7 +351,7 @@ class BaseLLM(ABC): **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 @@ -370,9 +370,9 @@ class BaseLLM(ABC): sleep_time = self.request_interval - elapsed await asyncio.sleep(sleep_time) 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], @@ -386,7 +386,7 @@ class BaseLLM(ABC): """Internal implementation of chat with retry and error handling logic.""" # 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 - + for i in range(self.max_retries): try: result = await self._chat( @@ -407,7 +407,7 @@ class BaseLLM(ABC): "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) @@ -424,12 +424,12 @@ class BaseLLM(ABC): 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})") await asyncio.sleep(60) continue - + logger.exception(f"chat with model={effective_model} encounter error with e={e.args}") if i == self.max_retries - 1: @@ -451,7 +451,7 @@ class BaseLLM(ABC): **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 @@ -463,7 +463,7 @@ class BaseLLM(ABC): """ # 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 - + for i in range(self.max_retries): try: result = self._chat_sync( @@ -484,7 +484,7 @@ class BaseLLM(ABC): "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) @@ -501,12 +501,12 @@ class BaseLLM(ABC): 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})") time.sleep(60) continue - + logger.exception(f"chat sync with model={effective_model} encounter error with e={e.args}") if i == self.max_retries - 1: diff --git a/reme_ai/core/llm/lite_llm.py b/reme_ai/core_old/llm/lite_llm.py similarity index 99% rename from reme_ai/core/llm/lite_llm.py rename to reme_ai/core_old/llm/lite_llm.py index 934fa52d..458fe1c0 100644 --- a/reme_ai/core/llm/lite_llm.py +++ b/reme_ai/core_old/llm/lite_llm.py @@ -40,7 +40,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 +50,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, diff --git a/reme_ai/core/llm/lite_llm_sync.py b/reme_ai/core_old/llm/lite_llm_sync.py similarity index 100% rename from reme_ai/core/llm/lite_llm_sync.py rename to reme_ai/core_old/llm/lite_llm_sync.py diff --git a/reme_ai/core/llm/openai_llm.py b/reme_ai/core_old/llm/openai_llm.py similarity index 99% rename from reme_ai/core/llm/openai_llm.py rename to reme_ai/core_old/llm/openai_llm.py index 7b645036..f15dcb8e 100644 --- a/reme_ai/core/llm/openai_llm.py +++ b/reme_ai/core_old/llm/openai_llm.py @@ -45,7 +45,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 +55,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, diff --git a/reme_ai/core/llm/openai_llm_sync.py b/reme_ai/core_old/llm/openai_llm_sync.py similarity index 100% rename from reme_ai/core/llm/openai_llm_sync.py rename to reme_ai/core_old/llm/openai_llm_sync.py diff --git a/reme_ai/core/main.py b/reme_ai/core_old/main.py similarity index 100% rename from reme_ai/core/main.py rename to reme_ai/core_old/main.py diff --git a/reme_ai/core/op/__init__.py b/reme_ai/core_old/op/__init__.py similarity index 100% rename from reme_ai/core/op/__init__.py rename to reme_ai/core_old/op/__init__.py diff --git a/reme_ai/core/op/base_op.py b/reme_ai/core_old/op/base_op.py similarity index 100% rename from reme_ai/core/op/base_op.py rename to reme_ai/core_old/op/base_op.py diff --git a/reme_ai/core/op/base_ray_op.py b/reme_ai/core_old/op/base_ray_op.py similarity index 100% rename from reme_ai/core/op/base_ray_op.py rename to reme_ai/core_old/op/base_ray_op.py diff --git a/reme_ai/core/op/mcp_tool.py b/reme_ai/core_old/op/mcp_tool.py similarity index 100% rename from reme_ai/core/op/mcp_tool.py rename to reme_ai/core_old/op/mcp_tool.py diff --git a/reme_ai/core/op/parallel_op.py b/reme_ai/core_old/op/parallel_op.py similarity index 100% rename from reme_ai/core/op/parallel_op.py rename to reme_ai/core_old/op/parallel_op.py diff --git a/reme_ai/core/op/sequential_op.py b/reme_ai/core_old/op/sequential_op.py similarity index 100% rename from reme_ai/core/op/sequential_op.py rename to reme_ai/core_old/op/sequential_op.py diff --git a/reme_ai/reme.py b/reme_ai/core_old/reme.py similarity index 90% rename from reme_ai/reme.py rename to reme_ai/core_old/reme.py index d0a71e6f..b5581159 100644 --- a/reme_ai/reme.py +++ b/reme_ai/core_old/reme.py @@ -1,14 +1,14 @@ """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 @@ -177,7 +177,6 @@ class ReMe(Application): except Exception as e: print(f"Warning: reme_summarizer.call failed: {e}") return [] - else: raise NotImplementedError @@ -228,18 +227,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.""" @@ -293,14 +291,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.""" @@ -342,12 +340,12 @@ class ReMe(Application): raise NotImplementedError async def summary_v3( - 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 V3 workflow with user profile management.""" @@ -384,14 +382,14 @@ class ReMe(Application): 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, + 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.""" @@ -425,13 +423,13 @@ class ReMe(Application): raise NotImplementedError async def summary_v4( - self, - messages: list[dict], - description: str = "", - user_id: str = "", - assistant_id: str = "", - enable_thinking_params: bool = False, - **kwargs, + 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.""" @@ -464,15 +462,15 @@ class ReMe(Application): 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, + 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.""" diff --git a/reme_ai/core/schema/__init__.py b/reme_ai/core_old/schema/__init__.py similarity index 100% rename from reme_ai/core/schema/__init__.py rename to reme_ai/core_old/schema/__init__.py diff --git a/reme_ai/core/schema/memory_node.py b/reme_ai/core_old/schema/memory_node.py similarity index 100% rename from reme_ai/core/schema/memory_node.py rename to reme_ai/core_old/schema/memory_node.py diff --git a/reme_ai/core/schema/message.py b/reme_ai/core_old/schema/message.py similarity index 100% rename from reme_ai/core/schema/message.py rename to reme_ai/core_old/schema/message.py diff --git a/reme_ai/core/schema/request.py b/reme_ai/core_old/schema/request.py similarity index 100% rename from reme_ai/core/schema/request.py rename to reme_ai/core_old/schema/request.py diff --git a/reme_ai/core/schema/response.py b/reme_ai/core_old/schema/response.py similarity index 100% rename from reme_ai/core/schema/response.py rename to reme_ai/core_old/schema/response.py diff --git a/reme_ai/core/schema/service_config.py b/reme_ai/core_old/schema/service_config.py similarity index 100% rename from reme_ai/core/schema/service_config.py rename to reme_ai/core_old/schema/service_config.py diff --git a/reme_ai/core/schema/stream_chunk.py b/reme_ai/core_old/schema/stream_chunk.py similarity index 100% rename from reme_ai/core/schema/stream_chunk.py rename to reme_ai/core_old/schema/stream_chunk.py diff --git a/reme_ai/core/schema/tool_call.py b/reme_ai/core_old/schema/tool_call.py similarity index 99% rename from reme_ai/core/schema/tool_call.py rename to reme_ai/core_old/schema/tool_call.py index 21dde495..3c01cdcd 100644 --- a/reme_ai/core/schema/tool_call.py +++ b/reme_ai/core_old/schema/tool_call.py @@ -177,7 +177,7 @@ class ToolCall(BaseModel): return True except Exception: return False - + def sanitize_and_check_argument(self) -> bool: """ Attempt to sanitize and validate arguments JSON. @@ -188,17 +188,17 @@ class ToolCall(BaseModel): """ 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: @@ -212,7 +212,7 @@ class ToolCall(BaseModel): sanitized = sanitized[:-1].rstrip() else: break - + return False def simple_output_dump(self) -> dict: diff --git a/reme_ai/core/schema/vector_node.py b/reme_ai/core_old/schema/vector_node.py similarity index 100% rename from reme_ai/core/schema/vector_node.py rename to reme_ai/core_old/schema/vector_node.py diff --git a/reme_ai/core/service/__init__.py b/reme_ai/core_old/service/__init__.py similarity index 100% rename from reme_ai/core/service/__init__.py rename to reme_ai/core_old/service/__init__.py diff --git a/reme_ai/core/service/base_service.py b/reme_ai/core_old/service/base_service.py similarity index 100% rename from reme_ai/core/service/base_service.py rename to reme_ai/core_old/service/base_service.py diff --git a/reme_ai/core/service/cmd_service.py b/reme_ai/core_old/service/cmd_service.py similarity index 100% rename from reme_ai/core/service/cmd_service.py rename to reme_ai/core_old/service/cmd_service.py diff --git a/reme_ai/core/service/http_service.py b/reme_ai/core_old/service/http_service.py similarity index 100% rename from reme_ai/core/service/http_service.py rename to reme_ai/core_old/service/http_service.py diff --git a/reme_ai/core/service/mcp_service.py b/reme_ai/core_old/service/mcp_service.py similarity index 100% rename from reme_ai/core/service/mcp_service.py rename to reme_ai/core_old/service/mcp_service.py diff --git a/reme_ai/core/token_counter/__init__.py b/reme_ai/core_old/token_counter/__init__.py similarity index 100% rename from reme_ai/core/token_counter/__init__.py rename to reme_ai/core_old/token_counter/__init__.py diff --git a/reme_ai/core/token_counter/base_token_counter.py b/reme_ai/core_old/token_counter/base_token_counter.py similarity index 100% rename from reme_ai/core/token_counter/base_token_counter.py rename to reme_ai/core_old/token_counter/base_token_counter.py diff --git a/reme_ai/core/token_counter/hf_token_counter.py b/reme_ai/core_old/token_counter/hf_token_counter.py similarity index 100% rename from reme_ai/core/token_counter/hf_token_counter.py rename to reme_ai/core_old/token_counter/hf_token_counter.py diff --git a/reme_ai/core/token_counter/openai_token_counter.py b/reme_ai/core_old/token_counter/openai_token_counter.py similarity index 100% rename from reme_ai/core/token_counter/openai_token_counter.py rename to reme_ai/core_old/token_counter/openai_token_counter.py diff --git a/reme_ai/core/utils/__init__.py b/reme_ai/core_old/utils/__init__.py similarity index 100% rename from reme_ai/core/utils/__init__.py rename to reme_ai/core_old/utils/__init__.py diff --git a/reme_ai/core/utils/cache_handler.py b/reme_ai/core_old/utils/cache_handler.py similarity index 100% rename from reme_ai/core/utils/cache_handler.py rename to reme_ai/core_old/utils/cache_handler.py diff --git a/reme_ai/core/utils/case_converter.py b/reme_ai/core_old/utils/case_converter.py similarity index 100% rename from reme_ai/core/utils/case_converter.py rename to reme_ai/core_old/utils/case_converter.py diff --git a/reme_ai/core/utils/common_utils.py b/reme_ai/core_old/utils/common_utils.py similarity index 100% rename from reme_ai/core/utils/common_utils.py rename to reme_ai/core_old/utils/common_utils.py diff --git a/reme_ai/core/utils/env_utils.py b/reme_ai/core_old/utils/env_utils.py similarity index 100% rename from reme_ai/core/utils/env_utils.py rename to reme_ai/core_old/utils/env_utils.py diff --git a/reme_ai/core/utils/execute_tuils.py b/reme_ai/core_old/utils/execute_tuils.py similarity index 100% rename from reme_ai/core/utils/execute_tuils.py rename to reme_ai/core_old/utils/execute_tuils.py diff --git a/reme_ai/core/utils/http_client.py b/reme_ai/core_old/utils/http_client.py similarity index 100% rename from reme_ai/core/utils/http_client.py rename to reme_ai/core_old/utils/http_client.py diff --git a/reme_ai/core/utils/llm_utils.py b/reme_ai/core_old/utils/llm_utils.py similarity index 100% rename from reme_ai/core/utils/llm_utils.py rename to reme_ai/core_old/utils/llm_utils.py diff --git a/reme_ai/core/utils/logger_utils.py b/reme_ai/core_old/utils/logger_utils.py similarity index 100% rename from reme_ai/core/utils/logger_utils.py rename to reme_ai/core_old/utils/logger_utils.py diff --git a/reme_ai/core/utils/logo_utils.py b/reme_ai/core_old/utils/logo_utils.py similarity index 100% rename from reme_ai/core/utils/logo_utils.py rename to reme_ai/core_old/utils/logo_utils.py diff --git a/reme_ai/core/utils/mcp_client.py b/reme_ai/core_old/utils/mcp_client.py similarity index 100% rename from reme_ai/core/utils/mcp_client.py rename to reme_ai/core_old/utils/mcp_client.py diff --git a/reme_ai/core/utils/pydantic_config_parser.py b/reme_ai/core_old/utils/pydantic_config_parser.py similarity index 100% rename from reme_ai/core/utils/pydantic_config_parser.py rename to reme_ai/core_old/utils/pydantic_config_parser.py diff --git a/reme_ai/core/utils/pydantic_utils.py b/reme_ai/core_old/utils/pydantic_utils.py similarity index 100% rename from reme_ai/core/utils/pydantic_utils.py rename to reme_ai/core_old/utils/pydantic_utils.py diff --git a/reme_ai/core/utils/singleton.py b/reme_ai/core_old/utils/singleton.py similarity index 100% rename from reme_ai/core/utils/singleton.py rename to reme_ai/core_old/utils/singleton.py diff --git a/reme_ai/core/utils/time.py b/reme_ai/core_old/utils/time.py similarity index 100% rename from reme_ai/core/utils/time.py rename to reme_ai/core_old/utils/time.py diff --git a/reme_ai/core/vector_store/__init__.py b/reme_ai/core_old/vector_store/__init__.py similarity index 100% rename from reme_ai/core/vector_store/__init__.py rename to reme_ai/core_old/vector_store/__init__.py diff --git a/reme_ai/core/vector_store/base_vector_store.py b/reme_ai/core_old/vector_store/base_vector_store.py similarity index 97% rename from reme_ai/core/vector_store/base_vector_store.py rename to reme_ai/core_old/vector_store/base_vector_store.py index a4a8ca8e..158af0d9 100644 --- a/reme_ai/core/vector_store/base_vector_store.py +++ b/reme_ai/core_old/vector_store/base_vector_store.py @@ -5,9 +5,9 @@ from abc import ABC, abstractmethod from collections.abc import Callable 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 reme_ai.core_old.context import C +from reme_ai.core_old.embedding import BaseEmbeddingModel +from reme_ai.core_old.schema import VectorNode class BaseVectorStore(ABC): diff --git a/reme_ai/core/vector_store/chroma_vector_store.py b/reme_ai/core_old/vector_store/chroma_vector_store.py similarity index 99% rename from reme_ai/core/vector_store/chroma_vector_store.py rename to reme_ai/core_old/vector_store/chroma_vector_store.py index 3a483c23..567ac3b1 100644 --- a/reme_ai/core/vector_store/chroma_vector_store.py +++ b/reme_ai/core_old/vector_store/chroma_vector_store.py @@ -118,7 +118,7 @@ class ChromaVectorStore(BaseVectorStore): @staticmethod def _generate_where_clause(filters: dict | None) -> dict | None: """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 @@ -128,7 +128,7 @@ class ChromaVectorStore(BaseVectorStore): 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) diff --git a/reme_ai/core/vector_store/es_vector_store.py b/reme_ai/core_old/vector_store/es_vector_store.py similarity index 100% rename from reme_ai/core/vector_store/es_vector_store.py rename to reme_ai/core_old/vector_store/es_vector_store.py diff --git a/reme_ai/core/vector_store/local_vector_store.py b/reme_ai/core_old/vector_store/local_vector_store.py similarity index 99% rename from reme_ai/core/vector_store/local_vector_store.py rename to reme_ai/core_old/vector_store/local_vector_store.py index cfeec74e..7533fa89 100644 --- a/reme_ai/core/vector_store/local_vector_store.py +++ b/reme_ai/core_old/vector_store/local_vector_store.py @@ -92,7 +92,7 @@ class LocalVectorStore(BaseVectorStore): @staticmethod def _match_filters(node: VectorNode, filters: dict | None) -> bool: """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 diff --git a/reme_ai/core/vector_store/pgvector_store.py b/reme_ai/core_old/vector_store/pgvector_store.py similarity index 99% rename from reme_ai/core/vector_store/pgvector_store.py rename to reme_ai/core_old/vector_store/pgvector_store.py index 576d8e6c..882ced3e 100644 --- a/reme_ai/core/vector_store/pgvector_store.py +++ b/reme_ai/core_old/vector_store/pgvector_store.py @@ -29,7 +29,7 @@ class PGVectorStore(BaseVectorStore): @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 @@ -279,11 +279,11 @@ 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. - + 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: @@ -297,7 +297,7 @@ class PGVectorStore(BaseVectorStore): # 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] diff --git a/reme_ai/core/vector_store/qdrant_vector_store.py b/reme_ai/core_old/vector_store/qdrant_vector_store.py similarity index 99% rename from reme_ai/core/vector_store/qdrant_vector_store.py rename to reme_ai/core_old/vector_store/qdrant_vector_store.py index 5d9fa4f5..d0b4a8aa 100644 --- a/reme_ai/core/vector_store/qdrant_vector_store.py +++ b/reme_ai/core_old/vector_store/qdrant_vector_store.py @@ -247,7 +247,7 @@ class QdrantVectorStore(BaseVectorStore): @staticmethod def _create_filter(filters: dict) -> Filter | None: """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 @@ -295,7 +295,7 @@ class QdrantVectorStore(BaseVectorStore): f"Qdrant range filter for key '{key}' requires numeric lte value, got {type(value['lte']).__name__}. Skipping." ) continue - + if range_params: # Only add condition if we have valid numeric parameters conditions.append( FieldCondition( diff --git a/reme_ai/mem_agent/base_memory_agent.py b/reme_ai/mem_agent/base_memory_agent.py index d53cc8d9..5ba0cad3 100644 --- a/reme_ai/mem_agent/base_memory_agent.py +++ b/reme_ai/mem_agent/base_memory_agent.py @@ -6,9 +6,9 @@ from abc import ABCMeta from loguru import logger -from ..core.enumeration import Role, MemoryType -from ..core.op import BaseOp -from ..core.schema import Message, ToolCall, MemoryNode +from ..core_old.enumeration import Role, MemoryType +from ..core_old.op import BaseOp +from ..core_old.schema import Message, ToolCall, MemoryNode from ..mem_tool import BaseMemoryTool, ThinkTool @@ -155,13 +155,13 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): tool_call_id=op.tool_call.id, ) tool_result_messages.append(tool_message) - + # # 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 diff --git a/reme_ai/mem_agent/chat/remy_agent.py b/reme_ai/mem_agent/chat/remy_agent.py index 1eaa9ac7..c7806617 100644 --- a/reme_ai/mem_agent/chat/remy_agent.py +++ b/reme_ai/mem_agent/chat/remy_agent.py @@ -3,10 +3,10 @@ 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 +from ...core_old.context import C +from ...core_old.enumeration import Role +from ...core_old.schema import Message +from ...core_old.utils import get_now_time @C.register_op() diff --git a/reme_ai/mem_agent/chat/simple_chat.py b/reme_ai/mem_agent/chat/simple_chat.py index 8a71c7a8..9e09ffec 100644 --- a/reme_ai/mem_agent/chat/simple_chat.py +++ b/reme_ai/mem_agent/chat/simple_chat.py @@ -2,10 +2,10 @@ from loguru import logger -from ...core.context import C -from ...core.enumeration import Role -from ...core.op import BaseOp -from ...core.schema import Message, ToolCall +from ...core_old.context import C +from ...core_old.enumeration import Role +from ...core_old.op import BaseOp +from ...core_old.schema import Message, ToolCall @C.register_op() diff --git a/reme_ai/mem_agent/chat/stream_chat.py b/reme_ai/mem_agent/chat/stream_chat.py index 470e4647..2b121446 100644 --- a/reme_ai/mem_agent/chat/stream_chat.py +++ b/reme_ai/mem_agent/chat/stream_chat.py @@ -2,10 +2,10 @@ from loguru import logger -from ...core.context import C -from ...core.enumeration import Role, ChunkEnum -from ...core.op import BaseOp -from ...core.schema import Message, ToolCall +from ...core_old.context import C +from ...core_old.enumeration import Role, ChunkEnum +from ...core_old.op import BaseOp +from ...core_old.schema import Message, ToolCall @C.register_op() diff --git a/reme_ai/mem_agent/retriever/reme_retriever.py b/reme_ai/mem_agent/retriever/reme_retriever.py index f3700b1e..400e0d01 100644 --- a/reme_ai/mem_agent/retriever/reme_retriever.py +++ b/reme_ai/mem_agent/retriever/reme_retriever.py @@ -3,10 +3,10 @@ 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, format_messages +from ...core_old.context import C +from ...core_old.enumeration import Role +from ...core_old.schema import Message +from ...core_old.utils import get_now_time, format_messages @C.register_op() 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..fedfe188 100644 --- a/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.py +++ b/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.py @@ -3,16 +3,16 @@ 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 format_messages +from ...core_old.context import C +from ...core_old.enumeration import Role +from ...core_old.schema import Message +from ...core_old.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/identity_summarizer.py b/reme_ai/mem_agent/summarizer/identity_summarizer.py index be571cdc..0a9e1410 100644 --- a/reme_ai/mem_agent/summarizer/identity_summarizer.py +++ b/reme_ai/mem_agent/summarizer/identity_summarizer.py @@ -1,10 +1,10 @@ """Specialized agent for extracting and updating agent self-cognition memories.""" from ..base_memory_agent import BaseMemoryAgent -from ...core.context import C -from ...core.enumeration import Role, MemoryType -from ...core.schema import Message -from ...core.utils import get_now_time, format_messages +from ...core_old.context import C +from ...core_old.enumeration import Role, MemoryType +from ...core_old.schema import Message +from ...core_old.utils import get_now_time, format_messages @C.register_op() diff --git a/reme_ai/mem_agent/summarizer/personal_summarizer.py b/reme_ai/mem_agent/summarizer/personal_summarizer.py index 352fe1ee..1c32f499 100644 --- a/reme_ai/mem_agent/summarizer/personal_summarizer.py +++ b/reme_ai/mem_agent/summarizer/personal_summarizer.py @@ -1,10 +1,10 @@ """Specialized agent for extracting and managing personal memories about specific individuals.""" from ..base_memory_agent import BaseMemoryAgent -from ...core.context import C -from ...core.enumeration import Role, MemoryType -from ...core.schema import Message, ToolCall -from ...core.utils import get_now_time, format_messages +from ...core_old.context import C +from ...core_old.enumeration import Role, MemoryType +from ...core_old.schema import Message, ToolCall +from ...core_old.utils import get_now_time, format_messages @C.register_op() diff --git a/reme_ai/mem_agent/summarizer/procedural_summarizer.py b/reme_ai/mem_agent/summarizer/procedural_summarizer.py index e31422e4..24a75339 100644 --- a/reme_ai/mem_agent/summarizer/procedural_summarizer.py +++ b/reme_ai/mem_agent/summarizer/procedural_summarizer.py @@ -1,10 +1,10 @@ """Specialized agent for extracting and managing procedural knowledge and workflows.""" from ..base_memory_agent import BaseMemoryAgent -from ...core.context import C -from ...core.enumeration import Role, MemoryType -from ...core.schema import Message -from ...core.utils import get_now_time, format_messages +from ...core_old.context import C +from ...core_old.enumeration import Role, MemoryType +from ...core_old.schema import Message +from ...core_old.utils import get_now_time, format_messages @C.register_op() diff --git a/reme_ai/mem_agent/summarizer/reme_summarizer.py b/reme_ai/mem_agent/summarizer/reme_summarizer.py index 9eb06d84..5521762c 100644 --- a/reme_ai/mem_agent/summarizer/reme_summarizer.py +++ b/reme_ai/mem_agent/summarizer/reme_summarizer.py @@ -3,10 +3,10 @@ from loguru import logger from ..base_memory_agent import BaseMemoryAgent -from ...core.context import C -from ...core.enumeration import Role, MemoryType -from ...core.schema import Message, MemoryNode, ToolCall -from ...core.utils import get_now_time, format_messages +from ...core_old.context import C +from ...core_old.enumeration import Role, MemoryType +from ...core_old.schema import Message, MemoryNode, ToolCall +from ...core_old.utils import get_now_time, format_messages @C.register_op() @@ -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/tool_summarizer.py b/reme_ai/mem_agent/summarizer/tool_summarizer.py index 50399bba..1e9e33c0 100644 --- a/reme_ai/mem_agent/summarizer/tool_summarizer.py +++ b/reme_ai/mem_agent/summarizer/tool_summarizer.py @@ -1,10 +1,10 @@ """Specialized agent for extracting and managing tool usage guidelines and best practices.""" from ..base_memory_agent import BaseMemoryAgent -from ...core.context import C -from ...core.enumeration import Role, MemoryType -from ...core.schema import Message -from ...core.utils import get_now_time, format_messages +from ...core_old.context import C +from ...core_old.enumeration import Role, MemoryType +from ...core_old.schema import Message +from ...core_old.utils import get_now_time, format_messages @C.register_op() 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..cc2794b3 100644 --- a/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.py +++ b/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.py @@ -1,10 +1,10 @@ """Simplified personal memory summarizer using v2 memory tools.""" from ..base_memory_agent import BaseMemoryAgent -from ...core.context import C -from ...core.enumeration import Role, MemoryType -from ...core.schema import Message, ToolCall -from ...core.utils import format_messages +from ...core_old.context import C +from ...core_old.enumeration import Role, MemoryType +from ...core_old.schema import Message, ToolCall +from ...core_old.utils import format_messages @C.register_op() @@ -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.py b/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.py index 1aae4ad4..aa680da9 100644 --- a/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.py +++ b/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.py @@ -3,10 +3,10 @@ from loguru import logger from ..base_memory_agent import BaseMemoryAgent -from ...core.context import C -from ...core.enumeration import Role, MemoryType -from ...core.schema import Message, MemoryNode, ToolCall -from ...core.utils import format_messages +from ...core_old.context import C +from ...core_old.enumeration import Role, MemoryType +from ...core_old.schema import Message, MemoryNode, ToolCall +from ...core_old.utils import format_messages @C.register_op() 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/personal_summarizer_v3.py b/reme_ai/mem_agent/v3/personal_summarizer_v3.py index 0093884d..3f2c9ca9 100644 --- a/reme_ai/mem_agent/v3/personal_summarizer_v3.py +++ b/reme_ai/mem_agent/v3/personal_summarizer_v3.py @@ -1,7 +1,7 @@ from ..base_memory_agent import BaseMemoryAgent -from ...core.enumeration import Role, MemoryType -from ...core.schema import Message, ToolCall -from ...core.utils import format_messages +from ...core_old.enumeration import Role, MemoryType +from ...core_old.schema import Message, ToolCall +from ...core_old.utils import format_messages class PersonalSummarizerV3(BaseMemoryAgent): diff --git a/reme_ai/mem_agent/v3/reme_retriever_v3.py b/reme_ai/mem_agent/v3/reme_retriever_v3.py index 8f5c62dc..020b5ad2 100644 --- a/reme_ai/mem_agent/v3/reme_retriever_v3.py +++ b/reme_ai/mem_agent/v3/reme_retriever_v3.py @@ -3,9 +3,9 @@ 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 +from ...core_old.enumeration import Role +from ...core_old.schema import Message +from ...core_old.utils import format_messages class ReMeRetrieverV3(BaseMemoryAgent): diff --git a/reme_ai/mem_agent/v3/reme_retriever_v3.yaml b/reme_ai/mem_agent/v3/reme_retriever_v3.yaml index 7a3e575f..925bee68 100644 --- a/reme_ai/mem_agent/v3/reme_retriever_v3.yaml +++ b/reme_ai/mem_agent/v3/reme_retriever_v3.yaml @@ -26,13 +26,13 @@ system_prompt: | * 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)** diff --git a/reme_ai/mem_agent/v3/reme_summarizer_v3.py b/reme_ai/mem_agent/v3/reme_summarizer_v3.py index a0f466b9..3e1f17f7 100644 --- a/reme_ai/mem_agent/v3/reme_summarizer_v3.py +++ b/reme_ai/mem_agent/v3/reme_summarizer_v3.py @@ -1,9 +1,9 @@ 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 +from ...core_old.enumeration import Role, MemoryType +from ...core_old.schema import Message, MemoryNode, ToolCall +from ...core_old.utils import format_messages class ReMeSummarizerV3(BaseMemoryAgent): diff --git a/reme_ai/mem_agent/v3/reme_summarizer_v3.yaml b/reme_ai/mem_agent/v3/reme_summarizer_v3.yaml index 30792a08..58080cab 100644 --- a/reme_ai/mem_agent/v3/reme_summarizer_v3.yaml +++ b/reme_ai/mem_agent/v3/reme_summarizer_v3.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/v4/personal_retriever_v4.py b/reme_ai/mem_agent/v4/personal_retriever_v4.py index 5135acd9..2ba0dba5 100644 --- a/reme_ai/mem_agent/v4/personal_retriever_v4.py +++ b/reme_ai/mem_agent/v4/personal_retriever_v4.py @@ -1,7 +1,7 @@ from ..base_memory_agent import BaseMemoryAgent -from ...core.enumeration import Role, MemoryType -from ...core.schema import Message -from ...core.utils import format_messages +from ...core_old.enumeration import Role, MemoryType +from ...core_old.schema import Message +from ...core_old.utils import format_messages from ...mem_tool.v4 import ReadUserProfile @@ -41,7 +41,7 @@ class PersonalRetrieverV4(BaseMemoryAgent): 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: diff --git a/reme_ai/mem_agent/v4/personal_summarizer_v4.py b/reme_ai/mem_agent/v4/personal_summarizer_v4.py index c0e1d4f1..8cda9b24 100644 --- a/reme_ai/mem_agent/v4/personal_summarizer_v4.py +++ b/reme_ai/mem_agent/v4/personal_summarizer_v4.py @@ -1,8 +1,8 @@ from loguru import logger from ..base_memory_agent import BaseMemoryAgent -from ...core.enumeration import Role, MemoryType -from ...core.schema import Message, MemoryNode +from ...core_old.enumeration import Role, MemoryType +from ...core_old.schema import Message, MemoryNode class PersonalSummarizerV4(BaseMemoryAgent): diff --git a/reme_ai/mem_agent/v4/personal_summarizer_v4.yaml b/reme_ai/mem_agent/v4/personal_summarizer_v4.yaml index 5e7d12b7..3663c992 100644 --- a/reme_ai/mem_agent/v4/personal_summarizer_v4.yaml +++ b/reme_ai/mem_agent/v4/personal_summarizer_v4.yaml @@ -13,7 +13,7 @@ user_message_phase1: | **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) @@ -34,13 +34,13 @@ user_message_phase2: | {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 diff --git a/reme_ai/mem_agent/v4/reme_retriever_v4.py b/reme_ai/mem_agent/v4/reme_retriever_v4.py index 7dbf2b8a..db5f92ae 100644 --- a/reme_ai/mem_agent/v4/reme_retriever_v4.py +++ b/reme_ai/mem_agent/v4/reme_retriever_v4.py @@ -1,9 +1,9 @@ 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 +from ...core_old.enumeration import Role +from ...core_old.schema import Message +from ...core_old.utils import format_messages class ReMeRetrieverV4(BaseMemoryAgent): @@ -47,7 +47,7 @@ class ReMeRetrieverV4(BaseMemoryAgent): 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 [] @@ -98,16 +98,16 @@ class ReMeRetrieverV4(BaseMemoryAgent): 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 = [] diff --git a/reme_ai/mem_agent/v4/reme_summarizer_v4.py b/reme_ai/mem_agent/v4/reme_summarizer_v4.py index a4069c85..7787d035 100644 --- a/reme_ai/mem_agent/v4/reme_summarizer_v4.py +++ b/reme_ai/mem_agent/v4/reme_summarizer_v4.py @@ -1,9 +1,9 @@ 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 +from ...core_old.enumeration import Role, MemoryType +from ...core_old.schema import Message, MemoryNode +from ...core_old.utils import format_messages class ReMeSummarizerV4(BaseMemoryAgent): diff --git a/reme_ai/mem_agent/v4/reme_summarizer_v4.yaml b/reme_ai/mem_agent/v4/reme_summarizer_v4.yaml index a6d322ef..4cb6f54f 100644 --- a/reme_ai/mem_agent/v4/reme_summarizer_v4.yaml +++ b/reme_ai/mem_agent/v4/reme_summarizer_v4.yaml @@ -19,7 +19,7 @@ system_prompt: | - 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: | diff --git a/reme_ai/mem_agent/wk/personal_summarizer_wk.py b/reme_ai/mem_agent/wk/personal_summarizer_wk.py index e974f99a..c95ac36b 100644 --- a/reme_ai/mem_agent/wk/personal_summarizer_wk.py +++ b/reme_ai/mem_agent/wk/personal_summarizer_wk.py @@ -1,7 +1,7 @@ from ..base_memory_agent import BaseMemoryAgent -from ...core.enumeration import Role, MemoryType -from ...core.schema import Message, ToolCall -from ...core.utils import format_messages +from ...core_old.enumeration import Role, MemoryType +from ...core_old.schema import Message, ToolCall +from ...core_old.utils import format_messages class PersonalSummarizerWk(BaseMemoryAgent): diff --git a/reme_ai/mem_agent/wk/reme_retriever_wk.py b/reme_ai/mem_agent/wk/reme_retriever_wk.py index c98de937..21a5403a 100644 --- a/reme_ai/mem_agent/wk/reme_retriever_wk.py +++ b/reme_ai/mem_agent/wk/reme_retriever_wk.py @@ -3,9 +3,9 @@ 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 +from ...core_old.enumeration import Role +from ...core_old.schema import Message +from ...core_old.utils import format_messages class ReMeRetrieverV2(BaseMemoryAgent): 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.py b/reme_ai/mem_agent/wk/reme_summarizer_wk.py index 02a7dbf3..a04d230d 100644 --- a/reme_ai/mem_agent/wk/reme_summarizer_wk.py +++ b/reme_ai/mem_agent/wk/reme_summarizer_wk.py @@ -1,9 +1,9 @@ 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 +from ...core_old.enumeration import Role, MemoryType +from ...core_old.schema import Message, MemoryNode, ToolCall +from ...core_old.utils import format_messages class ReMeSummarizerWk(BaseMemoryAgent): 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 index 8b124496..827066b9 100644 --- a/reme_ai/mem_tool/base_memory_tool.py +++ b/reme_ai/mem_tool/base_memory_tool.py @@ -3,10 +3,10 @@ from abc import ABCMeta from pathlib import Path -from ..core.enumeration import MemoryType -from ..core.op import BaseOp -from ..core.schema import ToolCall, MemoryNode -from ..core.utils import CacheHandler +from ..core_old.enumeration import MemoryType +from ..core_old.op import BaseOp +from ..core_old.schema import ToolCall, MemoryNode +from ..core_old.utils import CacheHandler class BaseMemoryTool(BaseOp, metaclass=ABCMeta): diff --git a/reme_ai/mem_tool/hands_off_tool.py b/reme_ai/mem_tool/hands_off_tool.py index 2cb28d60..7ab65cd7 100644 --- a/reme_ai/mem_tool/hands_off_tool.py +++ b/reme_ai/mem_tool/hands_off_tool.py @@ -6,8 +6,8 @@ from typing import TYPE_CHECKING from loguru import logger from .base_memory_tool import BaseMemoryTool -from ..core.context import C -from ..core.enumeration import MemoryType +from ..core_old.context import C +from ..core_old.enumeration import MemoryType if TYPE_CHECKING: from ..mem_agent import BaseMemoryAgent diff --git a/reme_ai/mem_tool/history/add_history_memory.py b/reme_ai/mem_tool/history/add_history_memory.py index a92deca5..85e0181a 100644 --- a/reme_ai/mem_tool/history/add_history_memory.py +++ b/reme_ai/mem_tool/history/add_history_memory.py @@ -3,10 +3,10 @@ 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 +from ...core_old.context import C +from ...core_old.enumeration import MemoryType +from ...core_old.schema import ToolCall, Message +from ...core_old.utils import format_messages @C.register_op() diff --git a/reme_ai/mem_tool/history/read_history_memory.py b/reme_ai/mem_tool/history/read_history_memory.py index def2ff24..24cc0366 100644 --- a/reme_ai/mem_tool/history/read_history_memory.py +++ b/reme_ai/mem_tool/history/read_history_memory.py @@ -3,8 +3,8 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C -from ...core.schema import MemoryNode +from ...core_old.context import C +from ...core_old.schema import MemoryNode @C.register_op() diff --git a/reme_ai/mem_tool/identity/read_identity_memory.py b/reme_ai/mem_tool/identity/read_identity_memory.py index bd9f8031..dfede68f 100644 --- a/reme_ai/mem_tool/identity/read_identity_memory.py +++ b/reme_ai/mem_tool/identity/read_identity_memory.py @@ -3,7 +3,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C +from ...core_old.context import C @C.register_op() diff --git a/reme_ai/mem_tool/identity/update_identity_memory.py b/reme_ai/mem_tool/identity/update_identity_memory.py index b0211242..0883b1a9 100644 --- a/reme_ai/mem_tool/identity/update_identity_memory.py +++ b/reme_ai/mem_tool/identity/update_identity_memory.py @@ -3,7 +3,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C +from ...core_old.context import C @C.register_op() diff --git a/reme_ai/mem_tool/meta/add_meta_memory.py b/reme_ai/mem_tool/meta/add_meta_memory.py index d7b41254..99635698 100644 --- a/reme_ai/mem_tool/meta/add_meta_memory.py +++ b/reme_ai/mem_tool/meta/add_meta_memory.py @@ -5,8 +5,8 @@ import json from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C -from ...core.enumeration import MemoryType +from ...core_old.context import C +from ...core_old.enumeration import MemoryType @C.register_op() diff --git a/reme_ai/mem_tool/meta/read_meta_memory.py b/reme_ai/mem_tool/meta/read_meta_memory.py index 07ad1ecf..5ee59d54 100644 --- a/reme_ai/mem_tool/meta/read_meta_memory.py +++ b/reme_ai/mem_tool/meta/read_meta_memory.py @@ -3,8 +3,8 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C -from ...core.enumeration import MemoryType +from ...core_old.context import C +from ...core_old.enumeration import MemoryType @C.register_op() diff --git a/reme_ai/mem_tool/read_local_memories.py b/reme_ai/mem_tool/read_local_memories.py index 98d96187..f239d5d1 100644 --- a/reme_ai/mem_tool/read_local_memories.py +++ b/reme_ai/mem_tool/read_local_memories.py @@ -38,7 +38,7 @@ class ReadLocalMemories(BaseMemoryTool): 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) diff --git a/reme_ai/mem_tool/think_tool.py b/reme_ai/mem_tool/think_tool.py index 1d26446a..c54c7676 100644 --- a/reme_ai/mem_tool/think_tool.py +++ b/reme_ai/mem_tool/think_tool.py @@ -5,8 +5,8 @@ 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_old.context import C +from ..core_old.schema import ToolCall @C.register_op() diff --git a/reme_ai/mem_tool/v2/add_memory_drafts.py b/reme_ai/mem_tool/v2/add_memory_drafts.py index 93caad0c..b815be6e 100644 --- a/reme_ai/mem_tool/v2/add_memory_drafts.py +++ b/reme_ai/mem_tool/v2/add_memory_drafts.py @@ -3,7 +3,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C +from ...core_old.context import C @C.register_op() diff --git a/reme_ai/mem_tool/v2/read_history.py b/reme_ai/mem_tool/v2/read_history.py index 15aaada3..141989c0 100644 --- a/reme_ai/mem_tool/v2/read_history.py +++ b/reme_ai/mem_tool/v2/read_history.py @@ -3,20 +3,20 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C -from ...core.schema import MemoryNode +from ...core_old.context import C +from ...core_old.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.py b/reme_ai/mem_tool/v2/retrieve_memories.py index abdfc377..96d4ca99 100644 --- a/reme_ai/mem_tool/v2/retrieve_memories.py +++ b/reme_ai/mem_tool/v2/retrieve_memories.py @@ -3,9 +3,9 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C -from ...core.schema import MemoryNode, VectorNode -from ...core.utils import deduplicate_memories +from ...core_old.context import C +from ...core_old.schema import MemoryNode, VectorNode +from ...core_old.utils import deduplicate_memories @C.register_op() 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.py b/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.py index 38107ab0..cfed10c8 100644 --- a/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.py +++ b/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.py @@ -3,9 +3,9 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C -from ...core.schema import MemoryNode, VectorNode -from ...core.utils import deduplicate_memories +from ...core_old.context import C +from ...core_old.schema import MemoryNode, VectorNode +from ...core_old.utils import deduplicate_memories @C.register_op() 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/v2/summary_and_hands_off.py b/reme_ai/mem_tool/v2/summary_and_hands_off.py index 131b1101..631c7d6b 100644 --- a/reme_ai/mem_tool/v2/summary_and_hands_off.py +++ b/reme_ai/mem_tool/v2/summary_and_hands_off.py @@ -6,9 +6,9 @@ from typing import TYPE_CHECKING from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C -from ...core.enumeration import MemoryType -from ...core.schema import MemoryNode, Message +from ...core_old.context import C +from ...core_old.enumeration import MemoryType +from ...core_old.schema import MemoryNode, Message if TYPE_CHECKING: from ...mem_agent import BaseMemoryAgent diff --git a/reme_ai/mem_tool/v2/update_memories.py b/reme_ai/mem_tool/v2/update_memories.py index 03a9f494..cffe74ac 100644 --- a/reme_ai/mem_tool/v2/update_memories.py +++ b/reme_ai/mem_tool/v2/update_memories.py @@ -3,8 +3,8 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C -from ...core.schema import MemoryNode +from ...core_old.context import C +from ...core_old.schema import MemoryNode @C.register_op() diff --git a/reme_ai/mem_tool/v3/add_memory.py b/reme_ai/mem_tool/v3/add_memory.py index ee488639..4893c6a2 100644 --- a/reme_ai/mem_tool/v3/add_memory.py +++ b/reme_ai/mem_tool/v3/add_memory.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.schema import MemoryNode +from ...core_old.schema import MemoryNode class AddMemory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/v3/read_history.py b/reme_ai/mem_tool/v3/read_history.py index e9ab2a15..9506d88a 100644 --- a/reme_ai/mem_tool/v3/read_history.py +++ b/reme_ai/mem_tool/v3/read_history.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.schema import MemoryNode +from ...core_old.schema import MemoryNode class ReadHistory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/v3/read_user_profile.py b/reme_ai/mem_tool/v3/read_user_profile.py index 3dba2bcf..bab1d413 100644 --- a/reme_ai/mem_tool/v3/read_user_profile.py +++ b/reme_ai/mem_tool/v3/read_user_profile.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.schema.memory_node import MemoryNode +from ...core_old.schema.memory_node import MemoryNode class ReadUserProfile(BaseMemoryTool): diff --git a/reme_ai/mem_tool/v3/retrieve_memory.py b/reme_ai/mem_tool/v3/retrieve_memory.py index 32e526d7..d5b9a7bc 100644 --- a/reme_ai/mem_tool/v3/retrieve_memory.py +++ b/reme_ai/mem_tool/v3/retrieve_memory.py @@ -3,8 +3,8 @@ import json from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.schema import MemoryNode -from ...core.utils import deduplicate_memories +from ...core_old.schema import MemoryNode +from ...core_old.utils import deduplicate_memories class RetrieveMemory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/v3/summary_and_hands_off.py b/reme_ai/mem_tool/v3/summary_and_hands_off.py index 19d88744..e0b2756c 100644 --- a/reme_ai/mem_tool/v3/summary_and_hands_off.py +++ b/reme_ai/mem_tool/v3/summary_and_hands_off.py @@ -4,8 +4,8 @@ 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 +from ...core_old.enumeration import MemoryType +from ...core_old.schema import MemoryNode, Message if TYPE_CHECKING: from ...mem_agent import BaseMemoryAgent diff --git a/reme_ai/mem_tool/v3/update_user_profile.py b/reme_ai/mem_tool/v3/update_user_profile.py index 46879e2b..ef47f46b 100644 --- a/reme_ai/mem_tool/v3/update_user_profile.py +++ b/reme_ai/mem_tool/v3/update_user_profile.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.schema.memory_node import MemoryNode +from ...core_old.schema.memory_node import MemoryNode class UpdateUserProfile(BaseMemoryTool): diff --git a/reme_ai/mem_tool/v4/add_summary_memory.py b/reme_ai/mem_tool/v4/add_summary_memory.py index cc4be602..6fe30cda 100644 --- a/reme_ai/mem_tool/v4/add_summary_memory.py +++ b/reme_ai/mem_tool/v4/add_summary_memory.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.schema import MemoryNode +from ...core_old.schema import MemoryNode class AddSummaryMemory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/v4/hands_off.py b/reme_ai/mem_tool/v4/hands_off.py index 17a33429..2dab4531 100644 --- a/reme_ai/mem_tool/v4/hands_off.py +++ b/reme_ai/mem_tool/v4/hands_off.py @@ -3,8 +3,8 @@ 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 +from ...core_old.enumeration import MemoryType +from ...core_old.schema import Message if TYPE_CHECKING: from ...mem_agent import BaseMemoryAgent @@ -62,14 +62,14 @@ class HandsOff(BaseMemoryTool): 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, diff --git a/reme_ai/mem_tool/v4/read_history.py b/reme_ai/mem_tool/v4/read_history.py index 78a90eb4..00097234 100644 --- a/reme_ai/mem_tool/v4/read_history.py +++ b/reme_ai/mem_tool/v4/read_history.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.schema import MemoryNode +from ...core_old.schema import MemoryNode class ReadHistory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/v4/read_user_profile.py b/reme_ai/mem_tool/v4/read_user_profile.py index aa8a57f8..f45c0c34 100644 --- a/reme_ai/mem_tool/v4/read_user_profile.py +++ b/reme_ai/mem_tool/v4/read_user_profile.py @@ -2,7 +2,7 @@ from typing import Literal from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.schema.memory_node import MemoryNode +from ...core_old.schema.memory_node import MemoryNode class ReadUserProfile(BaseMemoryTool): @@ -39,11 +39,11 @@ class ReadUserProfile(BaseMemoryTool): "required": [], } - async def execute(self): + async def execute(self): # Determine which IDs to show show_profile_id = self.show_ids in ("both", "profile") show_history_id = self.show_ids in ("both", "history") - + cache_key = f"{self.memory_type}_{self.memory_target}".replace(" ", "_").lower() cached_data = self.meta_memory.load(cache_key, auto_clean=False) @@ -58,22 +58,22 @@ class ReadUserProfile(BaseMemoryTool): memory_formated = [] for node in memory_nodes: node_formated_parts = [] - + # Add profile_id if enabled if show_profile_id: node_formated_parts.append(f"profile_id={node.memory_id}") - + # Always add profile_content node_formated_parts.append(f"profile_content={node.content}") - + # Add conversation_time if available if "conversation_time" in node.metadata and node.metadata["conversation_time"]: node_formated_parts.append(f"conversation_time={node.metadata['conversation_time']}") - + # Add history_id if enabled and available if show_history_id and node.ref_memory_id: node_formated_parts.append(f"history_id={node.ref_memory_id}") - + node_formated = " ".join(node_formated_parts) memory_formated.append(node_formated.strip()) diff --git a/reme_ai/mem_tool/v4/retrieve_memory.py b/reme_ai/mem_tool/v4/retrieve_memory.py index a8138bdc..a6717bfa 100644 --- a/reme_ai/mem_tool/v4/retrieve_memory.py +++ b/reme_ai/mem_tool/v4/retrieve_memory.py @@ -3,8 +3,8 @@ import json from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.schema import MemoryNode -from ...core.utils import deduplicate_memories +from ...core_old.schema import MemoryNode +from ...core_old.utils import deduplicate_memories class RetrieveMemory(BaseMemoryTool): @@ -62,7 +62,7 @@ class RetrieveMemory(BaseMemoryTool): 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: diff --git a/reme_ai/mem_tool/v4/update_user_profile.py b/reme_ai/mem_tool/v4/update_user_profile.py index a8fa04f5..1ab3276c 100644 --- a/reme_ai/mem_tool/v4/update_user_profile.py +++ b/reme_ai/mem_tool/v4/update_user_profile.py @@ -1,8 +1,8 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.schema.memory_node import MemoryNode -from ...core.utils import deduplicate_memories +from ...core_old.schema.memory_node import MemoryNode +from ...core_old.utils import deduplicate_memories class UpdateUserProfile(BaseMemoryTool): diff --git a/reme_ai/mem_tool/vector_store/add_memory.py b/reme_ai/mem_tool/vector_store/add_memory.py index 1f87df4b..0937fe2a 100644 --- a/reme_ai/mem_tool/vector_store/add_memory.py +++ b/reme_ai/mem_tool/vector_store/add_memory.py @@ -3,8 +3,8 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C -from ...core.schema import MemoryNode +from ...core_old.context import C +from ...core_old.schema import MemoryNode @C.register_op() diff --git a/reme_ai/mem_tool/vector_store/add_summary_memory.py b/reme_ai/mem_tool/vector_store/add_summary_memory.py index ce1127ed..54abbd57 100644 --- a/reme_ai/mem_tool/vector_store/add_summary_memory.py +++ b/reme_ai/mem_tool/vector_store/add_summary_memory.py @@ -3,9 +3,9 @@ from loguru import logger from .add_memory import AddMemory -from ...core.context import C -from ...core.enumeration import MemoryType -from ...core.schema import MemoryNode +from ...core_old.context import C +from ...core_old.enumeration import MemoryType +from ...core_old.schema import MemoryNode @C.register_op() diff --git a/reme_ai/mem_tool/vector_store/delete_memory.py b/reme_ai/mem_tool/vector_store/delete_memory.py index 45b28632..95f4f210 100644 --- a/reme_ai/mem_tool/vector_store/delete_memory.py +++ b/reme_ai/mem_tool/vector_store/delete_memory.py @@ -3,7 +3,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C +from ...core_old.context import C @C.register_op() diff --git a/reme_ai/mem_tool/vector_store/retrieve_recent_memory.py b/reme_ai/mem_tool/vector_store/retrieve_recent_memory.py index 0896d75a..dfaa9240 100644 --- a/reme_ai/mem_tool/vector_store/retrieve_recent_memory.py +++ b/reme_ai/mem_tool/vector_store/retrieve_recent_memory.py @@ -3,9 +3,9 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C -from ...core.schema import MemoryNode, VectorNode -from ...core.utils import deduplicate_memories +from ...core_old.context import C +from ...core_old.schema import MemoryNode, VectorNode +from ...core_old.utils import deduplicate_memories @C.register_op() diff --git a/reme_ai/mem_tool/vector_store/update_memory.py b/reme_ai/mem_tool/vector_store/update_memory.py index 4873ce24..fe08fb90 100644 --- a/reme_ai/mem_tool/vector_store/update_memory.py +++ b/reme_ai/mem_tool/vector_store/update_memory.py @@ -3,8 +3,8 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C -from ...core.schema import MemoryNode +from ...core_old.context import C +from ...core_old.schema import MemoryNode @C.register_op() diff --git a/reme_ai/mem_tool/vector_store/vector_retrieve_memory.py b/reme_ai/mem_tool/vector_store/vector_retrieve_memory.py index 655de36e..24b05e8b 100644 --- a/reme_ai/mem_tool/vector_store/vector_retrieve_memory.py +++ b/reme_ai/mem_tool/vector_store/vector_retrieve_memory.py @@ -3,10 +3,10 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.context import C -from ...core.enumeration import MemoryType -from ...core.schema import MemoryNode, VectorNode -from ...core.utils import deduplicate_memories +from ...core_old.context import C +from ...core_old.enumeration import MemoryType +from ...core_old.schema import MemoryNode, VectorNode +from ...core_old.utils import deduplicate_memories @C.register_op() diff --git a/reme_ai/mem_tool/wk/add_memory.py b/reme_ai/mem_tool/wk/add_memory.py index 61cb154c..2af7ba71 100644 --- a/reme_ai/mem_tool/wk/add_memory.py +++ b/reme_ai/mem_tool/wk/add_memory.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.schema import MemoryNode +from ...core_old.schema import MemoryNode class AddMemory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/wk/read_history.py b/reme_ai/mem_tool/wk/read_history.py index 945c02e2..65e8a0cf 100644 --- a/reme_ai/mem_tool/wk/read_history.py +++ b/reme_ai/mem_tool/wk/read_history.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.schema import MemoryNode +from ...core_old.schema import MemoryNode class ReadHistory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/wk/summary_and_hands_off.py b/reme_ai/mem_tool/wk/summary_and_hands_off.py index 9a6a76cd..d384f16a 100644 --- a/reme_ai/mem_tool/wk/summary_and_hands_off.py +++ b/reme_ai/mem_tool/wk/summary_and_hands_off.py @@ -4,8 +4,8 @@ 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 +from ...core_old.enumeration import MemoryType +from ...core_old.schema import MemoryNode, Message if TYPE_CHECKING: from ...mem_agent import BaseMemoryAgent diff --git a/reme_ai/mem_tool/wk/update_memory.py b/reme_ai/mem_tool/wk/update_memory.py index c419261f..151eba53 100644 --- a/reme_ai/mem_tool/wk/update_memory.py +++ b/reme_ai/mem_tool/wk/update_memory.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.schema import MemoryNode +from ...core_old.schema import MemoryNode class UpdateMemory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/wk/vector_retrieve_memory.py b/reme_ai/mem_tool/wk/vector_retrieve_memory.py index de7498b4..2698cd6f 100644 --- a/reme_ai/mem_tool/wk/vector_retrieve_memory.py +++ b/reme_ai/mem_tool/wk/vector_retrieve_memory.py @@ -1,9 +1,9 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core.enumeration import MemoryType -from ...core.schema import MemoryNode, VectorNode -from ...core.utils import deduplicate_memories +from ...core_old.enumeration import MemoryType +from ...core_old.schema import MemoryNode, VectorNode +from ...core_old.utils import deduplicate_memories class VectorRetrieveMemory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/write_local_memories.py b/reme_ai/mem_tool/write_local_memories.py index ac1c5aae..f6336395 100644 --- a/reme_ai/mem_tool/write_local_memories.py +++ b/reme_ai/mem_tool/write_local_memories.py @@ -30,13 +30,13 @@ class WriteLocalMemories(BaseMemoryTool): 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) @@ -45,11 +45,11 @@ class WriteLocalMemories(BaseMemoryTool): 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}") diff --git a/reme_ai/tool/execute/execute_code.py b/reme_ai/tool/execute/execute_code.py index ea259487..f13aab3d 100644 --- a/reme_ai/tool/execute/execute_code.py +++ b/reme_ai/tool/execute/execute_code.py @@ -4,11 +4,11 @@ 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.schema import ToolCall +from ...core_old.context import C +from ...core_old.op import BaseOp +from ...core_old.schema import ToolCall -from ...core.utils import exec_code +from ...core_old.utils import exec_code @C.register_op() diff --git a/reme_ai/tool/execute/execute_shell.py b/reme_ai/tool/execute/execute_shell.py index 6e244ddb..1b235921 100644 --- a/reme_ai/tool/execute/execute_shell.py +++ b/reme_ai/tool/execute/execute_shell.py @@ -4,11 +4,11 @@ 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.schema import ToolCall +from ...core_old.context import C +from ...core_old.op import BaseOp +from ...core_old.schema import ToolCall -from ...core.utils import run_shell_command +from ...core_old.utils import run_shell_command @C.register_op() diff --git a/reme_ai/tool/search/dashscope_search.py b/reme_ai/tool/search/dashscope_search.py index 19bd8104..47c399ef 100644 --- a/reme_ai/tool/search/dashscope_search.py +++ b/reme_ai/tool/search/dashscope_search.py @@ -9,9 +9,9 @@ from typing import Literal from loguru import logger -from ...core.context import C -from ...core.op import BaseOp -from ...core.schema import ToolCall +from ...core_old.context import C +from ...core_old.op import BaseOp +from ...core_old.schema import ToolCall @C.register_op() diff --git a/reme_ai/tool/search/mock_search.py b/reme_ai/tool/search/mock_search.py index 187463dc..69ef3995 100644 --- a/reme_ai/tool/search/mock_search.py +++ b/reme_ai/tool/search/mock_search.py @@ -9,11 +9,11 @@ import random from loguru import logger -from ...core.context import C -from ...core.enumeration import Role -from ...core.op import BaseOp -from ...core.schema import ToolCall, Message -from ...core.utils import extract_content +from ...core_old.context import C +from ...core_old.enumeration import Role +from ...core_old.op import BaseOp +from ...core_old.schema import ToolCall, Message +from ...core_old.utils import extract_content @C.register_op() diff --git a/reme_ai/tool/search/tavily_search.py b/reme_ai/tool/search/tavily_search.py index 5c194bdc..bb000f16 100644 --- a/reme_ai/tool/search/tavily_search.py +++ b/reme_ai/tool/search/tavily_search.py @@ -9,9 +9,9 @@ import os from loguru import logger -from ...core.context import C -from ...core.op import BaseOp -from ...core.schema import ToolCall +from ...core_old.context import C +from ...core_old.op import BaseOp +from ...core_old.schema import ToolCall @C.register_op() 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/tests/mcp_servers_demo.json b/test/mcp_servers_demo.json similarity index 100% rename from tests/mcp_servers_demo.json rename to test/mcp_servers_demo.json 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/tests/test_base_context.py b/test/test_base_context.py similarity index 97% rename from tests/test_base_context.py rename to test/test_base_context.py index 316a6796..2b3ebc99 100644 --- a/tests/test_base_context.py +++ b/test/test_base_context.py @@ -4,7 +4,7 @@ Ensures attribute-style and dict-style access work interchangeably. """ import pickle -from reme_ai.core.context import BaseContext +from reme_ai.core_old.context import BaseContext def test_attribute_access(): diff --git a/tests/test_cache_handler.py b/test/test_cache_handler.py similarity index 97% rename from tests/test_cache_handler.py rename to test/test_cache_handler.py index ddcac86f..b741769a 100644 --- a/tests/test_cache_handler.py +++ b/test/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_ai.core_old.utils.cache_handler import CacheHandler def run_tests(): diff --git a/tests/test_embedding.py b/test/test_embedding.py similarity index 98% rename from tests/test_embedding.py rename to test/test_embedding.py index d1c404f5..b769d05f 100644 --- a/tests/test_embedding.py +++ b/test/test_embedding.py @@ -18,12 +18,12 @@ import asyncio import argparse from typing import Type, List -from reme_ai.core.utils import load_env +from reme_ai.core_old.utils import load_env load_env() -from reme_ai.core.embedding import OpenAIEmbeddingModel, BaseEmbeddingModel -from reme_ai.core.schema import VectorNode +from reme_ai.core_old.embedding import OpenAIEmbeddingModel, BaseEmbeddingModel +from reme_ai.core_old.schema import VectorNode def get_embedding_model(model_class: Type[BaseEmbeddingModel]) -> BaseEmbeddingModel: diff --git a/tests/test_embedding_sync.py b/test/test_embedding_sync.py similarity index 98% rename from tests/test_embedding_sync.py rename to test/test_embedding_sync.py index 361a42b3..f97e28ec 100644 --- a/tests/test_embedding_sync.py +++ b/test/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_ai.core_old.utils import load_env load_env() -from reme_ai.core.embedding import OpenAIEmbeddingModelSync, BaseEmbeddingModel -from reme_ai.core.schema import VectorNode +from reme_ai.core_old.embedding import OpenAIEmbeddingModelSync, BaseEmbeddingModel +from reme_ai.core_old.schema import VectorNode def get_embedding_model(model_class: Type[BaseEmbeddingModel]) -> BaseEmbeddingModel: diff --git a/tests/test_llm.py b/test/test_llm.py similarity index 98% rename from tests/test_llm.py rename to test/test_llm.py index 12c6eca3..819c2b2b 100644 --- a/tests/test_llm.py +++ b/test/test_llm.py @@ -18,13 +18,13 @@ import asyncio import argparse from typing import Type -from reme_ai.core.utils import load_env +from reme_ai.core_old.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_ai.core_old.llm import OpenAILLM, LiteLLM, BaseLLM +from reme_ai.core_old.schema import Message, ToolCall +from reme_ai.core_old.enumeration import Role, ChunkEnum def get_llm(llm_class: Type[BaseLLM]) -> BaseLLM: diff --git a/tests/test_llm_sync.py b/test/test_llm_sync.py similarity index 98% rename from tests/test_llm_sync.py rename to test/test_llm_sync.py index 98751f87..07d80ed1 100644 --- a/tests/test_llm_sync.py +++ b/test/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_ai.core_old.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_ai.core_old.llm import OpenAILLMSync, LiteLLMSync, BaseLLM +from reme_ai.core_old.schema import Message, ToolCall +from reme_ai.core_old.enumeration import Role, ChunkEnum def get_llm(llm_class: Type[BaseLLM]) -> BaseLLM: diff --git a/tests/test_logo.py b/test/test_logo.py similarity index 59% rename from tests/test_logo.py rename to test/test_logo.py index eeede81e..9b4cf089 100644 --- a/tests/test_logo.py +++ b/test/test_logo.py @@ -1,9 +1,9 @@ """test logo""" -from reme_ai.core.schema import ServiceConfig, MCPConfig +from reme_ai.core_old.schema import ServiceConfig, MCPConfig if __name__ == "__main__": - from reme_ai.core.utils import print_logo + from reme_ai.core_old.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/test/test_mcp_client.py similarity index 99% rename from tests/test_mcp_client.py rename to test/test_mcp_client.py index d2fef40e..0ae6fc54 100644 --- a/tests/test_mcp_client.py +++ b/test/test_mcp_client.py @@ -5,7 +5,7 @@ import asyncio import json -from reme_ai.core.utils import MCPClient +from reme_ai.core_old.utils import MCPClient async def main(): diff --git a/tests/test_mcp_server.py b/test/test_mcp_server.py similarity index 97% rename from tests/test_mcp_server.py rename to test/test_mcp_server.py index 67f9c542..4257b977 100644 --- a/tests/test_mcp_server.py +++ b/test/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_ai.core_old.schema import ToolCall +from reme_ai.core_old.utils import create_pydantic_model mcp = FastMCP("DynamicSchemaServer", port=8010) diff --git a/tests/test_memory_vector_conversion.py b/test/test_memory_vector_conversion.py similarity index 100% rename from tests/test_memory_vector_conversion.py rename to test/test_memory_vector_conversion.py diff --git a/tests/test_message.py b/test/test_message.py similarity index 98% rename from tests/test_message.py rename to test/test_message.py index 141174c5..e77e673b 100644 --- a/tests/test_message.py +++ b/test/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_ai.core_old.enumeration import Role +from reme_ai.core_old.schema import ToolAttr, ToolCall, ContentBlock, Message class TestModelDefinitions(unittest.TestCase): diff --git a/tests/test_op_composition.py b/test/test_op_composition.py similarity index 99% rename from tests/test_op_composition.py rename to test/test_op_composition.py index 8d32b51a..319981a3 100644 --- a/tests/test_op_composition.py +++ b/test/test_op_composition.py @@ -5,8 +5,8 @@ Tests asynchronous execution mode. import asyncio -from reme_ai.core.op import BaseOp -from reme_ai.core.schema import ToolCall, ToolAttr +from reme_ai.core_old.op import BaseOp +from reme_ai.core_old.schema import ToolCall, ToolAttr class AddOp(BaseOp): diff --git a/tests/test_reme.py b/test/test_reme.py similarity index 98% rename from tests/test_reme.py rename to test/test_reme.py index c9e6843e..b23aa63e 100644 --- a/tests/test_reme.py +++ b/test/test_reme.py @@ -2,7 +2,7 @@ import asyncio -from reme_ai.core.schema import VectorNode, MemoryNode +from reme_ai.core_old.schema import VectorNode, MemoryNode from reme_ai.reme import ReMe reme = ReMe( diff --git a/tests/test_timer.py b/test/test_timer.py similarity index 97% rename from tests/test_timer.py rename to test/test_timer.py index c9714e38..1c9de938 100644 --- a/tests/test_timer.py +++ b/test/test_timer.py @@ -7,7 +7,7 @@ import time from loguru import logger -from reme_ai.core.utils import timer +from reme_ai.core_old.utils import timer @timer diff --git a/tests/test_token_counter.py b/test/test_token_counter.py similarity index 99% rename from tests/test_token_counter.py rename to test/test_token_counter.py index 3c44a298..e67fb5c2 100644 --- a/tests/test_token_counter.py +++ b/test/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_ai.core_old.enumeration import Role +from reme_ai.core_old.schema import Message, ToolCall +from reme_ai.core_old.token_counter import BaseTokenCounter, OpenAITokenCounter, HFTokenCounter def get_token_counter(counter_class: Type[BaseTokenCounter], **kwargs) -> BaseTokenCounter: diff --git a/tests/test_tool.py b/test/test_tool.py similarity index 98% rename from tests/test_tool.py rename to test/test_tool.py index 9db5a3ed..765051b3 100644 --- a/tests/test_tool.py +++ b/test/test_tool.py @@ -169,8 +169,8 @@ async def test_stream_chat(): 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_ai.core_old.utils import execute_stream_task + from reme_ai.core_old.context import RuntimeContext from asyncio import Queue op = StreamChat() diff --git a/tests/test_tool_call.py b/test/test_tool_call.py similarity index 99% rename from tests/test_tool_call.py rename to test/test_tool_call.py index 30c0a37e..2ae7a619 100644 --- a/tests/test_tool_call.py +++ b/test/test_tool_call.py @@ -2,7 +2,7 @@ import json -from reme_ai.core.schema.tool_call import ToolCall +from reme_ai.core_old.schema.tool_call import ToolCall def test_simple_schema(): 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_vector_store.py b/test/test_vector_store.py similarity index 99% rename from tests/test_vector_store.py rename to test/test_vector_store.py index 6ef9ab0e..51edbd55 100644 --- a/tests/test_vector_store.py +++ b/test/test_vector_store.py @@ -23,9 +23,9 @@ 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_ai.core_old.embedding import OpenAIEmbeddingModel +from reme_ai.core_old.schema import VectorNode +from reme_ai.core_old.vector_store import ( BaseVectorStore, ChromaVectorStore, LocalVectorStore, @@ -388,7 +388,8 @@ async def test_search_with_multiple_filters(store: BaseVectorStore, _store_name: ) logger.info( - f"Multi-filter search (node_type=tech AND source=research AND priority=high) " 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") @@ -1452,11 +1453,11 @@ async def test_sql_injection_protection(store: BaseVectorStore, store_name: str) # Test 1: Invalid collection name (SQL injection attempt) try: - from reme_ai.core.vector_store import PGVectorStore - from reme_ai.core.embedding import OpenAIEmbeddingModel - + from reme_ai.core_old.vector_store import PGVectorStore + from reme_ai.core_old.embedding import OpenAIEmbeddingModel + embedding_model = OpenAIEmbeddingModel() - + # This should raise ValueError due to invalid table name try: invalid_store = PGVectorStore( @@ -1467,7 +1468,7 @@ async def test_sql_injection_protection(store: BaseVectorStore, store_name: str) assert False, "Should have raised ValueError for invalid collection name" except ValueError as e: logger.info(f"✓ Invalid collection name rejected: {e}") - + # Test 2: Invalid metadata key in filters try: results = await store.search( @@ -1481,9 +1482,9 @@ async def test_sql_injection_protection(store: BaseVectorStore, store_name: str) 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 From 32ff65ca0dc8150dbe3caaf1f2763aebf4dc139d Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 21 Jan 2026 17:05:56 +0800 Subject: [PATCH 10/19] feat(core): refactor context management and add schema definitions --- reme_ai/core/context/__init__.py | 9 +- reme_ai/core/context/prompt_handler.py | 106 ++++---- reme_ai/core/context/registry.py | 143 ----------- reme_ai/core/context/registry_factory.py | 45 ++++ reme_ai/core/context/runtime_context.py | 79 ------ reme_ai/core/schema/__init__.py | 42 ++++ reme_ai/core/schema/memory_node.py | 198 +++++++++++++++ reme_ai/core/schema/message.py | 165 +++++++++++++ reme_ai/core/schema/request.py | 11 + reme_ai/core/schema/response.py | 11 + reme_ai/core/schema/service_config.py | 113 +++++++++ reme_ai/core/schema/stream_chunk.py | 14 ++ reme_ai/core/schema/tool_call.py | 226 ++++++++++++++++++ reme_ai/core/schema/vector_node.py | 15 ++ reme_ai/core/utils/__init__.py | 7 + reme_ai/{core_old => core}/utils/singleton.py | 0 16 files changed, 897 insertions(+), 287 deletions(-) delete mode 100644 reme_ai/core/context/registry.py create mode 100644 reme_ai/core/context/registry_factory.py delete mode 100644 reme_ai/core/context/runtime_context.py create mode 100644 reme_ai/core/schema/__init__.py create mode 100644 reme_ai/core/schema/memory_node.py create mode 100644 reme_ai/core/schema/message.py create mode 100644 reme_ai/core/schema/request.py create mode 100644 reme_ai/core/schema/response.py create mode 100644 reme_ai/core/schema/service_config.py create mode 100644 reme_ai/core/schema/stream_chunk.py create mode 100644 reme_ai/core/schema/tool_call.py create mode 100644 reme_ai/core/schema/vector_node.py create mode 100644 reme_ai/core/utils/__init__.py rename reme_ai/{core_old => core}/utils/singleton.py (100%) diff --git a/reme_ai/core/context/__init__.py b/reme_ai/core/context/__init__.py index 7f26d600..27957fd2 100644 --- a/reme_ai/core/context/__init__.py +++ b/reme_ai/core/context/__init__.py @@ -2,15 +2,10 @@ from .base_context import BaseContext from .prompt_handler import PromptHandler -from .registry import Registry -from .runtime_context import RuntimeContext -from .service_context import ServiceContext, C +from .registry_factory import R __all__ = [ "BaseContext", "PromptHandler", - "Registry", - "RuntimeContext", - "ServiceContext", - "C", + "R", ] diff --git a/reme_ai/core/context/prompt_handler.py b/reme_ai/core/context/prompt_handler.py index 9639cbf8..e6b6d737 100644 --- a/reme_ai/core/context/prompt_handler.py +++ b/reme_ai/core/context/prompt_handler.py @@ -28,25 +28,24 @@ class PromptNotFoundError(KeyError): super().__init__( f"Prompt '{prompt_name}' not found. " f"Available prompts: {', '.join(available_prompts[:10])}" - f"{'...' if len(available_prompts) > 10 else ''}" + f"{'...' if len(available_prompts) > 10 else ''}", ) class PromptFormattingError(ValueError): """Exception raised when prompt formatting fails.""" - pass 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({ @@ -61,7 +60,7 @@ class PromptHandler(BaseContext): 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 @@ -72,24 +71,24 @@ class PromptHandler(BaseContext): self.language: str = language.strip() def load_prompt_by_file( - self, - prompt_file_path: Optional[Union[Path, str]] = None, - overwrite: bool = True + 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. @@ -115,8 +114,7 @@ class PromptHandler(BaseContext): prompt_dict = json.load(f) else: raise ValueError( - f"Unsupported file format: {suffix}. " - f"Supported formats: .yaml, .yml, .json" + f"Unsupported file format: {suffix}. " f"Supported formats: .yaml, .yml, .json", ) logger.info(f"Loaded {len(prompt_dict or {})} prompts from {prompt_file_path}") @@ -125,23 +123,23 @@ class PromptHandler(BaseContext): 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 + 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. """ @@ -156,8 +154,7 @@ class PromptHandler(BaseContext): if key in self: if overwrite: logger.warning( - f"Overwriting prompt '{key}': " - f"old length={len(self[key])}, new length={len(value)}" + f"Overwriting prompt '{key}': " f"old length={len(self[key])}, new length={len(value)}", ) self[key] = value else: @@ -170,20 +167,20 @@ class PromptHandler(BaseContext): 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. """ @@ -211,10 +208,10 @@ class PromptHandler(BaseContext): 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. """ @@ -226,11 +223,11 @@ class PromptHandler(BaseContext): 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. """ @@ -243,31 +240,27 @@ class PromptHandler(BaseContext): @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 - } + 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. """ @@ -288,38 +281,38 @@ class PromptHandler(BaseContext): elif flags[matched_flag]: # Flag is True, include without marker marker = f"[{matched_flag}]" - filtered_lines.append(line[len(marker):]) + 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 + 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}" @@ -347,7 +340,7 @@ class PromptHandler(BaseContext): if missing_fields: raise PromptFormattingError( f"Missing required format variables for prompt '{prompt_name}': " - f"{', '.join(sorted(missing_fields))}" + f"{', '.join(sorted(missing_fields))}", ) # Step 3: Format with variables @@ -356,18 +349,15 @@ class PromptHandler(BaseContext): prompt = prompt.format(**format_kwargs) except KeyError as e: raise PromptFormattingError( - f"Format error in prompt '{prompt_name}': missing variable {e}" + 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}" + 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)})" - ) + return f"PromptHandler(language='{self.language}', " f"num_prompts={len(self)})" diff --git a/reme_ai/core/context/registry.py b/reme_ai/core/context/registry.py deleted file mode 100644 index 9fa35d30..00000000 --- a/reme_ai/core/context/registry.py +++ /dev/null @@ -1,143 +0,0 @@ -"""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 ..enumeration import RegistryEnum -from ...core_old.utils import singleton - -T = TypeVar("T") - - -@singleton -class Registry(BaseContext): - """A singleton registry manager that maintains separate registries for different component types. - - This class serves as the central registry hub for the entire ReMe application, providing: - - Component registration for different types (LLMs, embeddings, vector stores, etc.) - - Convenient access methods for retrieving registered classes - - Decorator-based registration API - - The singleton pattern ensures only one instance exists throughout the application lifecycle, - accessible via the global `R` variable exported at the bottom of this module. - """ - - def __init__(self, **kwargs): - """Initialize the registry manager with separate registries for each component type.""" - super().__init__(**kwargs) - - # Registry system: stores class definitions for different component types - self.registry_dict: dict[RegistryEnum, dict] = { - v: {} for v in RegistryEnum.__members__.values() - } - - def register(self, name: str | type = "", register_type: RegistryEnum = None) -> Callable[[type[T]], type[T]] | type[T]: - """Return a decorator to register a component within a specific registry category. - - Can be used in multiple ways: - - @R.register_op() # with empty parentheses, uses class name - - @R.register_op # without parentheses, uses class name - - @R.register_op("custom_name") # with custom name - - Args: - name: Either a string name for the class, or the class itself when used without parentheses - register_type: The type of registry (LLM, EMBEDDING_MODEL, VECTOR_STORE, etc.) - - Returns: - Either a decorator function or the registered class itself - - Example: - @R.register("my_llm", RegistryEnum.LLM) - class MyLLM(BaseLLM): - pass - """ - if inspect.isclass(name): - # Used without parentheses: @R.register_op - self.registry_dict[register_type][name.__name__] = name - return name - else: - # Used with parentheses: @R.register_op() or @R.register_op("name") - def decorator(cls): - key = name if isinstance(name, str) and name else cls.__name__ - self.registry_dict[register_type][key] = cls - return cls - - return decorator - - 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) - - -# Export a global singleton instance for easy access across the application -# This is the primary way to access the registry throughout the codebase -R = Registry() diff --git a/reme_ai/core/context/registry_factory.py b/reme_ai/core/context/registry_factory.py new file mode 100644 index 00000000..28cb1c82 --- /dev/null +++ b/reme_ai/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_ai/core/context/runtime_context.py deleted file mode 100644 index d7112e1c..00000000 --- a/reme_ai/core/context/runtime_context.py +++ /dev/null @@ -1,79 +0,0 @@ -"""Runtime context for managing response states and asynchronous data streaming.""" - -import asyncio - -from .base_context import BaseContext -from ..enumeration import ChunkEnum -from ..schema import Response, StreamChunk - - -class RuntimeContext(BaseContext): - """Context for execution state, response metadata, and stream queues.""" - - def __init__( - self, - response: Response | None = None, - stream_queue: asyncio.Queue | None = None, - **kwargs, - ): - """Initialize the context with optional response and queue.""" - super().__init__(**kwargs) - self.response = response or Response() - self.stream_queue = stream_queue - - @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 - - async def _enqueue(self, chunk: StreamChunk) -> None: - """Internal helper to put a chunk into the queue if it exists.""" - if self.stream_queue: - await self.stream_queue.put(chunk) - - async def add_stream_string(self, chunk: str, chunk_type: ChunkEnum) -> "RuntimeContext": - """Enqueue a stream chunk from a raw string and type.""" - await self._enqueue(StreamChunk(chunk_type=chunk_type, chunk=chunk)) - return self - - async def add_stream_chunk(self, stream_chunk: StreamChunk) -> "RuntimeContext": - """Enqueue an existing stream chunk.""" - await self._enqueue(stream_chunk) - return self - - async def add_stream_done(self) -> "RuntimeContext": - """Enqueue a termination chunk to signal the end of the stream.""" - await self._enqueue(StreamChunk(chunk_type=ChunkEnum.DONE, chunk="", done=True)) - return self - - def add_response_error(self, e: Exception) -> "RuntimeContext": - """Record an exception into the response object.""" - self.response.success = False - self.response.answer = str(e) - return self - - def apply_mapping(self, mapping: dict[str, str]) -> "RuntimeContext": - """Copy internal values based on a source-to-target key map.""" - if not mapping: - return self - - for source, target in mapping.items(): - if source in self: - self[target] = self[source] - return self - - def validate_required_keys(self, required_keys: dict[str, bool], context_name: str = "context") -> "RuntimeContext": - """Ensure all required keys are present in the context. - - Args: - required_keys: Dictionary mapping key names to boolean indicating if required - context_name: Name of the context for error messages (e.g., operator name) - """ - for key, is_required in required_keys.items(): - if is_required and key not in self: - raise ValueError(f"{context_name}: missing required input '{key}'") - return self diff --git a/reme_ai/core/schema/__init__.py b/reme_ai/core/schema/__init__.py new file mode 100644 index 00000000..b7b73719 --- /dev/null +++ b/reme_ai/core/schema/__init__.py @@ -0,0 +1,42 @@ +"""schema""" + +from .memory_node import MemoryNode +from .message import ContentBlock, Message, Trajectory +from .request import Request +from .response import Response +from .service_config import ( + CmdConfig, + EmbeddingModelConfig, + FlowConfig, + HttpConfig, + LLMConfig, + MCPConfig, + ServiceConfig, + TokenCounterConfig, + VectorStoreConfig, +) +from .stream_chunk import StreamChunk +from .tool_call import ToolAttr, ToolCall +from .vector_node import VectorNode + +__all__ = [ + "MemoryNode", + "ContentBlock", + "EmbeddingModelConfig", + "FlowConfig", + "HttpConfig", + "LLMConfig", + "MCPConfig", + "Message", + "Request", + "Response", + "ServiceConfig", + "StreamChunk", + "TokenCounterConfig", + "Trajectory", + "ToolAttr", + "ToolCall", + "VectorNode", + "VectorStoreConfig", + "CmdConfig", +] diff --git a/reme_ai/core/schema/memory_node.py b/reme_ai/core/schema/memory_node.py new file mode 100644 index 00000000..67ed7c43 --- /dev/null +++ b/reme_ai/core/schema/memory_node.py @@ -0,0 +1,198 @@ +"""Memory schema module for the ReMe AI system. + +This module defines the MemoryNode class for storing and retrieving +memories in the ReMe system. +""" + +import datetime +import hashlib +from typing import Any + +from pydantic import BaseModel, Field, model_validator + +from .vector_node import VectorNode +from ..enumeration import MemoryType + + +def get_now_time() -> str: + """Get current timestamp in YYYY-MM-DD HH:MM:SS format. + + Returns: + str: Current timestamp string in format 'YYYY-MM-DD HH:MM:SS'. + """ + return datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S") + + +# Length of the memory ID (first N characters of SHA-256 hash) +MEMORY_ID_LENGTH: int = 16 + + +class MemoryNode(BaseModel): + """Memory node for storing memories in the ReMe system. + + Attributes: + memory_id: Unique identifier, auto-generated from content hash. + memory_type: Type of memory (e.g., SUMMARY, PERSONAL). + memory_target: Target or topic this memory relates to. + when_to_use: Condition description for vector retrieval. + content: Actual memory content. + ref_memory_id: Reference to related raw history memory. + time_created: Creation timestamp. + time_modified: Last modification timestamp. + author: Author or source of this memory. + score: Relevance or importance score. + metadata: Additional metadata for extensibility. + """ + + memory_id: str = Field(default="", description="Unique memory identifier") + memory_type: MemoryType = Field(default=..., description="Type of memory") + memory_target: str = Field(default="", description="Target or topic of the memory") + when_to_use: str = Field(default="", description="Condition description for vector retrieval") + content: str = Field(default="", description="Actual memory content") + ref_memory_id: str = Field(default="", description="Reference to related raw history memory ID") + + time_created: str = Field(default_factory=get_now_time, description="Creation timestamp") + time_modified: str = Field(default_factory=get_now_time, description="Last modification timestamp") + author: str = Field(default="", description="Author or source of the memory") + score: float = Field(default=0, description="Relevance or importance score") + + metadata: dict[str, Any] = Field(default_factory=dict, description="Additional metadata") + + def _update_modified_time(self) -> "MemoryNode": + """Update time_modified to current timestamp. + + Returns: + Self: Returns self for method chaining. + """ + self.time_modified = get_now_time() + return self + + def _update_memory_id(self) -> "MemoryNode": + """Generate memory_id from SHA-256 hash of content. + + Takes the first MEMORY_ID_LENGTH characters of the hash. + + Returns: + Self: Returns self for method chaining. + """ + if not self.content: + return self + + hash_obj = hashlib.sha256(self.content.encode("utf-8")) + hex_dig = hash_obj.hexdigest() + self.memory_id = hex_dig[:MEMORY_ID_LENGTH] + return self + + @model_validator(mode="after") + def _update_after_init(self) -> "MemoryNode": + """Post-initialization validator. + + Auto-generates memory_id from content if not provided. + + Returns: + Self: Returns self for method chaining. + """ + if not self.memory_id: + self._update_memory_id() + return self + + def __setattr__(self, name: str, value): + """Auto-update timestamps and memory_id when content or when_to_use changes. + + Args: + name: Attribute name being set. + value: New value for the attribute. + """ + should_update: bool = name in ("when_to_use", "content") and getattr(self, name, None) != value + super().__setattr__(name, value) + if should_update: + self._update_modified_time() + if name == "content": + self._update_memory_id() + + def to_vector_node(self) -> VectorNode: + """Convert to VectorNode for vector storage. + + When when_to_use is set, use it as vector content and store content in metadata. + When when_to_use is empty, use content as vector content directly. + + Returns: + VectorNode: Vector node representation of this memory. + """ + # Build base metadata (shared fields) + metadata: dict[str, Any] = { + "memory_type": self.memory_type.value, + "memory_target": self.memory_target, + "ref_memory_id": self.ref_memory_id, + "time_created": self.time_created, + "time_modified": self.time_modified, + "author": self.author, + "score": self.score, + **self.metadata, + } + + if self.when_to_use: + # Use when_to_use for vector embedding, store content in metadata + vector_content = self.when_to_use + metadata["content"] = self.content + else: + # Use content directly for vector embedding + vector_content = self.content + + return VectorNode( + vector_id=self.memory_id, + content=vector_content, + metadata=metadata, + ) + + @classmethod + def from_vector_node(cls, node: VectorNode) -> "MemoryNode": + """Reconstruct MemoryNode from VectorNode. + + Reverses the to_vector_node conversion: + - If metadata contains 'content': node.content -> when_to_use, metadata['content'] -> content + - Otherwise: node.content -> content, when_to_use remains empty + + Args: + node: VectorNode containing memory data. + + Returns: + Self: Reconstructed MemoryNode instance. + + Raises: + ValueError: If memory_type in metadata is invalid. + """ + metadata = node.metadata.copy() + memory_type_str = metadata.pop("memory_type", None) + + try: + memory_type: MemoryType = MemoryType(memory_type_str) + except ValueError as e: + raise ValueError( + f"Invalid memory_type '{memory_type_str}' in VectorNode metadata. " + f"Valid types are: {[t.value for t in MemoryType]}", + ) from e + + # Restore when_to_use and content based on metadata structure + if "content" in metadata: + # Original had when_to_use set + when_to_use = node.content + content = metadata.pop("content", "") + else: + # Original had empty when_to_use + when_to_use = "" + content = node.content + + return cls( + memory_id=node.vector_id, + memory_type=memory_type, + memory_target=metadata.pop("memory_target", ""), + when_to_use=when_to_use, + content=content, + ref_memory_id=metadata.pop("ref_memory_id", ""), + time_created=metadata.pop("time_created", ""), + time_modified=metadata.pop("time_modified", ""), + author=metadata.pop("author", ""), + score=metadata.pop("score", 0), + metadata=metadata, + ) diff --git a/reme_ai/core/schema/message.py b/reme_ai/core/schema/message.py new file mode 100644 index 00000000..321dd8c8 --- /dev/null +++ b/reme_ai/core/schema/message.py @@ -0,0 +1,165 @@ +"""Data models for multi-modal conversation history and LLM interaction trajectories.""" + +import datetime +import json +import re + +from pydantic import BaseModel, ConfigDict, Field, model_validator + +from .tool_call import ToolCall +from ..enumeration import Role + + +class ContentBlock(BaseModel): + """ + Individual unit of multi-modal content like text, images, or video. + examples: + { + "type": "image_url", + "image_url": { + "url": "https://img.alicdn.com/imgextra/i1/O1CN01gDEY8M1W114Hi3XcN_!!6000000002727-0-tps-1024-406.jpg" + }, + } + + { + "type": "video", + "video": [ + "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/xzsgiz/football1.jpg", + "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/tdescd/football2.jpg", + "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/zefdja/football3.jpg", + "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/aedbqh/football4.jpg", + ], + } + + { + "type": "text", + "text": "How do you solve this problem?" + } + """ + + model_config = ConfigDict(extra="allow") + + type: str = Field(default="") + content: str | dict | list = Field(default="") + + @model_validator(mode="before") + @classmethod + def init_block(cls, data: dict) -> dict: + """Dynamically maps the type-specific key to the content field.""" + content_type = data.get("type", "") + if content_type and content_type in data: + data["content"] = data[content_type] + return data + + def simple_dump(self) -> dict: + """Serializes the block into an API-compatible dictionary format.""" + return { + "type": self.type, + self.type: self.content, + **self.model_extra, + } + + +class Message(BaseModel): + """Data model for a single dialogue entry including roles and tool interactions.""" + + name: str | None = Field(default=None) + role: Role = Field(default=Role.USER) + content: str | list[ContentBlock] = Field(default="") + reasoning_content: str = Field(default="") + tool_calls: list[ToolCall] = Field(default_factory=list) + tool_call_id: str = Field(default="") + time_created: str = Field(default_factory=lambda: datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")) + metadata: dict = Field(default_factory=dict) + + def dump_content(self) -> str | list[dict]: + """Returns content as a raw string or a list of serialized blocks.""" + if isinstance(self.content, str): + return self.content + return [block.simple_dump() for block in self.content] + + def simple_dump( + self, + add_name: bool = False, + add_reasoning: bool = True, + add_time_created: bool = False, + add_metadata: bool = False, + enable_json_dump: bool = False, + ) -> dict | str: + """Transforms the message into a simplified dictionary for standard APIs.""" + result = {} + if add_name and self.name: + result["name"] = self.name + + result["role"] = self.role.value + result["content"] = self.dump_content() + + if add_reasoning and self.reasoning_content: + result["reasoning_content"] = self.reasoning_content + + if self.tool_calls: + result["tool_calls"] = [tc.simple_output_dump() for tc in self.tool_calls] + + if self.tool_call_id: + result["tool_call_id"] = self.tool_call_id + + if add_time_created: + result["time_created"] = self.time_created + + if add_metadata: + result["metadata"] = self.metadata + + if enable_json_dump: + return json.dumps(result, ensure_ascii=False) + else: + return result + + def format_message( + self, + index: int | None = None, + add_time: bool = False, + use_name: bool = False, + add_reasoning: bool = True, + add_tools: bool = True, + strip_markdown_headers: bool = False, + ) -> str: + """Generates a human-readable string representation of the message.""" + prefix = f"round{index} " if index is not None else "" + time_str = f"[{self.time_created}] " if add_time else "" + header = f"{self.name or self.role.value if use_name else self.role.value}:" + + lines = [f"{prefix}{time_str}{header}"] + + def strip_md_func(line): + if strip_markdown_headers: + line = re.sub(r"\n##+ +", "\n", line) + return line + + if add_reasoning and self.reasoning_content: + lines.append(self.reasoning_content) + + if isinstance(self.content, str): + lines.append(strip_md_func(self.content)) + + 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 = str(text) + lines.append(strip_md_func(text)) + + if add_tools and self.tool_calls: + for tc in self.tool_calls: + lines.append(f" - tool_call={tc.name} params={tc.arguments}") + + return " ".join(lines).strip() + + +class Trajectory(BaseModel): + """Sequence of messages representing a full conversation session and its evaluation.""" + + task_id: str = Field(default="") + messages: list[Message] = Field(default_factory=list) + score: float = Field(default=0.0) + metadata: dict = Field(default_factory=dict) diff --git a/reme_ai/core/schema/request.py b/reme_ai/core/schema/request.py new file mode 100644 index 00000000..ece942b9 --- /dev/null +++ b/reme_ai/core/schema/request.py @@ -0,0 +1,11 @@ +"""Defines the data structure for processing incoming user requests and message history.""" + +from pydantic import Field, BaseModel, ConfigDict + + +class Request(BaseModel): + """Represents a structured request payload containing a query, message list, and metadata.""" + + model_config = ConfigDict(extra="allow") + + metadata: dict = Field(default_factory=dict) diff --git a/reme_ai/core/schema/response.py b/reme_ai/core/schema/response.py new file mode 100644 index 00000000..3104bc6e --- /dev/null +++ b/reme_ai/core/schema/response.py @@ -0,0 +1,11 @@ +"""Defines the standardized data structure for model output responses.""" + +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="") + success: bool = Field(default=True) + metadata: dict = Field(default_factory=dict) diff --git a/reme_ai/core/schema/service_config.py b/reme_ai/core/schema/service_config.py new file mode 100644 index 00000000..e1ae9df7 --- /dev/null +++ b/reme_ai/core/schema/service_config.py @@ -0,0 +1,113 @@ +"""Configuration schemas for service components using Pydantic models.""" + +import os +from typing import Dict, List + +from pydantic import BaseModel, Field, ConfigDict + +from .tool_call import ToolCall + + +class MCPConfig(BaseModel): + """Configuration for Model Context Protocol transport and network settings.""" + + model_config = ConfigDict(extra="allow") + + transport: str = Field(default="stdio") + host: str = Field(default="0.0.0.0") + port: int = Field(default=8001) + + +class HttpConfig(BaseModel): + """Configuration for the HTTP server interface and connection lifecycle.""" + + model_config = ConfigDict(extra="allow") + + host: str = Field(default="0.0.0.0") + port: int = Field(default=8001) + timeout_keep_alive: int = Field(default=3600) + limit_concurrency: int = Field(default=1000) + + +class CmdConfig(BaseModel): + """Configuration for command-line flow execution parameters.""" + + model_config = ConfigDict(extra="allow") + + flow: str = Field(default="") + + +class FlowConfig(ToolCall): + """Configuration for workflow execution, caching, and error handling.""" + + model_config = ConfigDict(extra="allow") + + flow_content: str = Field(default="") + stream: bool = Field(default=False) + raise_exception: bool = Field(default=True) + enable_cache: bool = Field(default=False) + cache_path: str = Field(default="cache/flow") + cache_expire_hours: float = Field(default=0.1) + + +class LLMConfig(BaseModel): + """Configuration for Large Language Model backend and model identification.""" + + model_config = ConfigDict(extra="allow") + + backend: str = Field(default="") + model_name: str = Field(default="") + + +class EmbeddingModelConfig(BaseModel): + """Configuration for embedding model backends and identity.""" + + model_config = ConfigDict(extra="allow") + + backend: str = Field(default="") + model_name: str = Field(default="") + + +class VectorStoreConfig(BaseModel): + """Configuration for vector database storage and associated embeddings.""" + + model_config = ConfigDict(extra="allow") + + backend: str = Field(default="local") + collection_name: str = Field(default="reme") + embedding_model: str = Field(default="default") + + +class TokenCounterConfig(BaseModel): + """Configuration for token counting services and model mapping.""" + + model_config = ConfigDict(extra="allow") + + backend: str = Field(default="base") + model_name: str = Field(default="") + + +class ServiceConfig(BaseModel): + """Root configuration schema aggregating all service-level settings and components.""" + + model_config = ConfigDict(extra="allow") + + backend: str = Field(default="") + app_name: str = Field(default=os.getenv("APP_NAME", "ReMe")) + enable_logo: bool = Field(default=True) + language: str = Field(default="") + 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) + + 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) diff --git a/reme_ai/core/schema/stream_chunk.py b/reme_ai/core/schema/stream_chunk.py new file mode 100644 index 00000000..764981fd --- /dev/null +++ b/reme_ai/core/schema/stream_chunk.py @@ -0,0 +1,14 @@ +"""Defines the data structure for individual data packets in a streaming response.""" + +from pydantic import Field, BaseModel + +from ..enumeration import ChunkEnum + + +class StreamChunk(BaseModel): + """Represents a single chunk of streamed data including its type, content, and completion status.""" + + chunk_type: ChunkEnum = Field(default=ChunkEnum.ANSWER) + chunk: str | dict | list = Field(default="") + done: bool = Field(default=False) + metadata: dict = Field(default_factory=dict) diff --git a/reme_ai/core/schema/tool_call.py b/reme_ai/core/schema/tool_call.py new file mode 100644 index 00000000..70355d60 --- /dev/null +++ b/reme_ai/core/schema/tool_call.py @@ -0,0 +1,226 @@ +"""MCP Tool Schema definitions for recursive JSON Schema representation.""" + +import json +from typing import Any, Dict, List, Optional, Union + +from mcp.types import Tool +from pydantic import BaseModel, ConfigDict, Field, model_validator, field_validator + +from ..enumeration.json_schema_enum import JsonSchemaEnum + + +class ToolAttr(BaseModel): + """Recursive model representing JSON Schema attributes for tool parameters.""" + + model_config = ConfigDict(extra="allow") + + type: str = Field(default=str(JsonSchemaEnum.STRING), description="The data type of the attribute") + description: Optional[str] = Field(default=None, description="Description of the attribute") + required: Optional[List[str]] = Field(default=None, description="Required property names for object types") + properties: Optional[Dict[str, "ToolAttr"]] = Field(default=None, description="Child properties for objects") + items: Optional[Union[Dict[str, Any], "ToolAttr"]] = Field(default=None, description="Schema for array items") + enum: Optional[List[str]] = Field(default=None, description="Allowed values for the attribute") + + @field_validator("type") + @classmethod + def validate_type_is_valid_enum(cls, v: str) -> str: + """Validates that the provided type string exists within JsonSchemaEnum values.""" + valid_types = [str(e) for e in JsonSchemaEnum] + + if v not in valid_types: + raise ValueError(f"Invalid type: '{v}'. Must be one of {valid_types}") + return v + + def simple_input_dump(self) -> dict: + """Serializes the attribute into a standard JSON Schema dictionary.""" + res: dict = {"type": self.type} + if self.description: + res["description"] = self.description + if self.enum: + res["enum"] = self.enum + + 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 is not None: + res["required"] = self.required + + 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 + + +# Enable recursive type resolution +ToolAttr.model_rebuild() + + +class ToolCall(BaseModel): + """ + Model representing a tool definition and its call structure. + Supports parsing from standard JSON Schema formats and converting to MCP Tool objects. + input: + { + "type": "function", + "function": { + "name": "get_current_weather", + "description": "It is very useful when you want to check the weather of a specified city.", + "parameters": { + "type": "object", + "properties": { + "location": { + "type": "string", + "description": "Cities or counties, such as Beijing, Hangzhou, Yuhang District, etc.", + } + }, + "required": ["location"] + } + } + } + output: + { + "index": 0, + "id": "call_6596dafa2a6a46f7a217da", + "function": { + "arguments": "{\"location\": \"Beijing\"}", + "name": "get_current_weather" + }, + "type": "function", + } + """ + + index: int = 0 + id: str = "" + type: str = "function" + name: str = "" + description: str = "" + + arguments: str = Field(default="", description="JSON string of tool execution arguments") + + parameters: ToolAttr = Field( + default_factory=lambda: ToolAttr(type="object", properties={}, required=[]), + 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: + """Initializes the model by parsing tool-specific body data.""" + data = data.copy() + t_type = data.get("type", "function") + body = data.get(t_type, {}) + + # Extract basic metadata + data["name"] = body.get("name", data.get("name", "")) + data["arguments"] = body.get("arguments", data.get("arguments", "")) + data["description"] = body.get("description", data.get("description", "")) + + # Handle parameters mapping + if "parameters" in body: + params = body["parameters"] + # If parameters is already a dict, ensure it matches ToolAttr structure + if isinstance(params, dict): + data["parameters"] = ToolAttr(**params) + + # Handle output mapping (if provided in source) + if "output" in body and isinstance(body["output"], dict): + data["output"] = ToolAttr(**body["output"]) + + return data + + def simple_input_dump(self) -> dict: + """Returns a standardized tool definition dictionary.""" + return { + "type": self.type, + self.type: { + "name": self.name, + "description": self.description, + "parameters": self.parameters.simple_input_dump(), + }, + } + + 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.""" + # MCP Tool inputSchema maps directly to our parameters ToolAttr + return cls( + name=tool.name, + description=tool.description or "", + parameters=ToolAttr(**tool.inputSchema), + ) + + def to_mcp_tool(self) -> Tool: + """Converts the instance back into an MCP Tool object.""" + return Tool( + name=self.name, + description=self.description, + inputSchema=self.parameters.simple_input_dump(), + ) diff --git a/reme_ai/core/schema/vector_node.py b/reme_ai/core/schema/vector_node.py new file mode 100644 index 00000000..937ef4be --- /dev/null +++ b/reme_ai/core/schema/vector_node.py @@ -0,0 +1,15 @@ +"""Defines the data structure for individual vector embedding nodes within a retrieval system.""" + +from typing import List, Dict +from uuid import uuid4 + +from pydantic import BaseModel, Field + + +class VectorNode(BaseModel): + """Represents a discrete unit of text content paired with its corresponding vector embedding and metadata.""" + + vector_id: str = Field(default_factory=lambda: uuid4().hex) + content: str = Field(default="") + vector: List[float] | None = Field(default=None) + metadata: Dict[str, str | bool | int | float] = Field(default_factory=dict) diff --git a/reme_ai/core/utils/__init__.py b/reme_ai/core/utils/__init__.py new file mode 100644 index 00000000..ce9d9b2d --- /dev/null +++ b/reme_ai/core/utils/__init__.py @@ -0,0 +1,7 @@ +"""utils""" + +from .singleton import singleton + +__all__ = [ + "singleton", +] diff --git a/reme_ai/core_old/utils/singleton.py b/reme_ai/core/utils/singleton.py similarity index 100% rename from reme_ai/core_old/utils/singleton.py rename to reme_ai/core/utils/singleton.py From 4560ff09add675c2845ac145e811763db7ef55ef Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 21 Jan 2026 17:10:05 +0800 Subject: [PATCH 11/19] refactor(core): migrate core modules and update imports --- bench/eval_reme_old.py | 4 +- bench/halumem/eval_reme.py | 4 +- bench/halumem/eval_reme_simple.py | 4 +- bench/halumem/eval_reme_simple_v3.py | 4 +- bench/halumem/eval_reme_simple_v4.py | 4 +- bench/halumem/llms.py | 4 +- bench/human_in_the_loop/reevaluate_qa.py | 4 +- bench/human_in_the_loop2/reevaluate_qa.py | 4 +- reme/__init__.py | 0 reme/core/__init__.py | 0 reme/core/context/__init__.py | 11 + .../core}/context/base_context.py | 0 reme/core/context/prompt_handler.py | 363 +++++++++++++++++ .../core/context/registry_factory.py | 0 .../core}/enumeration/__init__.py | 0 .../core}/enumeration/chunk_enum.py | 0 .../core}/enumeration/http_enum.py | 0 reme/core/enumeration/json_schema_enum.py | 38 ++ reme/core/enumeration/memory_type.py | 33 ++ .../core}/enumeration/registry_enum.py | 0 .../core}/enumeration/role.py | 0 .../core_old => reme/core}/schema/__init__.py | 0 .../core}/schema/memory_node.py | 25 -- .../core_old => reme/core}/schema/message.py | 7 +- .../core_old => reme/core}/schema/request.py | 0 .../core_old => reme/core}/schema/response.py | 0 .../core}/schema/service_config.py | 2 +- .../core}/schema/stream_chunk.py | 0 .../core}/schema/tool_call.py | 62 ++- .../core}/schema/vector_node.py | 0 reme/core/utils/__init__.py | 7 + {reme_ai => reme}/core/utils/singleton.py | 0 reme_ai/core/__init__.py | 17 + reme_ai/{core_old => core}/application.py | 0 reme_ai/{core_old => core}/config/__init__.py | 0 .../{core_old => core}/config/default.yaml | 0 .../config/reme_config_parser.py | 0 reme_ai/core/context/__init__.py | 9 +- reme_ai/core/context/prompt_handler.py | 368 +++--------------- .../{core_old => core}/context/registry.py | 0 .../context/runtime_context.py | 0 .../context/service_context.py | 0 .../{core_old => core}/embedding/__init__.py | 0 .../embedding/base_embedding_model.py | 0 .../embedding/openai_embedding_model.py | 0 .../embedding/openai_embedding_model_sync.py | 0 reme_ai/core/enumeration/json_schema_enum.py | 26 +- reme_ai/core/enumeration/memory_type.py | 30 +- reme_ai/{core_old => core}/flow/__init__.py | 0 reme_ai/{core_old => core}/flow/base_flow.py | 0 reme_ai/{core_old => core}/flow/cmd_flow.py | 0 .../flow/expression_flow.py | 0 .../{core_old => core}/flow/simple_flow.py | 0 reme_ai/{core_old => core}/llm/__init__.py | 0 reme_ai/{core_old => core}/llm/base_llm.py | 0 reme_ai/{core_old => core}/llm/lite_llm.py | 0 .../{core_old => core}/llm/lite_llm_sync.py | 0 reme_ai/{core_old => core}/llm/openai_llm.py | 0 .../{core_old => core}/llm/openai_llm_sync.py | 0 reme_ai/{core_old => core}/main.py | 0 reme_ai/{core_old => core}/op/__init__.py | 0 reme_ai/{core_old => core}/op/base_op.py | 0 reme_ai/{core_old => core}/op/base_ray_op.py | 0 reme_ai/{core_old => core}/op/mcp_tool.py | 0 reme_ai/{core_old => core}/op/parallel_op.py | 0 .../{core_old => core}/op/sequential_op.py | 0 reme_ai/{core_old => core}/reme.py | 0 reme_ai/core/schema/memory_node.py | 25 ++ reme_ai/core/schema/message.py | 7 +- reme_ai/core/schema/service_config.py | 2 +- reme_ai/core/schema/tool_call.py | 62 +-- .../{core_old => core}/service/__init__.py | 0 .../service/base_service.py | 0 .../{core_old => core}/service/cmd_service.py | 0 .../service/http_service.py | 0 .../{core_old => core}/service/mcp_service.py | 0 .../token_counter/__init__.py | 0 .../token_counter/base_token_counter.py | 0 .../token_counter/hf_token_counter.py | 0 .../token_counter/openai_token_counter.py | 0 reme_ai/core/utils/__init__.py | 40 ++ .../{core_old => core}/utils/cache_handler.py | 0 .../utils/case_converter.py | 0 .../{core_old => core}/utils/common_utils.py | 0 reme_ai/{core_old => core}/utils/env_utils.py | 0 .../{core_old => core}/utils/execute_tuils.py | 0 .../{core_old => core}/utils/http_client.py | 0 reme_ai/{core_old => core}/utils/llm_utils.py | 0 .../{core_old => core}/utils/logger_utils.py | 0 .../{core_old => core}/utils/logo_utils.py | 0 .../{core_old => core}/utils/mcp_client.py | 0 .../utils/pydantic_config_parser.py | 0 .../utils/pydantic_utils.py | 0 reme_ai/{core_old => core}/utils/time.py | 0 .../vector_store/__init__.py | 0 .../vector_store/base_vector_store.py | 6 +- .../vector_store/chroma_vector_store.py | 0 .../vector_store/es_vector_store.py | 0 .../vector_store/local_vector_store.py | 0 .../vector_store/pgvector_store.py | 0 .../vector_store/qdrant_vector_store.py | 0 reme_ai/core_old/__init__.py | 17 - reme_ai/core_old/context/__init__.py | 16 - reme_ai/core_old/context/prompt_handler.py | 95 ----- .../core_old/enumeration/json_schema_enum.py | 18 - reme_ai/core_old/enumeration/memory_type.py | 25 -- reme_ai/core_old/utils/__init__.py | 47 --- reme_ai/mem_agent/base_memory_agent.py | 6 +- reme_ai/mem_agent/chat/remy_agent.py | 8 +- reme_ai/mem_agent/chat/simple_chat.py | 8 +- reme_ai/mem_agent/chat/stream_chat.py | 8 +- reme_ai/mem_agent/retriever/reme_retriever.py | 8 +- .../retriever_v2/reme_retriever_v2.py | 8 +- .../summarizer/identity_summarizer.py | 8 +- .../summarizer/personal_summarizer.py | 8 +- .../summarizer/procedural_summarizer.py | 8 +- .../mem_agent/summarizer/reme_summarizer.py | 8 +- .../mem_agent/summarizer/tool_summarizer.py | 8 +- .../summarizer_v2/personal_summarizer_v2.py | 8 +- .../summarizer_v2/reme_summarizer_v2.py | 8 +- .../mem_agent/v3/personal_summarizer_v3.py | 6 +- reme_ai/mem_agent/v3/reme_retriever_v3.py | 6 +- reme_ai/mem_agent/v3/reme_summarizer_v3.py | 6 +- reme_ai/mem_agent/v4/personal_retriever_v4.py | 6 +- .../mem_agent/v4/personal_summarizer_v4.py | 4 +- reme_ai/mem_agent/v4/reme_retriever_v4.py | 6 +- reme_ai/mem_agent/v4/reme_summarizer_v4.py | 6 +- .../mem_agent/wk/personal_summarizer_wk.py | 6 +- reme_ai/mem_agent/wk/reme_retriever_wk.py | 6 +- reme_ai/mem_agent/wk/reme_summarizer_wk.py | 6 +- reme_ai/mem_tool/base_memory_tool.py | 8 +- reme_ai/mem_tool/hands_off_tool.py | 4 +- .../mem_tool/history/add_history_memory.py | 8 +- .../mem_tool/history/read_history_memory.py | 4 +- .../mem_tool/identity/read_identity_memory.py | 2 +- .../identity/update_identity_memory.py | 2 +- reme_ai/mem_tool/meta/add_meta_memory.py | 4 +- reme_ai/mem_tool/meta/read_meta_memory.py | 4 +- reme_ai/mem_tool/think_tool.py | 4 +- reme_ai/mem_tool/v2/add_memory_drafts.py | 2 +- reme_ai/mem_tool/v2/read_history.py | 4 +- reme_ai/mem_tool/v2/retrieve_memories.py | 6 +- .../retrieve_recent_and_similar_memories.py | 6 +- reme_ai/mem_tool/v2/summary_and_hands_off.py | 6 +- reme_ai/mem_tool/v2/update_memories.py | 4 +- reme_ai/mem_tool/v3/add_memory.py | 2 +- reme_ai/mem_tool/v3/read_history.py | 2 +- reme_ai/mem_tool/v3/read_user_profile.py | 2 +- reme_ai/mem_tool/v3/retrieve_memory.py | 4 +- reme_ai/mem_tool/v3/summary_and_hands_off.py | 4 +- reme_ai/mem_tool/v3/update_user_profile.py | 2 +- reme_ai/mem_tool/v4/add_summary_memory.py | 2 +- reme_ai/mem_tool/v4/hands_off.py | 4 +- reme_ai/mem_tool/v4/read_history.py | 2 +- reme_ai/mem_tool/v4/read_user_profile.py | 2 +- reme_ai/mem_tool/v4/retrieve_memory.py | 4 +- reme_ai/mem_tool/v4/update_user_profile.py | 4 +- reme_ai/mem_tool/vector_store/add_memory.py | 4 +- .../vector_store/add_summary_memory.py | 6 +- .../mem_tool/vector_store/delete_memory.py | 2 +- .../vector_store/retrieve_recent_memory.py | 6 +- .../mem_tool/vector_store/update_memory.py | 4 +- .../vector_store/vector_retrieve_memory.py | 8 +- reme_ai/mem_tool/wk/add_memory.py | 2 +- reme_ai/mem_tool/wk/read_history.py | 2 +- reme_ai/mem_tool/wk/summary_and_hands_off.py | 4 +- reme_ai/mem_tool/wk/update_memory.py | 2 +- reme_ai/mem_tool/wk/vector_retrieve_memory.py | 6 +- reme_ai/tool/execute/execute_code.py | 8 +- reme_ai/tool/execute/execute_shell.py | 8 +- reme_ai/tool/search/dashscope_search.py | 6 +- reme_ai/tool/search/mock_search.py | 10 +- reme_ai/tool/search/tavily_search.py | 6 +- test/test_base_context.py | 2 +- test/test_cache_handler.py | 2 +- test/test_embedding.py | 6 +- test/test_embedding_sync.py | 6 +- test/test_llm.py | 8 +- test/test_llm_sync.py | 8 +- test/test_logo.py | 4 +- test/test_mcp_client.py | 2 +- test/test_mcp_server.py | 4 +- test/test_message.py | 4 +- test/test_op_composition.py | 4 +- test/test_reme.py | 2 +- test/test_timer.py | 2 +- test/test_token_counter.py | 6 +- test/test_tool.py | 4 +- test/test_tool_call.py | 2 +- test/test_vector_store.py | 10 +- 190 files changed, 906 insertions(+), 906 deletions(-) create mode 100644 reme/__init__.py create mode 100644 reme/core/__init__.py create mode 100644 reme/core/context/__init__.py rename {reme_ai/core_old => reme/core}/context/base_context.py (100%) create mode 100644 reme/core/context/prompt_handler.py rename {reme_ai => reme}/core/context/registry_factory.py (100%) rename {reme_ai/core_old => reme/core}/enumeration/__init__.py (100%) rename {reme_ai/core_old => reme/core}/enumeration/chunk_enum.py (100%) rename {reme_ai/core_old => reme/core}/enumeration/http_enum.py (100%) create mode 100644 reme/core/enumeration/json_schema_enum.py create mode 100644 reme/core/enumeration/memory_type.py rename {reme_ai/core_old => reme/core}/enumeration/registry_enum.py (100%) rename {reme_ai/core_old => reme/core}/enumeration/role.py (100%) rename {reme_ai/core_old => reme/core}/schema/__init__.py (100%) rename {reme_ai/core_old => reme/core}/schema/memory_node.py (91%) rename {reme_ai/core_old => reme/core}/schema/message.py (96%) rename {reme_ai/core_old => reme/core}/schema/request.py (100%) rename {reme_ai/core_old => reme/core}/schema/response.py (100%) rename {reme_ai/core_old => reme/core}/schema/service_config.py (96%) rename {reme_ai/core_old => reme/core}/schema/stream_chunk.py (100%) rename {reme_ai/core_old => reme/core}/schema/tool_call.py (98%) rename {reme_ai/core_old => reme/core}/schema/vector_node.py (100%) create mode 100644 reme/core/utils/__init__.py rename {reme_ai => reme}/core/utils/singleton.py (100%) rename reme_ai/{core_old => core}/application.py (100%) rename reme_ai/{core_old => core}/config/__init__.py (100%) rename reme_ai/{core_old => core}/config/default.yaml (100%) rename reme_ai/{core_old => core}/config/reme_config_parser.py (100%) rename reme_ai/{core_old => core}/context/registry.py (100%) rename reme_ai/{core_old => core}/context/runtime_context.py (100%) rename reme_ai/{core_old => core}/context/service_context.py (100%) rename reme_ai/{core_old => core}/embedding/__init__.py (100%) rename reme_ai/{core_old => core}/embedding/base_embedding_model.py (100%) rename reme_ai/{core_old => core}/embedding/openai_embedding_model.py (100%) rename reme_ai/{core_old => core}/embedding/openai_embedding_model_sync.py (100%) rename reme_ai/{core_old => core}/flow/__init__.py (100%) rename reme_ai/{core_old => core}/flow/base_flow.py (100%) rename reme_ai/{core_old => core}/flow/cmd_flow.py (100%) rename reme_ai/{core_old => core}/flow/expression_flow.py (100%) rename reme_ai/{core_old => core}/flow/simple_flow.py (100%) rename reme_ai/{core_old => core}/llm/__init__.py (100%) rename reme_ai/{core_old => core}/llm/base_llm.py (100%) rename reme_ai/{core_old => core}/llm/lite_llm.py (100%) rename reme_ai/{core_old => core}/llm/lite_llm_sync.py (100%) rename reme_ai/{core_old => core}/llm/openai_llm.py (100%) rename reme_ai/{core_old => core}/llm/openai_llm_sync.py (100%) rename reme_ai/{core_old => core}/main.py (100%) rename reme_ai/{core_old => core}/op/__init__.py (100%) rename reme_ai/{core_old => core}/op/base_op.py (100%) rename reme_ai/{core_old => core}/op/base_ray_op.py (100%) rename reme_ai/{core_old => core}/op/mcp_tool.py (100%) rename reme_ai/{core_old => core}/op/parallel_op.py (100%) rename reme_ai/{core_old => core}/op/sequential_op.py (100%) rename reme_ai/{core_old => core}/reme.py (100%) rename reme_ai/{core_old => core}/service/__init__.py (100%) rename reme_ai/{core_old => core}/service/base_service.py (100%) rename reme_ai/{core_old => core}/service/cmd_service.py (100%) rename reme_ai/{core_old => core}/service/http_service.py (100%) rename reme_ai/{core_old => core}/service/mcp_service.py (100%) rename reme_ai/{core_old => core}/token_counter/__init__.py (100%) rename reme_ai/{core_old => core}/token_counter/base_token_counter.py (100%) rename reme_ai/{core_old => core}/token_counter/hf_token_counter.py (100%) rename reme_ai/{core_old => core}/token_counter/openai_token_counter.py (100%) rename reme_ai/{core_old => core}/utils/cache_handler.py (100%) rename reme_ai/{core_old => core}/utils/case_converter.py (100%) rename reme_ai/{core_old => core}/utils/common_utils.py (100%) rename reme_ai/{core_old => core}/utils/env_utils.py (100%) rename reme_ai/{core_old => core}/utils/execute_tuils.py (100%) rename reme_ai/{core_old => core}/utils/http_client.py (100%) rename reme_ai/{core_old => core}/utils/llm_utils.py (100%) rename reme_ai/{core_old => core}/utils/logger_utils.py (100%) rename reme_ai/{core_old => core}/utils/logo_utils.py (100%) rename reme_ai/{core_old => core}/utils/mcp_client.py (100%) rename reme_ai/{core_old => core}/utils/pydantic_config_parser.py (100%) rename reme_ai/{core_old => core}/utils/pydantic_utils.py (100%) rename reme_ai/{core_old => core}/utils/time.py (100%) rename reme_ai/{core_old => core}/vector_store/__init__.py (100%) rename reme_ai/{core_old => core}/vector_store/base_vector_store.py (97%) rename reme_ai/{core_old => core}/vector_store/chroma_vector_store.py (100%) rename reme_ai/{core_old => core}/vector_store/es_vector_store.py (100%) rename reme_ai/{core_old => core}/vector_store/local_vector_store.py (100%) rename reme_ai/{core_old => core}/vector_store/pgvector_store.py (100%) rename reme_ai/{core_old => core}/vector_store/qdrant_vector_store.py (100%) delete mode 100644 reme_ai/core_old/__init__.py delete mode 100644 reme_ai/core_old/context/__init__.py delete mode 100644 reme_ai/core_old/context/prompt_handler.py delete mode 100644 reme_ai/core_old/enumeration/json_schema_enum.py delete mode 100644 reme_ai/core_old/enumeration/memory_type.py delete mode 100644 reme_ai/core_old/utils/__init__.py diff --git a/bench/eval_reme_old.py b/bench/eval_reme_old.py index 41fde3b1..51265da7 100644 --- a/bench/eval_reme_old.py +++ b/bench/eval_reme_old.py @@ -10,8 +10,8 @@ from datetime import datetime, timezone from tqdm import tqdm -from reme_ai.core_old.enumeration import Role -from reme_ai.core_old.schema import Message, MemoryNode +from reme_ai.core.enumeration import Role +from reme_ai.core.schema import Message, MemoryNode from reme_ai.reme import ReMe TEMPLATE_REME = """Memories for user {user_id}: diff --git a/bench/halumem/eval_reme.py b/bench/halumem/eval_reme.py index 3254c332..2b943143 100644 --- a/bench/halumem/eval_reme.py +++ b/bench/halumem/eval_reme.py @@ -35,8 +35,8 @@ from eval_tools import ( evaluation_for_update_memory, ) from llms import llm_request -from reme_ai.core_old.enumeration import MemoryType -from reme_ai.core_old.schema import MemoryNode +from reme_ai.core.enumeration import MemoryType +from reme_ai.core.schema import MemoryNode from reme_ai.reme import ReMe # Template for formatting memories (from shared YAML config) diff --git a/bench/halumem/eval_reme_simple.py b/bench/halumem/eval_reme_simple.py index 4635c556..17671def 100644 --- a/bench/halumem/eval_reme_simple.py +++ b/bench/halumem/eval_reme_simple.py @@ -26,8 +26,8 @@ from typing import Any from loguru import logger from eval_tools import evaluation_for_question2 -from reme_ai.core_old.enumeration import MemoryType -from reme_ai.core_old.schema import MemoryNode +from reme_ai.core.enumeration import MemoryType +from reme_ai.core.schema import MemoryNode from reme_ai.reme import ReMe diff --git a/bench/halumem/eval_reme_simple_v3.py b/bench/halumem/eval_reme_simple_v3.py index 43cee5d9..469102ba 100644 --- a/bench/halumem/eval_reme_simple_v3.py +++ b/bench/halumem/eval_reme_simple_v3.py @@ -27,8 +27,8 @@ from typing import Any from loguru import logger from eval_tools import evaluation_for_question2 -from reme_ai.core_old.enumeration import MemoryType -from reme_ai.core_old.schema import MemoryNode +from reme_ai.core.enumeration import MemoryType +from reme_ai.core.schema import MemoryNode from reme_ai.reme import ReMe diff --git a/bench/halumem/eval_reme_simple_v4.py b/bench/halumem/eval_reme_simple_v4.py index 63cb4f37..c68ee4d1 100644 --- a/bench/halumem/eval_reme_simple_v4.py +++ b/bench/halumem/eval_reme_simple_v4.py @@ -27,8 +27,8 @@ from typing import Any from loguru import logger from eval_tools import evaluation_for_question2, answer_question_with_memories -from reme_ai.core_old.enumeration import MemoryType -from reme_ai.core_old.schema import MemoryNode +from reme_ai.core.enumeration import MemoryType +from reme_ai.core.schema import MemoryNode from reme_ai.reme import ReMe diff --git a/bench/halumem/llms.py b/bench/halumem/llms.py index 5a59b2b2..73577575 100644 --- a/bench/halumem/llms.py +++ b/bench/halumem/llms.py @@ -5,8 +5,8 @@ import re from tenacity import retry, stop_after_attempt, wait_random_exponential, before_sleep_log -from reme_ai.core_old.schema import Message -from reme_ai.core_old.utils import load_env +from reme_ai.core.schema import Message +from reme_ai.core.utils import load_env from reme_ai.reme import ReMe logger = logging.getLogger(__name__) diff --git a/bench/human_in_the_loop/reevaluate_qa.py b/bench/human_in_the_loop/reevaluate_qa.py index a013616c..6b9cac8f 100644 --- a/bench/human_in_the_loop/reevaluate_qa.py +++ b/bench/human_in_the_loop/reevaluate_qa.py @@ -16,8 +16,8 @@ from collections import defaultdict from pathlib import Path from typing import Any -from reme_ai.core_old.schema import Message -from reme_ai.core_old.utils import load_env +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 diff --git a/bench/human_in_the_loop2/reevaluate_qa.py b/bench/human_in_the_loop2/reevaluate_qa.py index a5e6d0fd..fbc44015 100644 --- a/bench/human_in_the_loop2/reevaluate_qa.py +++ b/bench/human_in_the_loop2/reevaluate_qa.py @@ -16,8 +16,8 @@ from collections import defaultdict from pathlib import Path from typing import Any -from reme_ai.core_old.schema import Message -from reme_ai.core_old.utils import load_env +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 diff --git a/reme/__init__.py b/reme/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/reme/core/__init__.py b/reme/core/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/reme/core/context/__init__.py b/reme/core/context/__init__.py new file mode 100644 index 00000000..27957fd2 --- /dev/null +++ b/reme/core/context/__init__.py @@ -0,0 +1,11 @@ +"""context""" + +from .base_context import BaseContext +from .prompt_handler import PromptHandler +from .registry_factory import R + +__all__ = [ + "BaseContext", + "PromptHandler", + "R", +] diff --git a/reme_ai/core_old/context/base_context.py b/reme/core/context/base_context.py similarity index 100% rename from reme_ai/core_old/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_ai/core/context/registry_factory.py b/reme/core/context/registry_factory.py similarity index 100% rename from reme_ai/core/context/registry_factory.py rename to reme/core/context/registry_factory.py diff --git a/reme_ai/core_old/enumeration/__init__.py b/reme/core/enumeration/__init__.py similarity index 100% rename from reme_ai/core_old/enumeration/__init__.py rename to reme/core/enumeration/__init__.py diff --git a/reme_ai/core_old/enumeration/chunk_enum.py b/reme/core/enumeration/chunk_enum.py similarity index 100% rename from reme_ai/core_old/enumeration/chunk_enum.py rename to reme/core/enumeration/chunk_enum.py diff --git a/reme_ai/core_old/enumeration/http_enum.py b/reme/core/enumeration/http_enum.py similarity index 100% rename from reme_ai/core_old/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_old/enumeration/registry_enum.py b/reme/core/enumeration/registry_enum.py similarity index 100% rename from reme_ai/core_old/enumeration/registry_enum.py rename to reme/core/enumeration/registry_enum.py diff --git a/reme_ai/core_old/enumeration/role.py b/reme/core/enumeration/role.py similarity index 100% rename from reme_ai/core_old/enumeration/role.py rename to reme/core/enumeration/role.py diff --git a/reme_ai/core_old/schema/__init__.py b/reme/core/schema/__init__.py similarity index 100% rename from reme_ai/core_old/schema/__init__.py rename to reme/core/schema/__init__.py diff --git a/reme_ai/core_old/schema/memory_node.py b/reme/core/schema/memory_node.py similarity index 91% rename from reme_ai/core_old/schema/memory_node.py rename to reme/core/schema/memory_node.py index 304dfe5c..67ed7c43 100644 --- a/reme_ai/core_old/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_old/schema/message.py b/reme/core/schema/message.py similarity index 96% rename from reme_ai/core_old/schema/message.py rename to reme/core/schema/message.py index 6c3299e7..321dd8c8 100644 --- a/reme_ai/core_old/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_old/schema/request.py b/reme/core/schema/request.py similarity index 100% rename from reme_ai/core_old/schema/request.py rename to reme/core/schema/request.py diff --git a/reme_ai/core_old/schema/response.py b/reme/core/schema/response.py similarity index 100% rename from reme_ai/core_old/schema/response.py rename to reme/core/schema/response.py diff --git a/reme_ai/core_old/schema/service_config.py b/reme/core/schema/service_config.py similarity index 96% rename from reme_ai/core_old/schema/service_config.py rename to reme/core/schema/service_config.py index 4c6eb543..e1ae9df7 100644 --- a/reme_ai/core_old/schema/service_config.py +++ b/reme/core/schema/service_config.py @@ -101,7 +101,7 @@ class ServiceConfig(BaseModel): 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") + mcp_servers: Dict[str, dict] = Field(default_factory=dict) mcp: MCPConfig = Field(default_factory=MCPConfig) http: HttpConfig = Field(default_factory=HttpConfig) diff --git a/reme_ai/core_old/schema/stream_chunk.py b/reme/core/schema/stream_chunk.py similarity index 100% rename from reme_ai/core_old/schema/stream_chunk.py rename to reme/core/schema/stream_chunk.py diff --git a/reme_ai/core_old/schema/tool_call.py b/reme/core/schema/tool_call.py similarity index 98% rename from reme_ai/core_old/schema/tool_call.py rename to reme/core/schema/tool_call.py index 3c01cdcd..70355d60 100644 --- a/reme_ai/core_old/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 @@ -147,23 +145,17 @@ class ToolCall(BaseModel): }, } - @classmethod - def from_mcp_tool(cls, tool: Tool) -> "ToolCall": - """Creates a ToolCall instance from an MCP Tool object.""" - # MCP Tool inputSchema maps directly to our parameters ToolAttr - return cls( - name=tool.name, - description=tool.description or "", - parameters=ToolAttr(**tool.inputSchema), - ) - - def to_mcp_tool(self) -> Tool: - """Converts the instance back into an MCP Tool object.""" - return Tool( - name=self.name, - description=self.description, - inputSchema=self.parameters.simple_input_dump(), - ) + 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: @@ -208,21 +200,27 @@ class ToolCall(BaseModel): return True except json.JSONDecodeError: # Try removing last character - if sanitized[-1] in ']}': + if sanitized[-1] in "]}": sanitized = sanitized[:-1].rstrip() else: break 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, - } + @classmethod + def from_mcp_tool(cls, tool: Tool) -> "ToolCall": + """Creates a ToolCall instance from an MCP Tool object.""" + # MCP Tool inputSchema maps directly to our parameters ToolAttr + return cls( + name=tool.name, + description=tool.description or "", + parameters=ToolAttr(**tool.inputSchema), + ) + + def to_mcp_tool(self) -> Tool: + """Converts the instance back into an MCP Tool object.""" + return Tool( + name=self.name, + description=self.description, + inputSchema=self.parameters.simple_input_dump(), + ) diff --git a/reme_ai/core_old/schema/vector_node.py b/reme/core/schema/vector_node.py similarity index 100% rename from reme_ai/core_old/schema/vector_node.py rename to reme/core/schema/vector_node.py diff --git a/reme/core/utils/__init__.py b/reme/core/utils/__init__.py new file mode 100644 index 00000000..ce9d9b2d --- /dev/null +++ b/reme/core/utils/__init__.py @@ -0,0 +1,7 @@ +"""utils""" + +from .singleton import singleton + +__all__ = [ + "singleton", +] 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/__init__.py b/reme_ai/core/__init__.py index e69de29b..8eab5792 100644 --- a/reme_ai/core/__init__.py +++ b/reme_ai/core/__init__.py @@ -0,0 +1,17 @@ +"""Core module for ReMe AI framework.""" + +# pylint: disable=wrong-import-position +# flake8: noqa: F401 + +from . import config +from . import context +from . import embedding +from . import enumeration +from . import flow +from . import llm +from . import op +from . import schema +from . import service +from . import token_counter +from . import utils +from . import vector_store diff --git a/reme_ai/core_old/application.py b/reme_ai/core/application.py similarity index 100% rename from reme_ai/core_old/application.py rename to reme_ai/core/application.py diff --git a/reme_ai/core_old/config/__init__.py b/reme_ai/core/config/__init__.py similarity index 100% rename from reme_ai/core_old/config/__init__.py rename to reme_ai/core/config/__init__.py diff --git a/reme_ai/core_old/config/default.yaml b/reme_ai/core/config/default.yaml similarity index 100% rename from reme_ai/core_old/config/default.yaml rename to reme_ai/core/config/default.yaml diff --git a/reme_ai/core_old/config/reme_config_parser.py b/reme_ai/core/config/reme_config_parser.py similarity index 100% rename from reme_ai/core_old/config/reme_config_parser.py rename to reme_ai/core/config/reme_config_parser.py diff --git a/reme_ai/core/context/__init__.py b/reme_ai/core/context/__init__.py index 27957fd2..7f26d600 100644 --- a/reme_ai/core/context/__init__.py +++ b/reme_ai/core/context/__init__.py @@ -2,10 +2,15 @@ from .base_context import BaseContext from .prompt_handler import PromptHandler -from .registry_factory import R +from .registry import Registry +from .runtime_context import RuntimeContext +from .service_context import ServiceContext, C __all__ = [ "BaseContext", "PromptHandler", - "R", + "Registry", + "RuntimeContext", + "ServiceContext", + "C", ] diff --git a/reme_ai/core/context/prompt_handler.py b/reme_ai/core/context/prompt_handler.py index e6b6d737..b428b163 100644 --- a/reme_ai/core/context/prompt_handler.py +++ b/reme_ai/core/context/prompt_handler.py @@ -1,99 +1,24 @@ -"""Module for managing and formatting prompt templates from files or dictionaries. +"""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.""" +from .service_context import C 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!' - """ + """A context-aware handler for loading, retrieving, and formatting prompt templates.""" 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. - """ + """Initialize the handler with a specific language and optional context data.""" super().__init__(**kwargs) - self.language: str = language.strip() + self.language: str = language or C.language - 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. - """ + 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 @@ -101,263 +26,70 @@ class PromptHandler(BaseContext): 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 - + 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: 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. - """ + 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 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 + 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"Skipping existing prompt: key={key}") - else: - logger.debug(f"Adding new prompt: key={key}, length={len(value)}") + logger.debug(f"Adding new prompt key={key}, value={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. + 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() - 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 + assert key in self, f"prompt_name={key} not found." + return self[key].strip() - 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 + 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 format variables + # Separate boolean flags from string formatting arguments 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)} + other_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) + 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 - # Step 2: Validate required fields if requested - if validate: - required_fields = self._extract_format_fields(prompt) - missing_fields = required_fields - set(format_kwargs.keys()) + hit = True + hit_flag = flag + # Remove the flag prefix from the line + line = line.strip(f"[{key}]") + break - if missing_fields: - raise PromptFormattingError( - f"Missing required format variables for prompt '{prompt_name}': " - f"{', '.join(sorted(missing_fields))}", - ) + # 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) - # 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 + prompt = "\n".join(split_prompt) - return prompt.strip() + if other_kwargs: + # Apply standard Python string formatting + prompt = prompt.format(**other_kwargs) - def __repr__(self) -> str: - """Return a string representation of the PromptHandler.""" - return f"PromptHandler(language='{self.language}', " f"num_prompts={len(self)})" + return prompt diff --git a/reme_ai/core_old/context/registry.py b/reme_ai/core/context/registry.py similarity index 100% rename from reme_ai/core_old/context/registry.py rename to reme_ai/core/context/registry.py diff --git a/reme_ai/core_old/context/runtime_context.py b/reme_ai/core/context/runtime_context.py similarity index 100% rename from reme_ai/core_old/context/runtime_context.py rename to reme_ai/core/context/runtime_context.py diff --git a/reme_ai/core_old/context/service_context.py b/reme_ai/core/context/service_context.py similarity index 100% rename from reme_ai/core_old/context/service_context.py rename to reme_ai/core/context/service_context.py diff --git a/reme_ai/core_old/embedding/__init__.py b/reme_ai/core/embedding/__init__.py similarity index 100% rename from reme_ai/core_old/embedding/__init__.py rename to reme_ai/core/embedding/__init__.py diff --git a/reme_ai/core_old/embedding/base_embedding_model.py b/reme_ai/core/embedding/base_embedding_model.py similarity index 100% rename from reme_ai/core_old/embedding/base_embedding_model.py rename to reme_ai/core/embedding/base_embedding_model.py diff --git a/reme_ai/core_old/embedding/openai_embedding_model.py b/reme_ai/core/embedding/openai_embedding_model.py similarity index 100% rename from reme_ai/core_old/embedding/openai_embedding_model.py rename to reme_ai/core/embedding/openai_embedding_model.py diff --git a/reme_ai/core_old/embedding/openai_embedding_model_sync.py b/reme_ai/core/embedding/openai_embedding_model_sync.py similarity index 100% rename from reme_ai/core_old/embedding/openai_embedding_model_sync.py rename to reme_ai/core/embedding/openai_embedding_model_sync.py diff --git a/reme_ai/core/enumeration/json_schema_enum.py b/reme_ai/core/enumeration/json_schema_enum.py index d66882e2..507645f4 100644 --- a/reme_ai/core/enumeration/json_schema_enum.py +++ b/reme_ai/core/enumeration/json_schema_enum.py @@ -1,38 +1,18 @@ -"""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.). -""" +"""Defines the standard data types supported by JSON Schema.""" from enum import Enum class JsonSchemaEnum(Enum): - """Enumeration of valid JSON Schema data types. + """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.""" + """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 index b9f5ed29..22d35481 100644 --- a/reme_ai/core/enumeration/memory_type.py +++ b/reme_ai/core/enumeration/memory_type.py @@ -1,33 +1,25 @@ -"""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.). -""" +"""Memory type enumeration for the three-layer memory architecture.""" from enum import Enum class MemoryType(str, Enum): - """Enumeration of memory categories used by the memory subsystem. + """ + Three-layer memory architecture for agent memory management. - These types describe *what* a piece of memory is about, which guides - storage, retrieval, and summarization strategies. + 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 """ - # 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_old/flow/__init__.py b/reme_ai/core/flow/__init__.py similarity index 100% rename from reme_ai/core_old/flow/__init__.py rename to reme_ai/core/flow/__init__.py diff --git a/reme_ai/core_old/flow/base_flow.py b/reme_ai/core/flow/base_flow.py similarity index 100% rename from reme_ai/core_old/flow/base_flow.py rename to reme_ai/core/flow/base_flow.py diff --git a/reme_ai/core_old/flow/cmd_flow.py b/reme_ai/core/flow/cmd_flow.py similarity index 100% rename from reme_ai/core_old/flow/cmd_flow.py rename to reme_ai/core/flow/cmd_flow.py diff --git a/reme_ai/core_old/flow/expression_flow.py b/reme_ai/core/flow/expression_flow.py similarity index 100% rename from reme_ai/core_old/flow/expression_flow.py rename to reme_ai/core/flow/expression_flow.py diff --git a/reme_ai/core_old/flow/simple_flow.py b/reme_ai/core/flow/simple_flow.py similarity index 100% rename from reme_ai/core_old/flow/simple_flow.py rename to reme_ai/core/flow/simple_flow.py diff --git a/reme_ai/core_old/llm/__init__.py b/reme_ai/core/llm/__init__.py similarity index 100% rename from reme_ai/core_old/llm/__init__.py rename to reme_ai/core/llm/__init__.py diff --git a/reme_ai/core_old/llm/base_llm.py b/reme_ai/core/llm/base_llm.py similarity index 100% rename from reme_ai/core_old/llm/base_llm.py rename to reme_ai/core/llm/base_llm.py diff --git a/reme_ai/core_old/llm/lite_llm.py b/reme_ai/core/llm/lite_llm.py similarity index 100% rename from reme_ai/core_old/llm/lite_llm.py rename to reme_ai/core/llm/lite_llm.py diff --git a/reme_ai/core_old/llm/lite_llm_sync.py b/reme_ai/core/llm/lite_llm_sync.py similarity index 100% rename from reme_ai/core_old/llm/lite_llm_sync.py rename to reme_ai/core/llm/lite_llm_sync.py diff --git a/reme_ai/core_old/llm/openai_llm.py b/reme_ai/core/llm/openai_llm.py similarity index 100% rename from reme_ai/core_old/llm/openai_llm.py rename to reme_ai/core/llm/openai_llm.py diff --git a/reme_ai/core_old/llm/openai_llm_sync.py b/reme_ai/core/llm/openai_llm_sync.py similarity index 100% rename from reme_ai/core_old/llm/openai_llm_sync.py rename to reme_ai/core/llm/openai_llm_sync.py diff --git a/reme_ai/core_old/main.py b/reme_ai/core/main.py similarity index 100% rename from reme_ai/core_old/main.py rename to reme_ai/core/main.py diff --git a/reme_ai/core_old/op/__init__.py b/reme_ai/core/op/__init__.py similarity index 100% rename from reme_ai/core_old/op/__init__.py rename to reme_ai/core/op/__init__.py diff --git a/reme_ai/core_old/op/base_op.py b/reme_ai/core/op/base_op.py similarity index 100% rename from reme_ai/core_old/op/base_op.py rename to reme_ai/core/op/base_op.py diff --git a/reme_ai/core_old/op/base_ray_op.py b/reme_ai/core/op/base_ray_op.py similarity index 100% rename from reme_ai/core_old/op/base_ray_op.py rename to reme_ai/core/op/base_ray_op.py diff --git a/reme_ai/core_old/op/mcp_tool.py b/reme_ai/core/op/mcp_tool.py similarity index 100% rename from reme_ai/core_old/op/mcp_tool.py rename to reme_ai/core/op/mcp_tool.py diff --git a/reme_ai/core_old/op/parallel_op.py b/reme_ai/core/op/parallel_op.py similarity index 100% rename from reme_ai/core_old/op/parallel_op.py rename to reme_ai/core/op/parallel_op.py diff --git a/reme_ai/core_old/op/sequential_op.py b/reme_ai/core/op/sequential_op.py similarity index 100% rename from reme_ai/core_old/op/sequential_op.py rename to reme_ai/core/op/sequential_op.py diff --git a/reme_ai/core_old/reme.py b/reme_ai/core/reme.py similarity index 100% rename from reme_ai/core_old/reme.py rename to reme_ai/core/reme.py diff --git a/reme_ai/core/schema/memory_node.py b/reme_ai/core/schema/memory_node.py index 67ed7c43..304dfe5c 100644 --- a/reme_ai/core/schema/memory_node.py +++ b/reme_ai/core/schema/memory_node.py @@ -6,6 +6,7 @@ memories in the ReMe system. import datetime import hashlib +import json from typing import Any from pydantic import BaseModel, Field, model_validator @@ -145,6 +146,30 @@ 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_ai/core/schema/message.py index 321dd8c8..6c3299e7 100644 --- a/reme_ai/core/schema/message.py +++ b/reme_ai/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,9 +143,8 @@ 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/service_config.py b/reme_ai/core/schema/service_config.py index e1ae9df7..4c6eb543 100644 --- a/reme_ai/core/schema/service_config.py +++ b/reme_ai/core/schema/service_config.py @@ -101,7 +101,7 @@ class ServiceConfig(BaseModel): 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) + mcp_servers: Dict[str, dict] = Field(default_factory=dict, description="External MCP Server configuration") mcp: MCPConfig = Field(default_factory=MCPConfig) http: HttpConfig = Field(default_factory=HttpConfig) diff --git a/reme_ai/core/schema/tool_call.py b/reme_ai/core/schema/tool_call.py index 70355d60..3c01cdcd 100644 --- a/reme_ai/core/schema/tool_call.py +++ b/reme_ai/core/schema/tool_call.py @@ -1,4 +1,6 @@ -"""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 @@ -145,17 +147,23 @@ 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, - } + @classmethod + def from_mcp_tool(cls, tool: Tool) -> "ToolCall": + """Creates a ToolCall instance from an MCP Tool object.""" + # MCP Tool inputSchema maps directly to our parameters ToolAttr + return cls( + name=tool.name, + description=tool.description or "", + parameters=ToolAttr(**tool.inputSchema), + ) + + def to_mcp_tool(self) -> Tool: + """Converts the instance back into an MCP Tool object.""" + return Tool( + name=self.name, + description=self.description, + inputSchema=self.parameters.simple_input_dump(), + ) @property def argument_dict(self) -> dict: @@ -200,27 +208,21 @@ class ToolCall(BaseModel): return True except json.JSONDecodeError: # Try removing last character - if sanitized[-1] in "]}": + 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.""" - # MCP Tool inputSchema maps directly to our parameters ToolAttr - return cls( - name=tool.name, - description=tool.description or "", - parameters=ToolAttr(**tool.inputSchema), - ) - - def to_mcp_tool(self) -> Tool: - """Converts the instance back into an MCP Tool object.""" - return Tool( - name=self.name, - description=self.description, - inputSchema=self.parameters.simple_input_dump(), - ) + 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_old/service/__init__.py b/reme_ai/core/service/__init__.py similarity index 100% rename from reme_ai/core_old/service/__init__.py rename to reme_ai/core/service/__init__.py diff --git a/reme_ai/core_old/service/base_service.py b/reme_ai/core/service/base_service.py similarity index 100% rename from reme_ai/core_old/service/base_service.py rename to reme_ai/core/service/base_service.py diff --git a/reme_ai/core_old/service/cmd_service.py b/reme_ai/core/service/cmd_service.py similarity index 100% rename from reme_ai/core_old/service/cmd_service.py rename to reme_ai/core/service/cmd_service.py diff --git a/reme_ai/core_old/service/http_service.py b/reme_ai/core/service/http_service.py similarity index 100% rename from reme_ai/core_old/service/http_service.py rename to reme_ai/core/service/http_service.py diff --git a/reme_ai/core_old/service/mcp_service.py b/reme_ai/core/service/mcp_service.py similarity index 100% rename from reme_ai/core_old/service/mcp_service.py rename to reme_ai/core/service/mcp_service.py diff --git a/reme_ai/core_old/token_counter/__init__.py b/reme_ai/core/token_counter/__init__.py similarity index 100% rename from reme_ai/core_old/token_counter/__init__.py rename to reme_ai/core/token_counter/__init__.py diff --git a/reme_ai/core_old/token_counter/base_token_counter.py b/reme_ai/core/token_counter/base_token_counter.py similarity index 100% rename from reme_ai/core_old/token_counter/base_token_counter.py rename to reme_ai/core/token_counter/base_token_counter.py diff --git a/reme_ai/core_old/token_counter/hf_token_counter.py b/reme_ai/core/token_counter/hf_token_counter.py similarity index 100% rename from reme_ai/core_old/token_counter/hf_token_counter.py rename to reme_ai/core/token_counter/hf_token_counter.py diff --git a/reme_ai/core_old/token_counter/openai_token_counter.py b/reme_ai/core/token_counter/openai_token_counter.py similarity index 100% rename from reme_ai/core_old/token_counter/openai_token_counter.py rename to reme_ai/core/token_counter/openai_token_counter.py diff --git a/reme_ai/core/utils/__init__.py b/reme_ai/core/utils/__init__.py index ce9d9b2d..23396f97 100644 --- a/reme_ai/core/utils/__init__.py +++ b/reme_ai/core/utils/__init__.py @@ -1,7 +1,47 @@ """utils""" +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 .http_client import HttpClient +from .llm_utils import extract_content, format_messages, deduplicate_memories +from .logger_utils import init_logger +from .logo_utils import print_logo + +# Make MCPClient import optional to avoid breaking if MCP dependencies are not available +try: + from .mcp_client import MCPClient + _HAS_MCP = True +except ImportError: + MCPClient = None + _HAS_MCP = False + +from .pydantic_config_parser import PydanticConfigParser +from .pydantic_utils import create_pydantic_model from .singleton import singleton +from .time import timer, get_now_time __all__ = [ + "CacheHandler", + "snake_to_camel", + "camel_to_snake", + "run_coro_safely", + "execute_stream_task", + "load_env", + "exec_code", + "run_shell_command", + "HttpClient", + "extract_content", + "format_messages", + "deduplicate_memories", + "init_logger", + "print_logo", + "MCPClient", + "PydanticConfigParser", + "create_pydantic_model", "singleton", + "timer", + "get_now_time", ] diff --git a/reme_ai/core_old/utils/cache_handler.py b/reme_ai/core/utils/cache_handler.py similarity index 100% rename from reme_ai/core_old/utils/cache_handler.py rename to reme_ai/core/utils/cache_handler.py diff --git a/reme_ai/core_old/utils/case_converter.py b/reme_ai/core/utils/case_converter.py similarity index 100% rename from reme_ai/core_old/utils/case_converter.py rename to reme_ai/core/utils/case_converter.py diff --git a/reme_ai/core_old/utils/common_utils.py b/reme_ai/core/utils/common_utils.py similarity index 100% rename from reme_ai/core_old/utils/common_utils.py rename to reme_ai/core/utils/common_utils.py diff --git a/reme_ai/core_old/utils/env_utils.py b/reme_ai/core/utils/env_utils.py similarity index 100% rename from reme_ai/core_old/utils/env_utils.py rename to reme_ai/core/utils/env_utils.py diff --git a/reme_ai/core_old/utils/execute_tuils.py b/reme_ai/core/utils/execute_tuils.py similarity index 100% rename from reme_ai/core_old/utils/execute_tuils.py rename to reme_ai/core/utils/execute_tuils.py diff --git a/reme_ai/core_old/utils/http_client.py b/reme_ai/core/utils/http_client.py similarity index 100% rename from reme_ai/core_old/utils/http_client.py rename to reme_ai/core/utils/http_client.py diff --git a/reme_ai/core_old/utils/llm_utils.py b/reme_ai/core/utils/llm_utils.py similarity index 100% rename from reme_ai/core_old/utils/llm_utils.py rename to reme_ai/core/utils/llm_utils.py diff --git a/reme_ai/core_old/utils/logger_utils.py b/reme_ai/core/utils/logger_utils.py similarity index 100% rename from reme_ai/core_old/utils/logger_utils.py rename to reme_ai/core/utils/logger_utils.py diff --git a/reme_ai/core_old/utils/logo_utils.py b/reme_ai/core/utils/logo_utils.py similarity index 100% rename from reme_ai/core_old/utils/logo_utils.py rename to reme_ai/core/utils/logo_utils.py diff --git a/reme_ai/core_old/utils/mcp_client.py b/reme_ai/core/utils/mcp_client.py similarity index 100% rename from reme_ai/core_old/utils/mcp_client.py rename to reme_ai/core/utils/mcp_client.py diff --git a/reme_ai/core_old/utils/pydantic_config_parser.py b/reme_ai/core/utils/pydantic_config_parser.py similarity index 100% rename from reme_ai/core_old/utils/pydantic_config_parser.py rename to reme_ai/core/utils/pydantic_config_parser.py diff --git a/reme_ai/core_old/utils/pydantic_utils.py b/reme_ai/core/utils/pydantic_utils.py similarity index 100% rename from reme_ai/core_old/utils/pydantic_utils.py rename to reme_ai/core/utils/pydantic_utils.py diff --git a/reme_ai/core_old/utils/time.py b/reme_ai/core/utils/time.py similarity index 100% rename from reme_ai/core_old/utils/time.py rename to reme_ai/core/utils/time.py diff --git a/reme_ai/core_old/vector_store/__init__.py b/reme_ai/core/vector_store/__init__.py similarity index 100% rename from reme_ai/core_old/vector_store/__init__.py rename to reme_ai/core/vector_store/__init__.py diff --git a/reme_ai/core_old/vector_store/base_vector_store.py b/reme_ai/core/vector_store/base_vector_store.py similarity index 97% rename from reme_ai/core_old/vector_store/base_vector_store.py rename to reme_ai/core/vector_store/base_vector_store.py index 158af0d9..a4a8ca8e 100644 --- a/reme_ai/core_old/vector_store/base_vector_store.py +++ b/reme_ai/core/vector_store/base_vector_store.py @@ -5,9 +5,9 @@ from abc import ABC, abstractmethod from collections.abc import Callable from functools import partial -from reme_ai.core_old.context import C -from reme_ai.core_old.embedding import BaseEmbeddingModel -from reme_ai.core_old.schema import VectorNode +from reme_ai.core.context import C +from reme_ai.core.embedding import BaseEmbeddingModel +from reme_ai.core.schema import VectorNode class BaseVectorStore(ABC): diff --git a/reme_ai/core_old/vector_store/chroma_vector_store.py b/reme_ai/core/vector_store/chroma_vector_store.py similarity index 100% rename from reme_ai/core_old/vector_store/chroma_vector_store.py rename to reme_ai/core/vector_store/chroma_vector_store.py diff --git a/reme_ai/core_old/vector_store/es_vector_store.py b/reme_ai/core/vector_store/es_vector_store.py similarity index 100% rename from reme_ai/core_old/vector_store/es_vector_store.py rename to reme_ai/core/vector_store/es_vector_store.py diff --git a/reme_ai/core_old/vector_store/local_vector_store.py b/reme_ai/core/vector_store/local_vector_store.py similarity index 100% rename from reme_ai/core_old/vector_store/local_vector_store.py rename to reme_ai/core/vector_store/local_vector_store.py diff --git a/reme_ai/core_old/vector_store/pgvector_store.py b/reme_ai/core/vector_store/pgvector_store.py similarity index 100% rename from reme_ai/core_old/vector_store/pgvector_store.py rename to reme_ai/core/vector_store/pgvector_store.py diff --git a/reme_ai/core_old/vector_store/qdrant_vector_store.py b/reme_ai/core/vector_store/qdrant_vector_store.py similarity index 100% rename from reme_ai/core_old/vector_store/qdrant_vector_store.py rename to reme_ai/core/vector_store/qdrant_vector_store.py diff --git a/reme_ai/core_old/__init__.py b/reme_ai/core_old/__init__.py deleted file mode 100644 index 8eab5792..00000000 --- a/reme_ai/core_old/__init__.py +++ /dev/null @@ -1,17 +0,0 @@ -"""Core module for ReMe AI framework.""" - -# pylint: disable=wrong-import-position -# flake8: noqa: F401 - -from . import config -from . import context -from . import embedding -from . import enumeration -from . import flow -from . import llm -from . import op -from . import schema -from . import service -from . import token_counter -from . import utils -from . import vector_store diff --git a/reme_ai/core_old/context/__init__.py b/reme_ai/core_old/context/__init__.py deleted file mode 100644 index 7f26d600..00000000 --- a/reme_ai/core_old/context/__init__.py +++ /dev/null @@ -1,16 +0,0 @@ -"""context""" - -from .base_context import BaseContext -from .prompt_handler import PromptHandler -from .registry import Registry -from .runtime_context import RuntimeContext -from .service_context import ServiceContext, C - -__all__ = [ - "BaseContext", - "PromptHandler", - "Registry", - "RuntimeContext", - "ServiceContext", - "C", -] diff --git a/reme_ai/core_old/context/prompt_handler.py b/reme_ai/core_old/context/prompt_handler.py deleted file mode 100644 index b428b163..00000000 --- a/reme_ai/core_old/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].strip() - - 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_old/enumeration/json_schema_enum.py b/reme_ai/core_old/enumeration/json_schema_enum.py deleted file mode 100644 index 507645f4..00000000 --- a/reme_ai/core_old/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_old/enumeration/memory_type.py b/reme_ai/core_old/enumeration/memory_type.py deleted file mode 100644 index 22d35481..00000000 --- a/reme_ai/core_old/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_old/utils/__init__.py b/reme_ai/core_old/utils/__init__.py deleted file mode 100644 index 23396f97..00000000 --- a/reme_ai/core_old/utils/__init__.py +++ /dev/null @@ -1,47 +0,0 @@ -"""utils""" - -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 .http_client import HttpClient -from .llm_utils import extract_content, format_messages, deduplicate_memories -from .logger_utils import init_logger -from .logo_utils import print_logo - -# Make MCPClient import optional to avoid breaking if MCP dependencies are not available -try: - from .mcp_client import MCPClient - _HAS_MCP = True -except ImportError: - MCPClient = None - _HAS_MCP = False - -from .pydantic_config_parser import PydanticConfigParser -from .pydantic_utils import create_pydantic_model -from .singleton import singleton -from .time import timer, get_now_time - -__all__ = [ - "CacheHandler", - "snake_to_camel", - "camel_to_snake", - "run_coro_safely", - "execute_stream_task", - "load_env", - "exec_code", - "run_shell_command", - "HttpClient", - "extract_content", - "format_messages", - "deduplicate_memories", - "init_logger", - "print_logo", - "MCPClient", - "PydanticConfigParser", - "create_pydantic_model", - "singleton", - "timer", - "get_now_time", -] diff --git a/reme_ai/mem_agent/base_memory_agent.py b/reme_ai/mem_agent/base_memory_agent.py index 5ba0cad3..7c21bb19 100644 --- a/reme_ai/mem_agent/base_memory_agent.py +++ b/reme_ai/mem_agent/base_memory_agent.py @@ -6,9 +6,9 @@ from abc import ABCMeta from loguru import logger -from ..core_old.enumeration import Role, MemoryType -from ..core_old.op import BaseOp -from ..core_old.schema import Message, ToolCall, MemoryNode +from ..core.enumeration import Role, MemoryType +from ..core.op import BaseOp +from ..core.schema import Message, ToolCall, MemoryNode from ..mem_tool import BaseMemoryTool, ThinkTool diff --git a/reme_ai/mem_agent/chat/remy_agent.py b/reme_ai/mem_agent/chat/remy_agent.py index c7806617..1eaa9ac7 100644 --- a/reme_ai/mem_agent/chat/remy_agent.py +++ b/reme_ai/mem_agent/chat/remy_agent.py @@ -3,10 +3,10 @@ from typing import List from ..base_memory_agent import BaseMemoryAgent -from ...core_old.context import C -from ...core_old.enumeration import Role -from ...core_old.schema import Message -from ...core_old.utils import get_now_time +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() diff --git a/reme_ai/mem_agent/chat/simple_chat.py b/reme_ai/mem_agent/chat/simple_chat.py index 9e09ffec..8a71c7a8 100644 --- a/reme_ai/mem_agent/chat/simple_chat.py +++ b/reme_ai/mem_agent/chat/simple_chat.py @@ -2,10 +2,10 @@ from loguru import logger -from ...core_old.context import C -from ...core_old.enumeration import Role -from ...core_old.op import BaseOp -from ...core_old.schema import Message, ToolCall +from ...core.context import C +from ...core.enumeration import Role +from ...core.op import BaseOp +from ...core.schema import Message, ToolCall @C.register_op() diff --git a/reme_ai/mem_agent/chat/stream_chat.py b/reme_ai/mem_agent/chat/stream_chat.py index 2b121446..470e4647 100644 --- a/reme_ai/mem_agent/chat/stream_chat.py +++ b/reme_ai/mem_agent/chat/stream_chat.py @@ -2,10 +2,10 @@ from loguru import logger -from ...core_old.context import C -from ...core_old.enumeration import Role, ChunkEnum -from ...core_old.op import BaseOp -from ...core_old.schema import Message, ToolCall +from ...core.context import C +from ...core.enumeration import Role, ChunkEnum +from ...core.op import BaseOp +from ...core.schema import Message, ToolCall @C.register_op() diff --git a/reme_ai/mem_agent/retriever/reme_retriever.py b/reme_ai/mem_agent/retriever/reme_retriever.py index 400e0d01..f3700b1e 100644 --- a/reme_ai/mem_agent/retriever/reme_retriever.py +++ b/reme_ai/mem_agent/retriever/reme_retriever.py @@ -3,10 +3,10 @@ from typing import List from ..base_memory_agent import BaseMemoryAgent -from ...core_old.context import C -from ...core_old.enumeration import Role -from ...core_old.schema import Message -from ...core_old.utils import get_now_time, format_messages +from ...core.context import C +from ...core.enumeration import Role +from ...core.schema import Message +from ...core.utils import get_now_time, format_messages @C.register_op() 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 fedfe188..3b934172 100644 --- a/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.py +++ b/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.py @@ -3,10 +3,10 @@ from typing import List from ..base_memory_agent import BaseMemoryAgent -from ...core_old.context import C -from ...core_old.enumeration import Role -from ...core_old.schema import Message -from ...core_old.utils import format_messages +from ...core.context import C +from ...core.enumeration import Role +from ...core.schema import Message +from ...core.utils import format_messages @C.register_op() diff --git a/reme_ai/mem_agent/summarizer/identity_summarizer.py b/reme_ai/mem_agent/summarizer/identity_summarizer.py index 0a9e1410..be571cdc 100644 --- a/reme_ai/mem_agent/summarizer/identity_summarizer.py +++ b/reme_ai/mem_agent/summarizer/identity_summarizer.py @@ -1,10 +1,10 @@ """Specialized agent for extracting and updating agent self-cognition memories.""" from ..base_memory_agent import BaseMemoryAgent -from ...core_old.context import C -from ...core_old.enumeration import Role, MemoryType -from ...core_old.schema import Message -from ...core_old.utils import get_now_time, format_messages +from ...core.context import C +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message +from ...core.utils import get_now_time, format_messages @C.register_op() diff --git a/reme_ai/mem_agent/summarizer/personal_summarizer.py b/reme_ai/mem_agent/summarizer/personal_summarizer.py index 1c32f499..352fe1ee 100644 --- a/reme_ai/mem_agent/summarizer/personal_summarizer.py +++ b/reme_ai/mem_agent/summarizer/personal_summarizer.py @@ -1,10 +1,10 @@ """Specialized agent for extracting and managing personal memories about specific individuals.""" from ..base_memory_agent import BaseMemoryAgent -from ...core_old.context import C -from ...core_old.enumeration import Role, MemoryType -from ...core_old.schema import Message, ToolCall -from ...core_old.utils import get_now_time, format_messages +from ...core.context import C +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message, ToolCall +from ...core.utils import get_now_time, format_messages @C.register_op() diff --git a/reme_ai/mem_agent/summarizer/procedural_summarizer.py b/reme_ai/mem_agent/summarizer/procedural_summarizer.py index 24a75339..e31422e4 100644 --- a/reme_ai/mem_agent/summarizer/procedural_summarizer.py +++ b/reme_ai/mem_agent/summarizer/procedural_summarizer.py @@ -1,10 +1,10 @@ """Specialized agent for extracting and managing procedural knowledge and workflows.""" from ..base_memory_agent import BaseMemoryAgent -from ...core_old.context import C -from ...core_old.enumeration import Role, MemoryType -from ...core_old.schema import Message -from ...core_old.utils import get_now_time, format_messages +from ...core.context import C +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message +from ...core.utils import get_now_time, format_messages @C.register_op() diff --git a/reme_ai/mem_agent/summarizer/reme_summarizer.py b/reme_ai/mem_agent/summarizer/reme_summarizer.py index 5521762c..b2418db5 100644 --- a/reme_ai/mem_agent/summarizer/reme_summarizer.py +++ b/reme_ai/mem_agent/summarizer/reme_summarizer.py @@ -3,10 +3,10 @@ from loguru import logger from ..base_memory_agent import BaseMemoryAgent -from ...core_old.context import C -from ...core_old.enumeration import Role, MemoryType -from ...core_old.schema import Message, MemoryNode, ToolCall -from ...core_old.utils import get_now_time, format_messages +from ...core.context import C +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message, MemoryNode, ToolCall +from ...core.utils import get_now_time, format_messages @C.register_op() diff --git a/reme_ai/mem_agent/summarizer/tool_summarizer.py b/reme_ai/mem_agent/summarizer/tool_summarizer.py index 1e9e33c0..50399bba 100644 --- a/reme_ai/mem_agent/summarizer/tool_summarizer.py +++ b/reme_ai/mem_agent/summarizer/tool_summarizer.py @@ -1,10 +1,10 @@ """Specialized agent for extracting and managing tool usage guidelines and best practices.""" from ..base_memory_agent import BaseMemoryAgent -from ...core_old.context import C -from ...core_old.enumeration import Role, MemoryType -from ...core_old.schema import Message -from ...core_old.utils import get_now_time, format_messages +from ...core.context import C +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message +from ...core.utils import get_now_time, format_messages @C.register_op() 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 cc2794b3..17bf4c66 100644 --- a/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.py +++ b/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.py @@ -1,10 +1,10 @@ """Simplified personal memory summarizer using v2 memory tools.""" from ..base_memory_agent import BaseMemoryAgent -from ...core_old.context import C -from ...core_old.enumeration import Role, MemoryType -from ...core_old.schema import Message, ToolCall -from ...core_old.utils import format_messages +from ...core.context import C +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message, ToolCall +from ...core.utils import format_messages @C.register_op() diff --git a/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.py b/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.py index aa680da9..1aae4ad4 100644 --- a/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.py +++ b/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.py @@ -3,10 +3,10 @@ from loguru import logger from ..base_memory_agent import BaseMemoryAgent -from ...core_old.context import C -from ...core_old.enumeration import Role, MemoryType -from ...core_old.schema import Message, MemoryNode, ToolCall -from ...core_old.utils import format_messages +from ...core.context import C +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message, MemoryNode, ToolCall +from ...core.utils import format_messages @C.register_op() diff --git a/reme_ai/mem_agent/v3/personal_summarizer_v3.py b/reme_ai/mem_agent/v3/personal_summarizer_v3.py index 3f2c9ca9..0093884d 100644 --- a/reme_ai/mem_agent/v3/personal_summarizer_v3.py +++ b/reme_ai/mem_agent/v3/personal_summarizer_v3.py @@ -1,7 +1,7 @@ from ..base_memory_agent import BaseMemoryAgent -from ...core_old.enumeration import Role, MemoryType -from ...core_old.schema import Message, ToolCall -from ...core_old.utils import format_messages +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message, ToolCall +from ...core.utils import format_messages class PersonalSummarizerV3(BaseMemoryAgent): diff --git a/reme_ai/mem_agent/v3/reme_retriever_v3.py b/reme_ai/mem_agent/v3/reme_retriever_v3.py index 020b5ad2..8f5c62dc 100644 --- a/reme_ai/mem_agent/v3/reme_retriever_v3.py +++ b/reme_ai/mem_agent/v3/reme_retriever_v3.py @@ -3,9 +3,9 @@ from typing import List from ..base_memory_agent import BaseMemoryAgent -from ...core_old.enumeration import Role -from ...core_old.schema import Message -from ...core_old.utils import format_messages +from ...core.enumeration import Role +from ...core.schema import Message +from ...core.utils import format_messages class ReMeRetrieverV3(BaseMemoryAgent): diff --git a/reme_ai/mem_agent/v3/reme_summarizer_v3.py b/reme_ai/mem_agent/v3/reme_summarizer_v3.py index 3e1f17f7..a0f466b9 100644 --- a/reme_ai/mem_agent/v3/reme_summarizer_v3.py +++ b/reme_ai/mem_agent/v3/reme_summarizer_v3.py @@ -1,9 +1,9 @@ from loguru import logger from ..base_memory_agent import BaseMemoryAgent -from ...core_old.enumeration import Role, MemoryType -from ...core_old.schema import Message, MemoryNode, ToolCall -from ...core_old.utils import format_messages +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message, MemoryNode, ToolCall +from ...core.utils import format_messages class ReMeSummarizerV3(BaseMemoryAgent): diff --git a/reme_ai/mem_agent/v4/personal_retriever_v4.py b/reme_ai/mem_agent/v4/personal_retriever_v4.py index 2ba0dba5..dedcf99a 100644 --- a/reme_ai/mem_agent/v4/personal_retriever_v4.py +++ b/reme_ai/mem_agent/v4/personal_retriever_v4.py @@ -1,7 +1,7 @@ from ..base_memory_agent import BaseMemoryAgent -from ...core_old.enumeration import Role, MemoryType -from ...core_old.schema import Message -from ...core_old.utils import format_messages +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message +from ...core.utils import format_messages from ...mem_tool.v4 import ReadUserProfile diff --git a/reme_ai/mem_agent/v4/personal_summarizer_v4.py b/reme_ai/mem_agent/v4/personal_summarizer_v4.py index 8cda9b24..c0e1d4f1 100644 --- a/reme_ai/mem_agent/v4/personal_summarizer_v4.py +++ b/reme_ai/mem_agent/v4/personal_summarizer_v4.py @@ -1,8 +1,8 @@ from loguru import logger from ..base_memory_agent import BaseMemoryAgent -from ...core_old.enumeration import Role, MemoryType -from ...core_old.schema import Message, MemoryNode +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message, MemoryNode class PersonalSummarizerV4(BaseMemoryAgent): diff --git a/reme_ai/mem_agent/v4/reme_retriever_v4.py b/reme_ai/mem_agent/v4/reme_retriever_v4.py index db5f92ae..48ab9f38 100644 --- a/reme_ai/mem_agent/v4/reme_retriever_v4.py +++ b/reme_ai/mem_agent/v4/reme_retriever_v4.py @@ -1,9 +1,9 @@ from loguru import logger from ..base_memory_agent import BaseMemoryAgent -from ...core_old.enumeration import Role -from ...core_old.schema import Message -from ...core_old.utils import format_messages +from ...core.enumeration import Role +from ...core.schema import Message +from ...core.utils import format_messages class ReMeRetrieverV4(BaseMemoryAgent): diff --git a/reme_ai/mem_agent/v4/reme_summarizer_v4.py b/reme_ai/mem_agent/v4/reme_summarizer_v4.py index 7787d035..a4069c85 100644 --- a/reme_ai/mem_agent/v4/reme_summarizer_v4.py +++ b/reme_ai/mem_agent/v4/reme_summarizer_v4.py @@ -1,9 +1,9 @@ from loguru import logger from ..base_memory_agent import BaseMemoryAgent -from ...core_old.enumeration import Role, MemoryType -from ...core_old.schema import Message, MemoryNode -from ...core_old.utils import format_messages +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message, MemoryNode +from ...core.utils import format_messages class ReMeSummarizerV4(BaseMemoryAgent): diff --git a/reme_ai/mem_agent/wk/personal_summarizer_wk.py b/reme_ai/mem_agent/wk/personal_summarizer_wk.py index c95ac36b..e974f99a 100644 --- a/reme_ai/mem_agent/wk/personal_summarizer_wk.py +++ b/reme_ai/mem_agent/wk/personal_summarizer_wk.py @@ -1,7 +1,7 @@ from ..base_memory_agent import BaseMemoryAgent -from ...core_old.enumeration import Role, MemoryType -from ...core_old.schema import Message, ToolCall -from ...core_old.utils import format_messages +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message, ToolCall +from ...core.utils import format_messages class PersonalSummarizerWk(BaseMemoryAgent): diff --git a/reme_ai/mem_agent/wk/reme_retriever_wk.py b/reme_ai/mem_agent/wk/reme_retriever_wk.py index 21a5403a..c98de937 100644 --- a/reme_ai/mem_agent/wk/reme_retriever_wk.py +++ b/reme_ai/mem_agent/wk/reme_retriever_wk.py @@ -3,9 +3,9 @@ from typing import List from ..base_memory_agent import BaseMemoryAgent -from ...core_old.enumeration import Role -from ...core_old.schema import Message -from ...core_old.utils import format_messages +from ...core.enumeration import Role +from ...core.schema import Message +from ...core.utils import format_messages class ReMeRetrieverV2(BaseMemoryAgent): diff --git a/reme_ai/mem_agent/wk/reme_summarizer_wk.py b/reme_ai/mem_agent/wk/reme_summarizer_wk.py index a04d230d..02a7dbf3 100644 --- a/reme_ai/mem_agent/wk/reme_summarizer_wk.py +++ b/reme_ai/mem_agent/wk/reme_summarizer_wk.py @@ -1,9 +1,9 @@ from loguru import logger from ..base_memory_agent import BaseMemoryAgent -from ...core_old.enumeration import Role, MemoryType -from ...core_old.schema import Message, MemoryNode, ToolCall -from ...core_old.utils import format_messages +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message, MemoryNode, ToolCall +from ...core.utils import format_messages class ReMeSummarizerWk(BaseMemoryAgent): diff --git a/reme_ai/mem_tool/base_memory_tool.py b/reme_ai/mem_tool/base_memory_tool.py index 827066b9..8b124496 100644 --- a/reme_ai/mem_tool/base_memory_tool.py +++ b/reme_ai/mem_tool/base_memory_tool.py @@ -3,10 +3,10 @@ from abc import ABCMeta from pathlib import Path -from ..core_old.enumeration import MemoryType -from ..core_old.op import BaseOp -from ..core_old.schema import ToolCall, MemoryNode -from ..core_old.utils import CacheHandler +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): diff --git a/reme_ai/mem_tool/hands_off_tool.py b/reme_ai/mem_tool/hands_off_tool.py index 7ab65cd7..2cb28d60 100644 --- a/reme_ai/mem_tool/hands_off_tool.py +++ b/reme_ai/mem_tool/hands_off_tool.py @@ -6,8 +6,8 @@ from typing import TYPE_CHECKING from loguru import logger from .base_memory_tool import BaseMemoryTool -from ..core_old.context import C -from ..core_old.enumeration import MemoryType +from ..core.context import C +from ..core.enumeration import MemoryType if TYPE_CHECKING: from ..mem_agent import BaseMemoryAgent diff --git a/reme_ai/mem_tool/history/add_history_memory.py b/reme_ai/mem_tool/history/add_history_memory.py index 85e0181a..a92deca5 100644 --- a/reme_ai/mem_tool/history/add_history_memory.py +++ b/reme_ai/mem_tool/history/add_history_memory.py @@ -3,10 +3,10 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C -from ...core_old.enumeration import MemoryType -from ...core_old.schema import ToolCall, Message -from ...core_old.utils import format_messages +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() diff --git a/reme_ai/mem_tool/history/read_history_memory.py b/reme_ai/mem_tool/history/read_history_memory.py index 24cc0366..def2ff24 100644 --- a/reme_ai/mem_tool/history/read_history_memory.py +++ b/reme_ai/mem_tool/history/read_history_memory.py @@ -3,8 +3,8 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C -from ...core_old.schema import MemoryNode +from ...core.context import C +from ...core.schema import MemoryNode @C.register_op() diff --git a/reme_ai/mem_tool/identity/read_identity_memory.py b/reme_ai/mem_tool/identity/read_identity_memory.py index dfede68f..bd9f8031 100644 --- a/reme_ai/mem_tool/identity/read_identity_memory.py +++ b/reme_ai/mem_tool/identity/read_identity_memory.py @@ -3,7 +3,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C +from ...core.context import C @C.register_op() diff --git a/reme_ai/mem_tool/identity/update_identity_memory.py b/reme_ai/mem_tool/identity/update_identity_memory.py index 0883b1a9..b0211242 100644 --- a/reme_ai/mem_tool/identity/update_identity_memory.py +++ b/reme_ai/mem_tool/identity/update_identity_memory.py @@ -3,7 +3,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C +from ...core.context import C @C.register_op() diff --git a/reme_ai/mem_tool/meta/add_meta_memory.py b/reme_ai/mem_tool/meta/add_meta_memory.py index 99635698..d7b41254 100644 --- a/reme_ai/mem_tool/meta/add_meta_memory.py +++ b/reme_ai/mem_tool/meta/add_meta_memory.py @@ -5,8 +5,8 @@ import json from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C -from ...core_old.enumeration import MemoryType +from ...core.context import C +from ...core.enumeration import MemoryType @C.register_op() diff --git a/reme_ai/mem_tool/meta/read_meta_memory.py b/reme_ai/mem_tool/meta/read_meta_memory.py index 5ee59d54..07ad1ecf 100644 --- a/reme_ai/mem_tool/meta/read_meta_memory.py +++ b/reme_ai/mem_tool/meta/read_meta_memory.py @@ -3,8 +3,8 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C -from ...core_old.enumeration import MemoryType +from ...core.context import C +from ...core.enumeration import MemoryType @C.register_op() diff --git a/reme_ai/mem_tool/think_tool.py b/reme_ai/mem_tool/think_tool.py index c54c7676..1d26446a 100644 --- a/reme_ai/mem_tool/think_tool.py +++ b/reme_ai/mem_tool/think_tool.py @@ -5,8 +5,8 @@ before taking actions, helping agents reason about their next steps. """ from .base_memory_tool import BaseMemoryTool -from ..core_old.context import C -from ..core_old.schema import ToolCall +from ..core.context import C +from ..core.schema import ToolCall @C.register_op() diff --git a/reme_ai/mem_tool/v2/add_memory_drafts.py b/reme_ai/mem_tool/v2/add_memory_drafts.py index b815be6e..93caad0c 100644 --- a/reme_ai/mem_tool/v2/add_memory_drafts.py +++ b/reme_ai/mem_tool/v2/add_memory_drafts.py @@ -3,7 +3,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C +from ...core.context import C @C.register_op() diff --git a/reme_ai/mem_tool/v2/read_history.py b/reme_ai/mem_tool/v2/read_history.py index 141989c0..7bb3b830 100644 --- a/reme_ai/mem_tool/v2/read_history.py +++ b/reme_ai/mem_tool/v2/read_history.py @@ -3,8 +3,8 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C -from ...core_old.schema import MemoryNode +from ...core.context import C +from ...core.schema import MemoryNode @C.register_op() diff --git a/reme_ai/mem_tool/v2/retrieve_memories.py b/reme_ai/mem_tool/v2/retrieve_memories.py index 96d4ca99..abdfc377 100644 --- a/reme_ai/mem_tool/v2/retrieve_memories.py +++ b/reme_ai/mem_tool/v2/retrieve_memories.py @@ -3,9 +3,9 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C -from ...core_old.schema import MemoryNode, VectorNode -from ...core_old.utils import deduplicate_memories +from ...core.context import C +from ...core.schema import MemoryNode, VectorNode +from ...core.utils import deduplicate_memories @C.register_op() diff --git a/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.py b/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.py index cfed10c8..38107ab0 100644 --- a/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.py +++ b/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.py @@ -3,9 +3,9 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C -from ...core_old.schema import MemoryNode, VectorNode -from ...core_old.utils import deduplicate_memories +from ...core.context import C +from ...core.schema import MemoryNode, VectorNode +from ...core.utils import deduplicate_memories @C.register_op() diff --git a/reme_ai/mem_tool/v2/summary_and_hands_off.py b/reme_ai/mem_tool/v2/summary_and_hands_off.py index 631c7d6b..131b1101 100644 --- a/reme_ai/mem_tool/v2/summary_and_hands_off.py +++ b/reme_ai/mem_tool/v2/summary_and_hands_off.py @@ -6,9 +6,9 @@ from typing import TYPE_CHECKING from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C -from ...core_old.enumeration import MemoryType -from ...core_old.schema import MemoryNode, Message +from ...core.context import C +from ...core.enumeration import MemoryType +from ...core.schema import MemoryNode, Message if TYPE_CHECKING: from ...mem_agent import BaseMemoryAgent diff --git a/reme_ai/mem_tool/v2/update_memories.py b/reme_ai/mem_tool/v2/update_memories.py index cffe74ac..03a9f494 100644 --- a/reme_ai/mem_tool/v2/update_memories.py +++ b/reme_ai/mem_tool/v2/update_memories.py @@ -3,8 +3,8 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C -from ...core_old.schema import MemoryNode +from ...core.context import C +from ...core.schema import MemoryNode @C.register_op() diff --git a/reme_ai/mem_tool/v3/add_memory.py b/reme_ai/mem_tool/v3/add_memory.py index 4893c6a2..ee488639 100644 --- a/reme_ai/mem_tool/v3/add_memory.py +++ b/reme_ai/mem_tool/v3/add_memory.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.schema import MemoryNode +from ...core.schema import MemoryNode class AddMemory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/v3/read_history.py b/reme_ai/mem_tool/v3/read_history.py index 9506d88a..e9ab2a15 100644 --- a/reme_ai/mem_tool/v3/read_history.py +++ b/reme_ai/mem_tool/v3/read_history.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.schema import MemoryNode +from ...core.schema import MemoryNode class ReadHistory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/v3/read_user_profile.py b/reme_ai/mem_tool/v3/read_user_profile.py index bab1d413..3dba2bcf 100644 --- a/reme_ai/mem_tool/v3/read_user_profile.py +++ b/reme_ai/mem_tool/v3/read_user_profile.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.schema.memory_node import MemoryNode +from ...core.schema.memory_node import MemoryNode class ReadUserProfile(BaseMemoryTool): diff --git a/reme_ai/mem_tool/v3/retrieve_memory.py b/reme_ai/mem_tool/v3/retrieve_memory.py index d5b9a7bc..32e526d7 100644 --- a/reme_ai/mem_tool/v3/retrieve_memory.py +++ b/reme_ai/mem_tool/v3/retrieve_memory.py @@ -3,8 +3,8 @@ import json from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.schema import MemoryNode -from ...core_old.utils import deduplicate_memories +from ...core.schema import MemoryNode +from ...core.utils import deduplicate_memories class RetrieveMemory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/v3/summary_and_hands_off.py b/reme_ai/mem_tool/v3/summary_and_hands_off.py index e0b2756c..19d88744 100644 --- a/reme_ai/mem_tool/v3/summary_and_hands_off.py +++ b/reme_ai/mem_tool/v3/summary_and_hands_off.py @@ -4,8 +4,8 @@ from typing import TYPE_CHECKING from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.enumeration import MemoryType -from ...core_old.schema import MemoryNode, Message +from ...core.enumeration import MemoryType +from ...core.schema import MemoryNode, Message if TYPE_CHECKING: from ...mem_agent import BaseMemoryAgent diff --git a/reme_ai/mem_tool/v3/update_user_profile.py b/reme_ai/mem_tool/v3/update_user_profile.py index ef47f46b..46879e2b 100644 --- a/reme_ai/mem_tool/v3/update_user_profile.py +++ b/reme_ai/mem_tool/v3/update_user_profile.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.schema.memory_node import MemoryNode +from ...core.schema.memory_node import MemoryNode class UpdateUserProfile(BaseMemoryTool): diff --git a/reme_ai/mem_tool/v4/add_summary_memory.py b/reme_ai/mem_tool/v4/add_summary_memory.py index 6fe30cda..cc4be602 100644 --- a/reme_ai/mem_tool/v4/add_summary_memory.py +++ b/reme_ai/mem_tool/v4/add_summary_memory.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.schema import MemoryNode +from ...core.schema import MemoryNode class AddSummaryMemory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/v4/hands_off.py b/reme_ai/mem_tool/v4/hands_off.py index 2dab4531..fd818b4f 100644 --- a/reme_ai/mem_tool/v4/hands_off.py +++ b/reme_ai/mem_tool/v4/hands_off.py @@ -3,8 +3,8 @@ from typing import TYPE_CHECKING from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.enumeration import MemoryType -from ...core_old.schema import Message +from ...core.enumeration import MemoryType +from ...core.schema import Message if TYPE_CHECKING: from ...mem_agent import BaseMemoryAgent diff --git a/reme_ai/mem_tool/v4/read_history.py b/reme_ai/mem_tool/v4/read_history.py index 00097234..78a90eb4 100644 --- a/reme_ai/mem_tool/v4/read_history.py +++ b/reme_ai/mem_tool/v4/read_history.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.schema import MemoryNode +from ...core.schema import MemoryNode class ReadHistory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/v4/read_user_profile.py b/reme_ai/mem_tool/v4/read_user_profile.py index f45c0c34..ff963ea2 100644 --- a/reme_ai/mem_tool/v4/read_user_profile.py +++ b/reme_ai/mem_tool/v4/read_user_profile.py @@ -2,7 +2,7 @@ from typing import Literal from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.schema.memory_node import MemoryNode +from ...core.schema.memory_node import MemoryNode class ReadUserProfile(BaseMemoryTool): diff --git a/reme_ai/mem_tool/v4/retrieve_memory.py b/reme_ai/mem_tool/v4/retrieve_memory.py index a6717bfa..7a902c71 100644 --- a/reme_ai/mem_tool/v4/retrieve_memory.py +++ b/reme_ai/mem_tool/v4/retrieve_memory.py @@ -3,8 +3,8 @@ import json from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.schema import MemoryNode -from ...core_old.utils import deduplicate_memories +from ...core.schema import MemoryNode +from ...core.utils import deduplicate_memories class RetrieveMemory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/v4/update_user_profile.py b/reme_ai/mem_tool/v4/update_user_profile.py index 1ab3276c..a8fa04f5 100644 --- a/reme_ai/mem_tool/v4/update_user_profile.py +++ b/reme_ai/mem_tool/v4/update_user_profile.py @@ -1,8 +1,8 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.schema.memory_node import MemoryNode -from ...core_old.utils import deduplicate_memories +from ...core.schema.memory_node import MemoryNode +from ...core.utils import deduplicate_memories class UpdateUserProfile(BaseMemoryTool): diff --git a/reme_ai/mem_tool/vector_store/add_memory.py b/reme_ai/mem_tool/vector_store/add_memory.py index 0937fe2a..1f87df4b 100644 --- a/reme_ai/mem_tool/vector_store/add_memory.py +++ b/reme_ai/mem_tool/vector_store/add_memory.py @@ -3,8 +3,8 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C -from ...core_old.schema import MemoryNode +from ...core.context import C +from ...core.schema import MemoryNode @C.register_op() diff --git a/reme_ai/mem_tool/vector_store/add_summary_memory.py b/reme_ai/mem_tool/vector_store/add_summary_memory.py index 54abbd57..ce1127ed 100644 --- a/reme_ai/mem_tool/vector_store/add_summary_memory.py +++ b/reme_ai/mem_tool/vector_store/add_summary_memory.py @@ -3,9 +3,9 @@ from loguru import logger from .add_memory import AddMemory -from ...core_old.context import C -from ...core_old.enumeration import MemoryType -from ...core_old.schema import MemoryNode +from ...core.context import C +from ...core.enumeration import MemoryType +from ...core.schema import MemoryNode @C.register_op() diff --git a/reme_ai/mem_tool/vector_store/delete_memory.py b/reme_ai/mem_tool/vector_store/delete_memory.py index 95f4f210..45b28632 100644 --- a/reme_ai/mem_tool/vector_store/delete_memory.py +++ b/reme_ai/mem_tool/vector_store/delete_memory.py @@ -3,7 +3,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C +from ...core.context import C @C.register_op() diff --git a/reme_ai/mem_tool/vector_store/retrieve_recent_memory.py b/reme_ai/mem_tool/vector_store/retrieve_recent_memory.py index dfaa9240..0896d75a 100644 --- a/reme_ai/mem_tool/vector_store/retrieve_recent_memory.py +++ b/reme_ai/mem_tool/vector_store/retrieve_recent_memory.py @@ -3,9 +3,9 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C -from ...core_old.schema import MemoryNode, VectorNode -from ...core_old.utils import deduplicate_memories +from ...core.context import C +from ...core.schema import MemoryNode, VectorNode +from ...core.utils import deduplicate_memories @C.register_op() diff --git a/reme_ai/mem_tool/vector_store/update_memory.py b/reme_ai/mem_tool/vector_store/update_memory.py index fe08fb90..4873ce24 100644 --- a/reme_ai/mem_tool/vector_store/update_memory.py +++ b/reme_ai/mem_tool/vector_store/update_memory.py @@ -3,8 +3,8 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C -from ...core_old.schema import MemoryNode +from ...core.context import C +from ...core.schema import MemoryNode @C.register_op() diff --git a/reme_ai/mem_tool/vector_store/vector_retrieve_memory.py b/reme_ai/mem_tool/vector_store/vector_retrieve_memory.py index 24b05e8b..655de36e 100644 --- a/reme_ai/mem_tool/vector_store/vector_retrieve_memory.py +++ b/reme_ai/mem_tool/vector_store/vector_retrieve_memory.py @@ -3,10 +3,10 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.context import C -from ...core_old.enumeration import MemoryType -from ...core_old.schema import MemoryNode, VectorNode -from ...core_old.utils import deduplicate_memories +from ...core.context import C +from ...core.enumeration import MemoryType +from ...core.schema import MemoryNode, VectorNode +from ...core.utils import deduplicate_memories @C.register_op() diff --git a/reme_ai/mem_tool/wk/add_memory.py b/reme_ai/mem_tool/wk/add_memory.py index 2af7ba71..61cb154c 100644 --- a/reme_ai/mem_tool/wk/add_memory.py +++ b/reme_ai/mem_tool/wk/add_memory.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.schema import MemoryNode +from ...core.schema import MemoryNode class AddMemory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/wk/read_history.py b/reme_ai/mem_tool/wk/read_history.py index 65e8a0cf..945c02e2 100644 --- a/reme_ai/mem_tool/wk/read_history.py +++ b/reme_ai/mem_tool/wk/read_history.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.schema import MemoryNode +from ...core.schema import MemoryNode class ReadHistory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/wk/summary_and_hands_off.py b/reme_ai/mem_tool/wk/summary_and_hands_off.py index d384f16a..9a6a76cd 100644 --- a/reme_ai/mem_tool/wk/summary_and_hands_off.py +++ b/reme_ai/mem_tool/wk/summary_and_hands_off.py @@ -4,8 +4,8 @@ from typing import TYPE_CHECKING from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.enumeration import MemoryType -from ...core_old.schema import MemoryNode, Message +from ...core.enumeration import MemoryType +from ...core.schema import MemoryNode, Message if TYPE_CHECKING: from ...mem_agent import BaseMemoryAgent diff --git a/reme_ai/mem_tool/wk/update_memory.py b/reme_ai/mem_tool/wk/update_memory.py index 151eba53..c419261f 100644 --- a/reme_ai/mem_tool/wk/update_memory.py +++ b/reme_ai/mem_tool/wk/update_memory.py @@ -1,7 +1,7 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.schema import MemoryNode +from ...core.schema import MemoryNode class UpdateMemory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/wk/vector_retrieve_memory.py b/reme_ai/mem_tool/wk/vector_retrieve_memory.py index 2698cd6f..de7498b4 100644 --- a/reme_ai/mem_tool/wk/vector_retrieve_memory.py +++ b/reme_ai/mem_tool/wk/vector_retrieve_memory.py @@ -1,9 +1,9 @@ from loguru import logger from ..base_memory_tool import BaseMemoryTool -from ...core_old.enumeration import MemoryType -from ...core_old.schema import MemoryNode, VectorNode -from ...core_old.utils import deduplicate_memories +from ...core.enumeration import MemoryType +from ...core.schema import MemoryNode, VectorNode +from ...core.utils import deduplicate_memories class VectorRetrieveMemory(BaseMemoryTool): diff --git a/reme_ai/tool/execute/execute_code.py b/reme_ai/tool/execute/execute_code.py index f13aab3d..ea259487 100644 --- a/reme_ai/tool/execute/execute_code.py +++ b/reme_ai/tool/execute/execute_code.py @@ -4,11 +4,11 @@ This module provides an operation that can execute Python code strings and return the output or error messages. """ -from ...core_old.context import C -from ...core_old.op import BaseOp -from ...core_old.schema import ToolCall +from ...core.context import C +from ...core.op import BaseOp +from ...core.schema import ToolCall -from ...core_old.utils import exec_code +from ...core.utils import exec_code @C.register_op() diff --git a/reme_ai/tool/execute/execute_shell.py b/reme_ai/tool/execute/execute_shell.py index 1b235921..6e244ddb 100644 --- a/reme_ai/tool/execute/execute_shell.py +++ b/reme_ai/tool/execute/execute_shell.py @@ -4,11 +4,11 @@ This module provides an operation that can execute shell commands asynchronously and return the output, error, and exit code. """ -from ...core_old.context import C -from ...core_old.op import BaseOp -from ...core_old.schema import ToolCall +from ...core.context import C +from ...core.op import BaseOp +from ...core.schema import ToolCall -from ...core_old.utils import run_shell_command +from ...core.utils import run_shell_command @C.register_op() diff --git a/reme_ai/tool/search/dashscope_search.py b/reme_ai/tool/search/dashscope_search.py index 47c399ef..19bd8104 100644 --- a/reme_ai/tool/search/dashscope_search.py +++ b/reme_ai/tool/search/dashscope_search.py @@ -9,9 +9,9 @@ from typing import Literal from loguru import logger -from ...core_old.context import C -from ...core_old.op import BaseOp -from ...core_old.schema import ToolCall +from ...core.context import C +from ...core.op import BaseOp +from ...core.schema import ToolCall @C.register_op() diff --git a/reme_ai/tool/search/mock_search.py b/reme_ai/tool/search/mock_search.py index 69ef3995..187463dc 100644 --- a/reme_ai/tool/search/mock_search.py +++ b/reme_ai/tool/search/mock_search.py @@ -9,11 +9,11 @@ import random from loguru import logger -from ...core_old.context import C -from ...core_old.enumeration import Role -from ...core_old.op import BaseOp -from ...core_old.schema import ToolCall, Message -from ...core_old.utils import extract_content +from ...core.context import C +from ...core.enumeration import Role +from ...core.op import BaseOp +from ...core.schema import ToolCall, Message +from ...core.utils import extract_content @C.register_op() diff --git a/reme_ai/tool/search/tavily_search.py b/reme_ai/tool/search/tavily_search.py index bb000f16..5c194bdc 100644 --- a/reme_ai/tool/search/tavily_search.py +++ b/reme_ai/tool/search/tavily_search.py @@ -9,9 +9,9 @@ import os from loguru import logger -from ...core_old.context import C -from ...core_old.op import BaseOp -from ...core_old.schema import ToolCall +from ...core.context import C +from ...core.op import BaseOp +from ...core.schema import ToolCall @C.register_op() diff --git a/test/test_base_context.py b/test/test_base_context.py index 2b3ebc99..316a6796 100644 --- a/test/test_base_context.py +++ b/test/test_base_context.py @@ -4,7 +4,7 @@ Ensures attribute-style and dict-style access work interchangeably. """ import pickle -from reme_ai.core_old.context import BaseContext +from reme_ai.core.context import BaseContext def test_attribute_access(): diff --git a/test/test_cache_handler.py b/test/test_cache_handler.py index b741769a..ddcac86f 100644 --- a/test/test_cache_handler.py +++ b/test/test_cache_handler.py @@ -9,7 +9,7 @@ from pathlib import Path import pandas as pd from loguru import logger -from reme_ai.core_old.utils.cache_handler import CacheHandler +from reme_ai.core.utils.cache_handler import CacheHandler def run_tests(): diff --git a/test/test_embedding.py b/test/test_embedding.py index b769d05f..d1c404f5 100644 --- a/test/test_embedding.py +++ b/test/test_embedding.py @@ -18,12 +18,12 @@ import asyncio import argparse from typing import Type, List -from reme_ai.core_old.utils import load_env +from reme_ai.core.utils import load_env load_env() -from reme_ai.core_old.embedding import OpenAIEmbeddingModel, BaseEmbeddingModel -from reme_ai.core_old.schema import VectorNode +from reme_ai.core.embedding import OpenAIEmbeddingModel, BaseEmbeddingModel +from reme_ai.core.schema import VectorNode def get_embedding_model(model_class: Type[BaseEmbeddingModel]) -> BaseEmbeddingModel: diff --git a/test/test_embedding_sync.py b/test/test_embedding_sync.py index f97e28ec..361a42b3 100644 --- a/test/test_embedding_sync.py +++ b/test/test_embedding_sync.py @@ -17,12 +17,12 @@ Usage: import argparse from typing import Type, List -from reme_ai.core_old.utils import load_env +from reme_ai.core.utils import load_env load_env() -from reme_ai.core_old.embedding import OpenAIEmbeddingModelSync, BaseEmbeddingModel -from reme_ai.core_old.schema import VectorNode +from reme_ai.core.embedding import OpenAIEmbeddingModelSync, BaseEmbeddingModel +from reme_ai.core.schema import VectorNode def get_embedding_model(model_class: Type[BaseEmbeddingModel]) -> BaseEmbeddingModel: diff --git a/test/test_llm.py b/test/test_llm.py index 819c2b2b..12c6eca3 100644 --- a/test/test_llm.py +++ b/test/test_llm.py @@ -18,13 +18,13 @@ import asyncio import argparse from typing import Type -from reme_ai.core_old.utils import load_env +from reme_ai.core.utils import load_env load_env() -from reme_ai.core_old.llm import OpenAILLM, LiteLLM, BaseLLM -from reme_ai.core_old.schema import Message, ToolCall -from reme_ai.core_old.enumeration import Role, ChunkEnum +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 def get_llm(llm_class: Type[BaseLLM]) -> BaseLLM: diff --git a/test/test_llm_sync.py b/test/test_llm_sync.py index 07d80ed1..98751f87 100644 --- a/test/test_llm_sync.py +++ b/test/test_llm_sync.py @@ -17,13 +17,13 @@ Usage: import argparse from typing import Type -from reme_ai.core_old.utils import load_env +from reme_ai.core.utils import load_env load_env() -from reme_ai.core_old.llm import OpenAILLMSync, LiteLLMSync, BaseLLM -from reme_ai.core_old.schema import Message, ToolCall -from reme_ai.core_old.enumeration import Role, ChunkEnum +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 def get_llm(llm_class: Type[BaseLLM]) -> BaseLLM: diff --git a/test/test_logo.py b/test/test_logo.py index 9b4cf089..eeede81e 100644 --- a/test/test_logo.py +++ b/test/test_logo.py @@ -1,9 +1,9 @@ """test logo""" -from reme_ai.core_old.schema import ServiceConfig, MCPConfig +from reme_ai.core.schema import ServiceConfig, MCPConfig if __name__ == "__main__": - from reme_ai.core_old.utils import print_logo + from reme_ai.core.utils import print_logo c = ServiceConfig(app_name="reme", backend="mcp", mcp=MCPConfig(transport="sse")) print_logo(service_config=c) diff --git a/test/test_mcp_client.py b/test/test_mcp_client.py index 0ae6fc54..d2fef40e 100644 --- a/test/test_mcp_client.py +++ b/test/test_mcp_client.py @@ -5,7 +5,7 @@ import asyncio import json -from reme_ai.core_old.utils import MCPClient +from reme_ai.core.utils import MCPClient async def main(): diff --git a/test/test_mcp_server.py b/test/test_mcp_server.py index 4257b977..67f9c542 100644 --- a/test/test_mcp_server.py +++ b/test/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_old.schema import ToolCall -from reme_ai.core_old.utils import create_pydantic_model +from reme_ai.core.schema import ToolCall +from reme_ai.core.utils import create_pydantic_model mcp = FastMCP("DynamicSchemaServer", port=8010) diff --git a/test/test_message.py b/test/test_message.py index e77e673b..141174c5 100644 --- a/test/test_message.py +++ b/test/test_message.py @@ -4,8 +4,8 @@ import unittest from mcp.types import Tool -from reme_ai.core_old.enumeration import Role -from reme_ai.core_old.schema import ToolAttr, ToolCall, ContentBlock, Message +from reme_ai.core.enumeration import Role +from reme_ai.core.schema import ToolAttr, ToolCall, ContentBlock, Message class TestModelDefinitions(unittest.TestCase): diff --git a/test/test_op_composition.py b/test/test_op_composition.py index 319981a3..8d32b51a 100644 --- a/test/test_op_composition.py +++ b/test/test_op_composition.py @@ -5,8 +5,8 @@ Tests asynchronous execution mode. import asyncio -from reme_ai.core_old.op import BaseOp -from reme_ai.core_old.schema import ToolCall, ToolAttr +from reme_ai.core.op import BaseOp +from reme_ai.core.schema import ToolCall, ToolAttr class AddOp(BaseOp): diff --git a/test/test_reme.py b/test/test_reme.py index b23aa63e..c9e6843e 100644 --- a/test/test_reme.py +++ b/test/test_reme.py @@ -2,7 +2,7 @@ import asyncio -from reme_ai.core_old.schema import VectorNode, MemoryNode +from reme_ai.core.schema import VectorNode, MemoryNode from reme_ai.reme import ReMe reme = ReMe( diff --git a/test/test_timer.py b/test/test_timer.py index 1c9de938..c9714e38 100644 --- a/test/test_timer.py +++ b/test/test_timer.py @@ -7,7 +7,7 @@ import time from loguru import logger -from reme_ai.core_old.utils import timer +from reme_ai.core.utils import timer @timer diff --git a/test/test_token_counter.py b/test/test_token_counter.py index e67fb5c2..3c44a298 100644 --- a/test/test_token_counter.py +++ b/test/test_token_counter.py @@ -14,9 +14,9 @@ Usage: import argparse from typing import Type, List -from reme_ai.core_old.enumeration import Role -from reme_ai.core_old.schema import Message, ToolCall -from reme_ai.core_old.token_counter import BaseTokenCounter, OpenAITokenCounter, HFTokenCounter +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 def get_token_counter(counter_class: Type[BaseTokenCounter], **kwargs) -> BaseTokenCounter: diff --git a/test/test_tool.py b/test/test_tool.py index 765051b3..9db5a3ed 100644 --- a/test/test_tool.py +++ b/test/test_tool.py @@ -169,8 +169,8 @@ async def test_stream_chat(): process and stream responses in real-time using async operations. """ from reme_ai.mem_agent.chat import StreamChat - from reme_ai.core_old.utils import execute_stream_task - from reme_ai.core_old.context import RuntimeContext + from reme_ai.core.utils import execute_stream_task + from reme_ai.core.context import RuntimeContext from asyncio import Queue op = StreamChat() diff --git a/test/test_tool_call.py b/test/test_tool_call.py index 2ae7a619..30c0a37e 100644 --- a/test/test_tool_call.py +++ b/test/test_tool_call.py @@ -2,7 +2,7 @@ import json -from reme_ai.core_old.schema.tool_call import ToolCall +from reme_ai.core.schema.tool_call import ToolCall def test_simple_schema(): diff --git a/test/test_vector_store.py b/test/test_vector_store.py index 51edbd55..00927d0c 100644 --- a/test/test_vector_store.py +++ b/test/test_vector_store.py @@ -23,9 +23,9 @@ from typing import List from loguru import logger -from reme_ai.core_old.embedding import OpenAIEmbeddingModel -from reme_ai.core_old.schema import VectorNode -from reme_ai.core_old.vector_store import ( +from reme_ai.core.embedding import OpenAIEmbeddingModel +from reme_ai.core.schema import VectorNode +from reme_ai.core.vector_store import ( BaseVectorStore, ChromaVectorStore, LocalVectorStore, @@ -1453,8 +1453,8 @@ async def test_sql_injection_protection(store: BaseVectorStore, store_name: str) # Test 1: Invalid collection name (SQL injection attempt) try: - from reme_ai.core_old.vector_store import PGVectorStore - from reme_ai.core_old.embedding import OpenAIEmbeddingModel + from reme_ai.core.vector_store import PGVectorStore + from reme_ai.core.embedding import OpenAIEmbeddingModel embedding_model = OpenAIEmbeddingModel() From 4348148b723fb9fd30d3e9fe17acf8147996d3a0 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 22 Jan 2026 16:25:20 +0800 Subject: [PATCH 12/19] refactor(core): restructure core modules and update pre-commit configuration --- .pre-commit-config.yaml | 13 +- {reme_ai/core => reme}/config/__init__.py | 0 {reme_ai/core => reme}/config/default.yaml | 5 +- .../config/reme_config_parser.py | 2 +- reme/core/context/__init__.py | 4 + .../core/context/runtime_context.py | 13 +- reme/core/context/service_context.py | 230 ++++++++ {reme_ai => reme}/core/embedding/__init__.py | 4 + .../core/embedding/base_embedding_model.py | 0 .../core/embedding/openai_embedding_model.py | 2 - .../embedding/openai_embedding_model_sync.py | 2 - {reme_ai => reme}/core/flow/__init__.py | 2 - {reme_ai => reme}/core/flow/base_flow.py | 67 +-- {reme_ai => reme}/core/flow/cmd_flow.py | 0 .../core/flow/expression_flow.py | 11 +- {reme_ai => reme}/core/llm/__init__.py | 6 + {reme_ai => reme}/core/llm/base_llm.py | 181 ++---- {reme_ai => reme}/core/llm/lite_llm.py | 2 - {reme_ai => reme}/core/llm/lite_llm_sync.py | 2 - {reme_ai => reme}/core/llm/openai_llm.py | 2 - {reme_ai => reme}/core/llm/openai_llm_sync.py | 2 - {reme_ai => reme}/core/op/__init__.py | 5 + {reme_ai => reme}/core/op/base_op.py | 227 +++----- {reme_ai => reme}/core/op/base_ray_op.py | 10 +- reme/core/op/base_tool.py | 61 ++ {reme_ai => reme}/core/op/mcp_tool.py | 40 +- {reme_ai => reme}/core/op/parallel_op.py | 4 +- {reme_ai => reme}/core/op/sequential_op.py | 8 +- reme/core/schema/response.py | 4 +- reme/core/schema/service_config.py | 17 +- reme/core/schema/tool_call.py | 5 - {reme_ai => reme}/core/service/__init__.py | 5 + .../core/service/base_service.py | 10 +- {reme_ai => reme}/core/service/cmd_service.py | 10 +- .../core/service/http_service.py | 6 +- {reme_ai => reme}/core/service/mcp_service.py | 14 +- .../core/token_counter/__init__.py | 5 + .../core/token_counter/base_token_counter.py | 3 +- .../core/token_counter/hf_token_counter.py | 2 - .../token_counter/openai_token_counter.py | 4 +- reme/core/utils/__init__.py | 32 ++ {reme_ai => reme}/core/utils/cache_handler.py | 0 .../core/utils/case_converter.py | 0 {reme_ai => reme}/core/utils/common_utils.py | 0 {reme_ai => reme}/core/utils/env_utils.py | 0 .../core/utils/execute_utils.py | 0 {reme_ai => reme}/core/utils/http_client.py | 0 {reme_ai => reme}/core/utils/llm_utils.py | 0 {reme_ai => reme}/core/utils/logger_utils.py | 0 {reme_ai => reme}/core/utils/logo_utils.py | 0 {reme_ai => reme}/core/utils/mcp_client.py | 0 .../core/utils/pydantic_config_parser.py | 0 .../core/utils/pydantic_utils.py | 0 {reme_ai => reme}/core/utils/time.py | 0 .../core/vector_store/__init__.py | 7 + .../core/vector_store/base_vector_store.py | 15 +- .../core/vector_store/chroma_vector_store.py | 8 +- .../core/vector_store/es_vector_store.py | 38 +- .../core/vector_store/local_vector_store.py | 4 +- .../core/vector_store/pgvector_store.py | 26 +- .../core/vector_store/qdrant_vector_store.py | 10 +- reme/reme_app.py | 90 +++ reme_ai/core/__init__.py | 17 - reme_ai/core/application.py | 221 ------- reme_ai/core/context/__init__.py | 16 - reme_ai/core/context/base_context.py | 41 -- reme_ai/core/context/prompt_handler.py | 95 ---- reme_ai/core/context/registry.py | 46 -- reme_ai/core/context/service_context.py | 537 ------------------ reme_ai/core/enumeration/__init__.py | 17 - reme_ai/core/enumeration/chunk_enum.py | 25 - reme_ai/core/enumeration/http_enum.py | 22 - reme_ai/core/enumeration/json_schema_enum.py | 18 - reme_ai/core/enumeration/memory_type.py | 25 - reme_ai/core/enumeration/registry_enum.py | 28 - reme_ai/core/enumeration/role.py | 19 - reme_ai/core/flow/simple_flow.py | 17 - reme_ai/core/main.py | 48 -- reme_ai/core/schema/__init__.py | 42 -- reme_ai/core/schema/memory_node.py | 223 -------- reme_ai/core/schema/message.py | 164 ------ reme_ai/core/schema/request.py | 11 - reme_ai/core/schema/response.py | 11 - reme_ai/core/schema/service_config.py | 113 ---- reme_ai/core/schema/stream_chunk.py | 14 - reme_ai/core/schema/tool_call.py | 228 -------- reme_ai/core/schema/vector_node.py | 15 - reme_ai/core/utils/__init__.py | 47 -- reme_ai/{core => }/reme.py | 1 - 89 files changed, 758 insertions(+), 2523 deletions(-) rename {reme_ai/core => reme}/config/__init__.py (100%) rename {reme_ai/core => reme}/config/default.yaml (88%) rename {reme_ai/core => reme}/config/reme_config_parser.py (76%) rename {reme_ai => reme}/core/context/runtime_context.py (88%) create mode 100644 reme/core/context/service_context.py rename {reme_ai => reme}/core/embedding/__init__.py (65%) rename {reme_ai => reme}/core/embedding/base_embedding_model.py (100%) rename {reme_ai => reme}/core/embedding/openai_embedding_model.py (96%) rename {reme_ai => reme}/core/embedding/openai_embedding_model_sync.py (94%) rename {reme_ai => reme}/core/flow/__init__.py (77%) rename {reme_ai => reme}/core/flow/base_flow.py (80%) rename {reme_ai => reme}/core/flow/cmd_flow.py (100%) rename {reme_ai => reme}/core/flow/expression_flow.py (76%) rename {reme_ai => reme}/core/llm/__init__.py (60%) rename {reme_ai => reme}/core/llm/base_llm.py (67%) rename {reme_ai => reme}/core/llm/lite_llm.py (98%) rename {reme_ai => reme}/core/llm/lite_llm_sync.py (97%) rename {reme_ai => reme}/core/llm/openai_llm.py (98%) rename {reme_ai => reme}/core/llm/openai_llm_sync.py (97%) rename {reme_ai => reme}/core/op/__init__.py (72%) rename {reme_ai => reme}/core/op/base_op.py (61%) rename {reme_ai => reme}/core/op/base_ray_op.py (94%) create mode 100644 reme/core/op/base_tool.py rename {reme_ai => reme}/core/op/mcp_tool.py (75%) rename {reme_ai => reme}/core/op/parallel_op.py (93%) rename {reme_ai => reme}/core/op/sequential_op.py (84%) rename {reme_ai => reme}/core/service/__init__.py (64%) rename {reme_ai => reme}/core/service/base_service.py (76%) rename {reme_ai => reme}/core/service/cmd_service.py (75%) rename {reme_ai => reme}/core/service/http_service.py (95%) rename {reme_ai => reme}/core/service/mcp_service.py (85%) rename {reme_ai => reme}/core/token_counter/__init__.py (58%) rename {reme_ai => reme}/core/token_counter/base_token_counter.py (96%) rename {reme_ai => reme}/core/token_counter/hf_token_counter.py (97%) rename {reme_ai => reme}/core/token_counter/openai_token_counter.py (97%) rename {reme_ai => reme}/core/utils/cache_handler.py (100%) rename {reme_ai => reme}/core/utils/case_converter.py (100%) rename {reme_ai => reme}/core/utils/common_utils.py (100%) rename {reme_ai => reme}/core/utils/env_utils.py (100%) rename reme_ai/core/utils/execute_tuils.py => reme/core/utils/execute_utils.py (100%) rename {reme_ai => reme}/core/utils/http_client.py (100%) rename {reme_ai => reme}/core/utils/llm_utils.py (100%) rename {reme_ai => reme}/core/utils/logger_utils.py (100%) rename {reme_ai => reme}/core/utils/logo_utils.py (100%) rename {reme_ai => reme}/core/utils/mcp_client.py (100%) rename {reme_ai => reme}/core/utils/pydantic_config_parser.py (100%) rename {reme_ai => reme}/core/utils/pydantic_utils.py (100%) rename {reme_ai => reme}/core/utils/time.py (100%) rename {reme_ai => reme}/core/vector_store/__init__.py (62%) rename {reme_ai => reme}/core/vector_store/base_vector_store.py (91%) rename {reme_ai => reme}/core/vector_store/chroma_vector_store.py (99%) rename {reme_ai => reme}/core/vector_store/es_vector_store.py (95%) rename {reme_ai => reme}/core/vector_store/local_vector_store.py (99%) rename {reme_ai => reme}/core/vector_store/pgvector_store.py (96%) rename {reme_ai => reme}/core/vector_store/qdrant_vector_store.py (98%) create mode 100644 reme/reme_app.py delete mode 100644 reme_ai/core/__init__.py delete mode 100644 reme_ai/core/application.py delete mode 100644 reme_ai/core/context/__init__.py delete mode 100644 reme_ai/core/context/base_context.py delete mode 100644 reme_ai/core/context/prompt_handler.py delete mode 100644 reme_ai/core/context/registry.py delete mode 100644 reme_ai/core/context/service_context.py delete mode 100644 reme_ai/core/enumeration/__init__.py delete mode 100644 reme_ai/core/enumeration/chunk_enum.py delete mode 100644 reme_ai/core/enumeration/http_enum.py delete mode 100644 reme_ai/core/enumeration/json_schema_enum.py delete mode 100644 reme_ai/core/enumeration/memory_type.py delete mode 100644 reme_ai/core/enumeration/registry_enum.py delete mode 100644 reme_ai/core/enumeration/role.py delete mode 100644 reme_ai/core/flow/simple_flow.py delete mode 100644 reme_ai/core/main.py delete mode 100644 reme_ai/core/schema/__init__.py delete mode 100644 reme_ai/core/schema/memory_node.py delete mode 100644 reme_ai/core/schema/message.py delete mode 100644 reme_ai/core/schema/request.py delete mode 100644 reme_ai/core/schema/response.py delete mode 100644 reme_ai/core/schema/service_config.py delete mode 100644 reme_ai/core/schema/stream_chunk.py delete mode 100644 reme_ai/core/schema/tool_call.py delete mode 100644 reme_ai/core/schema/vector_node.py delete mode 100644 reme_ai/core/utils/__init__.py rename reme_ai/{core => }/reme.py (99%) diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 3d6fcf96..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/|reme_ai/core_old/|reme_ai/mem_agent/|reme_ai/mem_tool/|bench) + 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/|reme_ai/core_old/|reme_ai/mem_agent/|reme_ai/mem_tool/|bench) + exclude: ^(test/|cookbook/|reme_ai/|bench) - repo: https://github.com/psf/black rev: 25.9.0 hooks: - id: black - exclude: ^(test/|cookbook/|reme_ai/core_old/|reme_ai/mem_agent/|reme_ai/mem_tool/|bench) + 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/|reme_ai/core_old/|reme_ai/mem_agent/|reme_ai/mem_tool/|bench) + exclude: ^(test/|cookbook/|reme_ai/|bench) args: [ "--extend-ignore=E203", "--max-line-length=120" @@ -44,9 +44,7 @@ repos: | \.demo$ | \.md$ | \.html$ - | reme_ai/core_old/ - | reme_ai/mem_agent/ - | reme_ai/mem_tool/ + | reme_ai/ | bench ) args: [ @@ -80,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/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 88% rename from reme_ai/core/config/default.yaml rename to reme/config/default.yaml index 00ef4062..836d866f 100644 --- a/reme_ai/core/config/default.yaml +++ b/reme/config/default.yaml @@ -15,15 +15,14 @@ http: llm: default: backend: openai -# model_name: qwen3-30b-a3b-instruct-2507 - model_name: qwen-flash + model_name: qwen3-30b-a3b-instruct-2507 +# model_name: qwen-flash request_interval: 1 temperature: 0.0001 qwen3_max_instruct: backend: openai model_name: qwen3-max -# temperature: 0.6 request_interval: 2 embedding_model: 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/core/context/__init__.py b/reme/core/context/__init__.py index 27957fd2..7bbd5869 100644 --- a/reme/core/context/__init__.py +++ b/reme/core/context/__init__.py @@ -3,9 +3,13 @@ from .base_context import BaseContext from .prompt_handler import PromptHandler from .registry_factory import R +from .runtime_context import RuntimeContext +from .service_context import ServiceContext __all__ = [ "BaseContext", "PromptHandler", "R", + "RuntimeContext", + "ServiceContext", ] diff --git a/reme_ai/core/context/runtime_context.py b/reme/core/context/runtime_context.py similarity index 88% rename from reme_ai/core/context/runtime_context.py rename to reme/core/context/runtime_context.py index d7112e1c..44167ca0 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,23 @@ 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: + 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/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 67% rename from reme_ai/core/llm/base_llm.py rename to reme/core/llm/base_llm.py index fcd543b6..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,16 +15,23 @@ 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 the LLM client with model configurations and retry policies. + def __init__( + self, + model_name: str, + max_retries: int = 10, + raise_exception: bool = False, + request_interval: float = 0.0, + **kwargs, + ): + """Initialize LLM client. 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 - request_interval: Minimum time interval (in seconds) between consecutive requests. Default is 0.0 (no interval). + 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 @@ -33,20 +40,17 @@ class BaseLLM(ABC): self.request_interval: float = request_interval self.kwargs: dict = kwargs - # Request rate control for async operations 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 @@ -58,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 [] @@ -69,10 +73,9 @@ class BaseLLM(ABC): if tool.name not in tool_dict: continue - # First try sanitizing arguments if not tool.sanitize_and_check_argument(): - logger.error(f"Tool call {tool.name} has invalid JSON arguments after sanitization attempt: {tool.arguments}") - raise ValueError(f"Tool call {tool.name} has invalid JSON arguments: {tool.arguments}") + 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 @@ -86,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, @@ -102,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( @@ -111,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( @@ -121,22 +116,13 @@ 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 request rate limiting if configured + """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: - sleep_time = self.request_interval - elapsed - await asyncio.sleep(sleep_time) + 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): @@ -149,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): @@ -159,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: @@ -177,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): @@ -193,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: @@ -212,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, @@ -231,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( @@ -280,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, @@ -299,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( @@ -350,28 +311,24 @@ 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 request rate limiting if configured + """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: - sleep_time = self.request_interval - elapsed - await asyncio.sleep(sleep_time) + 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) + return await self._chat_impl( + messages, + tools, + enable_stream_print, + callback_fn, + default_value, + model_name, + **kwargs, + ) async def _chat_impl( self, @@ -383,8 +340,7 @@ 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): @@ -399,19 +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() or - "exceeded your current quota" in error_message.lower() or - "insufficient_quota" in error_message.lower() + "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}]:") @@ -422,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: @@ -450,18 +404,7 @@ 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): @@ -476,19 +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() or - "exceeded your current quota" in error_message.lower() or - "insufficient_quota" in error_message.lower() + "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}]:") @@ -499,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: @@ -518,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 98% rename from reme_ai/core/llm/lite_llm.py rename to reme/core/llm/lite_llm.py index 458fe1c0..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.""" diff --git a/reme_ai/core/llm/lite_llm_sync.py b/reme/core/llm/lite_llm_sync.py similarity index 97% rename from reme_ai/core/llm/lite_llm_sync.py rename to reme/core/llm/lite_llm_sync.py index 11fbc7af..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.""" diff --git a/reme_ai/core/llm/openai_llm.py b/reme/core/llm/openai_llm.py similarity index 98% rename from reme_ai/core/llm/openai_llm.py rename to reme/core/llm/openai_llm.py index f15dcb8e..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.""" diff --git a/reme_ai/core/llm/openai_llm_sync.py b/reme/core/llm/openai_llm_sync.py similarity index 97% rename from reme_ai/core/llm/openai_llm_sync.py rename to reme/core/llm/openai_llm_sync.py index a2bcfee3..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.""" 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..9daf033a --- /dev/null +++ b/reme/core/op/base_tool.py @@ -0,0 +1,61 @@ +"""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 | None: + """Get the tool call schema.""" + 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 + 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/core/schema/response.py b/reme/core/schema/response.py index 3104bc6e..fe753232 100644 --- a/reme/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/core/schema/service_config.py b/reme/core/schema/service_config.py index e1ae9df7..95461196 100644 --- a/reme/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) + 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/core/schema/tool_call.py b/reme/core/schema/tool_call.py index 70355d60..e7ba9b78 100644 --- a/reme/core/schema/tool_call.py +++ b/reme/core/schema/tool_call.py @@ -103,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: 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/core/utils/__init__.py b/reme/core/utils/__init__.py index ce9d9b2d..3d54e4e4 100644 --- a/reme/core/utils/__init__.py +++ b/reme/core/utils/__init__.py @@ -1,7 +1,39 @@ """utils""" +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_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 +from .logo_utils import print_logo +from .mcp_client import MCPClient +from .pydantic_config_parser import PydanticConfigParser +from .pydantic_utils import create_pydantic_model from .singleton import singleton +from .time import timer, get_now_time __all__ = [ + "CacheHandler", + "snake_to_camel", + "camel_to_snake", + "run_coro_safely", + "execute_stream_task", + "load_env", + "exec_code", + "run_shell_command", + "HttpClient", + "extract_content", + "format_messages", + "deduplicate_memories", + "init_logger", + "print_logo", + "MCPClient", + "PydanticConfigParser", + "create_pydantic_model", "singleton", + "timer", + "get_now_time", ] diff --git a/reme_ai/core/utils/cache_handler.py b/reme/core/utils/cache_handler.py similarity index 100% rename from reme_ai/core/utils/cache_handler.py rename to reme/core/utils/cache_handler.py 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/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 99% rename from reme_ai/core/vector_store/chroma_vector_store.py rename to reme/core/vector_store/chroma_vector_store.py index 567ac3b1..eaddd0e7 100644 --- a/reme_ai/core/vector_store/chroma_vector_store.py +++ b/reme/core/vector_store/chroma_vector_store.py @@ -5,7 +5,6 @@ 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 +19,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.""" @@ -142,7 +140,7 @@ class ChromaVectorStore(BaseVectorStore): # ChromaDB requires separate conditions combined with $and return [ {k: {"$gte": v[0]}}, - {k: {"$lte": v[1]}} + {k: {"$lte": v[1]}}, ] if isinstance(v, dict): chroma_condition = {} @@ -239,8 +237,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 95% rename from reme_ai/core/vector_store/es_vector_store.py rename to reme/core/vector_store/es_vector_store.py index 39989277..49054afc 100644 --- a/reme_ai/core/vector_store/es_vector_store.py +++ b/reme/core/vector_store/es_vector_store.py @@ -9,7 +9,6 @@ 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 +23,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.""" @@ -265,14 +263,16 @@ class ESVectorStore(BaseVectorStore): # 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] - } - } - }) + filter_conditions.append( + { + "range": { + f"metadata.{key}": { + "gte": value[0], + "lte": value[1], + }, + }, + }, + ) else: # Exact match filter_conditions.append({"term": {f"metadata.{key}": value}}) @@ -461,14 +461,16 @@ class ESVectorStore(BaseVectorStore): # 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] - } - } - }) + filter_conditions.append( + { + "range": { + f"metadata.{key}": { + "gte": value[0], + "lte": value[1], + }, + }, + }, + ) else: # Exact match filter_conditions.append({"term": {f"metadata.{key}": value}}) diff --git a/reme_ai/core/vector_store/local_vector_store.py b/reme/core/vector_store/local_vector_store.py similarity index 99% rename from reme_ai/core/vector_store/local_vector_store.py rename to reme/core/vector_store/local_vector_store.py index 7533fa89..24e1977c 100644 --- a/reme_ai/core/vector_store/local_vector_store.py +++ b/reme/core/vector_store/local_vector_store.py @@ -6,12 +6,10 @@ 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.""" @@ -110,7 +108,7 @@ class LocalVectorStore(BaseVectorStore): return False try: # Try numeric comparison - if not (value[0] <= node_value <= value[1]): + if not value[0] <= node_value <= value[1]: return False except TypeError: # If comparison fails, the filter doesn't match diff --git a/reme_ai/core/vector_store/pgvector_store.py b/reme/core/vector_store/pgvector_store.py similarity index 96% rename from reme_ai/core/vector_store/pgvector_store.py rename to reme/core/vector_store/pgvector_store.py index 882ced3e..7d1b3401 100644 --- a/reme_ai/core/vector_store/pgvector_store.py +++ b/reme/core/vector_store/pgvector_store.py @@ -7,7 +7,6 @@ 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 @@ -22,7 +21,6 @@ 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.""" @@ -39,10 +37,10 @@ class PGVectorStore(BaseVectorStore): 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): + 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." + "and contain only alphanumeric characters and underscores.", ) def __init__( @@ -295,8 +293,10 @@ class PGVectorStore(BaseVectorStore): for key, value in filters.items(): # 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.") + 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: @@ -305,13 +305,12 @@ class PGVectorStore(BaseVectorStore): if isinstance(value[0], (int, float)) and isinstance(value[1], (int, float)): # Numeric range query conditions.append( - f"(metadata->>'{key}')::numeric >= ${param_idx} AND (metadata->>'{key}')::numeric <= ${param_idx + 1}" + 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}" - ) + conditions.append(f"metadata->>'{key}' >= ${param_idx} AND metadata->>'{key}' <= ${param_idx + 1}") params.extend([value[0], value[1]]) param_idx += 2 else: @@ -341,12 +340,9 @@ class PGVectorStore(BaseVectorStore): # Adjust parameter indices in filter clause to account for $1 being used by vector_str if filter_clause: - # Replace from highest index to lowest to avoid conflicts for i in range(len(filter_params), 0, -1): - old_placeholder = f"${i}" new_placeholder = f"${i + 1}" - # Use word boundary to ensure we only replace exact matches (e.g., $1 not $10) - filter_clause = re.sub(rf'\${i}\b', new_placeholder, filter_clause) + filter_clause = re.sub(rf"\${i}\b", new_placeholder, filter_clause) async with pool.acquire() as conn: sql = f""" @@ -417,7 +413,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 98% rename from reme_ai/core/vector_store/qdrant_vector_store.py rename to reme/core/vector_store/qdrant_vector_store.py index d0b4a8aa..5ca74cd0 100644 --- a/reme_ai/core/vector_store/qdrant_vector_store.py +++ b/reme/core/vector_store/qdrant_vector_store.py @@ -5,7 +5,6 @@ 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 +35,6 @@ except ImportError as e: VectorParams = None -@C.register_vector_store("qdrant") class QdrantVectorStore(BaseVectorStore): """Vector store implementation using Qdrant for dense vector search.""" @@ -274,7 +272,7 @@ class QdrantVectorStore(BaseVectorStore): 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." + f"Consider using numeric timestamps instead.", ) elif isinstance(value, dict) and ("gte" in value or "lte" in value): range_params = {} @@ -284,7 +282,8 @@ class QdrantVectorStore(BaseVectorStore): range_params["gte"] = value["gte"] else: logger.warning( - f"Qdrant range filter for key '{key}' requires numeric gte value, got {type(value['gte']).__name__}. Skipping." + f"Qdrant range filter for key '{key}' requires numeric gte value, " + f"got {type(value['gte']).__name__}. Skipping.", ) continue if "lte" in value: @@ -292,7 +291,8 @@ class QdrantVectorStore(BaseVectorStore): range_params["lte"] = value["lte"] else: logger.warning( - f"Qdrant range filter for key '{key}' requires numeric lte value, got {type(value['lte']).__name__}. Skipping." + f"Qdrant range filter for key '{key}' requires numeric lte value, " + f"got {type(value['lte']).__name__}. Skipping.", ) continue diff --git a/reme/reme_app.py b/reme/reme_app.py new file mode 100644 index 00000000..483cd5bb --- /dev/null +++ b/reme/reme_app.py @@ -0,0 +1,90 @@ +"""ReMe application classes for simplified configuration and execution.""" + +import asyncio +import sys + +from reme.core.utils import execute_stream_task +from .config import ReMeConfigParser +from .core.context import ServiceContext +from .core.flow import BaseFlow +from .core.schema import Response + + +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_ai/core/__init__.py b/reme_ai/core/__init__.py deleted file mode 100644 index 8eab5792..00000000 --- a/reme_ai/core/__init__.py +++ /dev/null @@ -1,17 +0,0 @@ -"""Core module for ReMe AI framework.""" - -# pylint: disable=wrong-import-position -# flake8: noqa: F401 - -from . import config -from . import context -from . import embedding -from . import enumeration -from . import flow -from . import llm -from . import op -from . import schema -from . import service -from . import token_counter -from . import utils -from . import vector_store 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/__init__.py b/reme_ai/core/context/__init__.py deleted file mode 100644 index 7f26d600..00000000 --- a/reme_ai/core/context/__init__.py +++ /dev/null @@ -1,16 +0,0 @@ -"""context""" - -from .base_context import BaseContext -from .prompt_handler import PromptHandler -from .registry import Registry -from .runtime_context import RuntimeContext -from .service_context import ServiceContext, C - -__all__ = [ - "BaseContext", - "PromptHandler", - "Registry", - "RuntimeContext", - "ServiceContext", - "C", -] diff --git a/reme_ai/core/context/base_context.py b/reme_ai/core/context/base_context.py deleted file mode 100644 index dabd8cdb..00000000 --- a/reme_ai/core/context/base_context.py +++ /dev/null @@ -1,41 +0,0 @@ -"""Module providing a dictionary subclass with attribute-style access and pickling support.""" - -from typing import Generic, TypeVar - -_KT = TypeVar("_KT") -_VT = TypeVar("_VT") - - -class BaseContext(dict, Generic[_KT, _VT]): - """A dictionary subclass that enables accessing and modifying keys as attributes.""" - - def __getattr__(self, name: str) -> _VT: - """Retrieve a dictionary item as an attribute.""" - try: - return self[name] - except KeyError as e: - raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") from e - - def __setattr__(self, name: str, value: _VT) -> None: - """Assign a value to a dictionary item using attribute syntax.""" - self[name] = value - - def __delattr__(self, name: str) -> None: - """Remove a dictionary item using attribute syntax.""" - try: - # Delete item from dict via key - del self[name] - except KeyError as e: - raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") from e - - def __getstate__(self) -> dict: - """Return the dictionary representation for pickling.""" - return dict(self) - - def __setstate__(self, state: dict) -> None: - """Restore the dictionary state from a pickled object.""" - self.update(state) - - def __reduce__(self): - """Define the reconstruction logic for pickling processes.""" - return self.__class__, (), self.__getstate__() diff --git a/reme_ai/core/context/prompt_handler.py b/reme_ai/core/context/prompt_handler.py deleted file mode 100644 index b428b163..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].strip() - - 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 08ea2271..00000000 --- a/reme_ai/core/context/registry.py +++ /dev/null @@ -1,46 +0,0 @@ -"""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 - -T = TypeVar('T') - - -class Registry(BaseContext): - """A registry container that uses decorators to map and store class references.""" - - def register(self, name: str | type = "", add_cls: bool = True) -> Callable[[type[T]], type[T]] | type[T]: - """Return a decorator that registers a class under a specific name in the registry. - - Can be used in three ways: - - @C.register_op() # with empty parentheses, uses class name - - @C.register_op # without parentheses, uses class name - - @C.register_op("custom_name") # with custom name - - Args: - name: Either a string name for the class, or the class itself when used without parentheses - add_cls: Whether to actually add the class to the registry - - Returns: - Either a decorator function or the registered class itself - """ - - def decorator(cls): - if add_cls: - # Use provided name or default to the class name as the key - key = name if isinstance(name, str) and name else cls.__name__ - self[key] = cls - return cls - - # If used without parentheses: @C.register_op - if inspect.isclass(name): - cls = name - # Register with class name as key - if add_cls: - self[cls.__name__] = cls - return cls - - # If used with parentheses: @C.register_op() or @C.register_op("name") - 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/__init__.py b/reme_ai/core/enumeration/__init__.py deleted file mode 100644 index 7202a949..00000000 --- a/reme_ai/core/enumeration/__init__.py +++ /dev/null @@ -1,17 +0,0 @@ -"""enumeration""" - -from .chunk_enum import ChunkEnum -from .http_enum import HttpEnum -from .json_schema_enum import JsonSchemaEnum -from .memory_type import MemoryType -from .registry_enum import RegistryEnum -from .role import Role - -__all__ = [ - "ChunkEnum", - "HttpEnum", - "JsonSchemaEnum", - "MemoryType", - "RegistryEnum", - "Role", -] diff --git a/reme_ai/core/enumeration/chunk_enum.py b/reme_ai/core/enumeration/chunk_enum.py deleted file mode 100644 index dbe37106..00000000 --- a/reme_ai/core/enumeration/chunk_enum.py +++ /dev/null @@ -1,25 +0,0 @@ -"""Defines the types of data chunks used in streaming responses.""" - -from enum import Enum - - -class ChunkEnum(str, Enum): - """Enumeration of possible chunk categories for stream processing.""" - - # Internal reasoning or chain-of-thought process - THINK = "think" - - # The final generated response content - ANSWER = "answer" - - # Metadata or calls related to external tools - TOOL = "tool" - - # Resource consumption and token usage statistics - USAGE = "usage" - - # Error messages or exception details - ERROR = "error" - - # Final signal indicating the completion of the stream - DONE = "done" diff --git a/reme_ai/core/enumeration/http_enum.py b/reme_ai/core/enumeration/http_enum.py deleted file mode 100644 index 19622242..00000000 --- a/reme_ai/core/enumeration/http_enum.py +++ /dev/null @@ -1,22 +0,0 @@ -"""Provides a collection of standard HTTP request methods.""" - -from enum import Enum - - -class HttpEnum(str, Enum): - """Enumeration of supported HTTP methods for network requests.""" - - # Retrieves data from a specified resource - GET = "get" - - # Submits data to be processed to a specified resource - POST = "post" - - # Identical to GET but only retrieves the response headers - HEAD = "head" - - # Uploads or replaces the representation of a target resource - PUT = "put" - - # Deletes the specified resource from the server - DELETE = "delete" 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/enumeration/registry_enum.py b/reme_ai/core/enumeration/registry_enum.py deleted file mode 100644 index 876c06b8..00000000 --- a/reme_ai/core/enumeration/registry_enum.py +++ /dev/null @@ -1,28 +0,0 @@ -"""Defines the registry categories for core components of the system.""" - -from enum import Enum - - -class RegistryEnum(str, Enum): - """Enumeration of component types registered within the application lifecycle.""" - - # Large Language Model interfaces - LLM = "llm" - - # Models used for generating vector embeddings - EMBEDDING_MODEL = "embedding_model" - - # Databases or storage systems for vector search - VECTOR_STORE = "vector_store" - - # Atomic operations or functional units - OP = "op" - - # Orchestrated sequences of operations or workflows - FLOW = "flow" - - # External APIs or shared internal services - SERVICE = "service" - - # Utilities for tracking and limiting token consumption - TOKEN_COUNTER = "token_counter" diff --git a/reme_ai/core/enumeration/role.py b/reme_ai/core/enumeration/role.py deleted file mode 100644 index 4acad7e5..00000000 --- a/reme_ai/core/enumeration/role.py +++ /dev/null @@ -1,19 +0,0 @@ -"""Defines the participant roles in a chat completion sequence.""" - -from enum import Enum - - -class Role(str, Enum): - """Enumeration of standard personas involved in a conversation flow.""" - - # High-level instructions to guide the model's behavior - SYSTEM = "system" - - # Input or queries provided by the human user - USER = "user" - - # Responses or messages generated by the AI model - ASSISTANT = "assistant" - - # Output or results returned from external tool executions - TOOL = "tool" 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/core/schema/__init__.py b/reme_ai/core/schema/__init__.py deleted file mode 100644 index b7b73719..00000000 --- a/reme_ai/core/schema/__init__.py +++ /dev/null @@ -1,42 +0,0 @@ -"""schema""" - -from .memory_node import MemoryNode -from .message import ContentBlock, Message, Trajectory -from .request import Request -from .response import Response -from .service_config import ( - CmdConfig, - EmbeddingModelConfig, - FlowConfig, - HttpConfig, - LLMConfig, - MCPConfig, - ServiceConfig, - TokenCounterConfig, - VectorStoreConfig, -) -from .stream_chunk import StreamChunk -from .tool_call import ToolAttr, ToolCall -from .vector_node import VectorNode - -__all__ = [ - "MemoryNode", - "ContentBlock", - "EmbeddingModelConfig", - "FlowConfig", - "HttpConfig", - "LLMConfig", - "MCPConfig", - "Message", - "Request", - "Response", - "ServiceConfig", - "StreamChunk", - "TokenCounterConfig", - "Trajectory", - "ToolAttr", - "ToolCall", - "VectorNode", - "VectorStoreConfig", - "CmdConfig", -] diff --git a/reme_ai/core/schema/memory_node.py b/reme_ai/core/schema/memory_node.py deleted file mode 100644 index 304dfe5c..00000000 --- a/reme_ai/core/schema/memory_node.py +++ /dev/null @@ -1,223 +0,0 @@ -"""Memory schema module for the ReMe AI system. - -This module defines the MemoryNode class for storing and retrieving -memories in the ReMe system. -""" - -import datetime -import hashlib -import json -from typing import Any - -from pydantic import BaseModel, Field, model_validator - -from .vector_node import VectorNode -from ..enumeration import MemoryType - - -def get_now_time() -> str: - """Get current timestamp in YYYY-MM-DD HH:MM:SS format. - - Returns: - str: Current timestamp string in format 'YYYY-MM-DD HH:MM:SS'. - """ - return datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S") - - -# Length of the memory ID (first N characters of SHA-256 hash) -MEMORY_ID_LENGTH: int = 16 - - -class MemoryNode(BaseModel): - """Memory node for storing memories in the ReMe system. - - Attributes: - memory_id: Unique identifier, auto-generated from content hash. - memory_type: Type of memory (e.g., SUMMARY, PERSONAL). - memory_target: Target or topic this memory relates to. - when_to_use: Condition description for vector retrieval. - content: Actual memory content. - ref_memory_id: Reference to related raw history memory. - time_created: Creation timestamp. - time_modified: Last modification timestamp. - author: Author or source of this memory. - score: Relevance or importance score. - metadata: Additional metadata for extensibility. - """ - - memory_id: str = Field(default="", description="Unique memory identifier") - memory_type: MemoryType = Field(default=..., description="Type of memory") - memory_target: str = Field(default="", description="Target or topic of the memory") - when_to_use: str = Field(default="", description="Condition description for vector retrieval") - content: str = Field(default="", description="Actual memory content") - ref_memory_id: str = Field(default="", description="Reference to related raw history memory ID") - - time_created: str = Field(default_factory=get_now_time, description="Creation timestamp") - time_modified: str = Field(default_factory=get_now_time, description="Last modification timestamp") - author: str = Field(default="", description="Author or source of the memory") - score: float = Field(default=0, description="Relevance or importance score") - - metadata: dict[str, Any] = Field(default_factory=dict, description="Additional metadata") - - def _update_modified_time(self) -> "MemoryNode": - """Update time_modified to current timestamp. - - Returns: - Self: Returns self for method chaining. - """ - self.time_modified = get_now_time() - return self - - def _update_memory_id(self) -> "MemoryNode": - """Generate memory_id from SHA-256 hash of content. - - Takes the first MEMORY_ID_LENGTH characters of the hash. - - Returns: - Self: Returns self for method chaining. - """ - if not self.content: - return self - - hash_obj = hashlib.sha256(self.content.encode("utf-8")) - hex_dig = hash_obj.hexdigest() - self.memory_id = hex_dig[:MEMORY_ID_LENGTH] - return self - - @model_validator(mode="after") - def _update_after_init(self) -> "MemoryNode": - """Post-initialization validator. - - Auto-generates memory_id from content if not provided. - - Returns: - Self: Returns self for method chaining. - """ - if not self.memory_id: - self._update_memory_id() - return self - - def __setattr__(self, name: str, value): - """Auto-update timestamps and memory_id when content or when_to_use changes. - - Args: - name: Attribute name being set. - value: New value for the attribute. - """ - should_update: bool = name in ("when_to_use", "content") and getattr(self, name, None) != value - super().__setattr__(name, value) - if should_update: - self._update_modified_time() - if name == "content": - self._update_memory_id() - - def to_vector_node(self) -> VectorNode: - """Convert to VectorNode for vector storage. - - When when_to_use is set, use it as vector content and store content in metadata. - When when_to_use is empty, use content as vector content directly. - - Returns: - VectorNode: Vector node representation of this memory. - """ - # Build base metadata (shared fields) - metadata: dict[str, Any] = { - "memory_type": self.memory_type.value, - "memory_target": self.memory_target, - "ref_memory_id": self.ref_memory_id, - "time_created": self.time_created, - "time_modified": self.time_modified, - "author": self.author, - "score": self.score, - **self.metadata, - } - - if self.when_to_use: - # Use when_to_use for vector embedding, store content in metadata - vector_content = self.when_to_use - metadata["content"] = self.content - else: - # Use content directly for vector embedding - vector_content = self.content - - return VectorNode( - vector_id=self.memory_id, - content=vector_content, - 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. - - Reverses the to_vector_node conversion: - - If metadata contains 'content': node.content -> when_to_use, metadata['content'] -> content - - Otherwise: node.content -> content, when_to_use remains empty - - Args: - node: VectorNode containing memory data. - - Returns: - Self: Reconstructed MemoryNode instance. - - Raises: - ValueError: If memory_type in metadata is invalid. - """ - metadata = node.metadata.copy() - memory_type_str = metadata.pop("memory_type", None) - - try: - memory_type: MemoryType = MemoryType(memory_type_str) - except ValueError as e: - raise ValueError( - f"Invalid memory_type '{memory_type_str}' in VectorNode metadata. " - f"Valid types are: {[t.value for t in MemoryType]}", - ) from e - - # Restore when_to_use and content based on metadata structure - if "content" in metadata: - # Original had when_to_use set - when_to_use = node.content - content = metadata.pop("content", "") - else: - # Original had empty when_to_use - when_to_use = "" - content = node.content - - return cls( - memory_id=node.vector_id, - memory_type=memory_type, - memory_target=metadata.pop("memory_target", ""), - when_to_use=when_to_use, - content=content, - ref_memory_id=metadata.pop("ref_memory_id", ""), - time_created=metadata.pop("time_created", ""), - time_modified=metadata.pop("time_modified", ""), - author=metadata.pop("author", ""), - score=metadata.pop("score", 0), - metadata=metadata, - ) diff --git a/reme_ai/core/schema/message.py b/reme_ai/core/schema/message.py deleted file mode 100644 index 6c3299e7..00000000 --- a/reme_ai/core/schema/message.py +++ /dev/null @@ -1,164 +0,0 @@ -"""Data models for multi-modal conversation history and LLM interaction trajectories.""" - -import datetime -import json -import re - -from pydantic import BaseModel, ConfigDict, Field, model_validator - -from .tool_call import ToolCall -from ..enumeration import Role - - -class ContentBlock(BaseModel): - """ - Individual unit of multi-modal content like text, images, or video. - examples: - { - "type": "image_url", - "image_url": { - "url": "https://img.alicdn.com/imgextra/i1/O1CN01gDEY8M1W114Hi3XcN_!!6000000002727-0-tps-1024-406.jpg" - }, - } - - { - "type": "video", - "video": [ - "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/xzsgiz/football1.jpg", - "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/tdescd/football2.jpg", - "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/zefdja/football3.jpg", - "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/aedbqh/football4.jpg", - ], - } - - { - "type": "text", - "text": "How do you solve this problem?" - } - """ - - model_config = ConfigDict(extra="allow") - - type: str = Field(default="") - content: str | dict | list = Field(default="") - - @model_validator(mode="before") - @classmethod - def init_block(cls, data: dict) -> dict: - """Dynamically maps the type-specific key to the content field.""" - content_type = data.get("type", "") - if content_type and content_type in data: - data["content"] = data[content_type] - return data - - def simple_dump(self) -> dict: - """Serializes the block into an API-compatible dictionary format.""" - return { - "type": self.type, - self.type: self.content, - **self.model_extra, - } - - -class Message(BaseModel): - """Data model for a single dialogue entry including roles and tool interactions.""" - - name: str | None = Field(default=None) - role: Role = Field(default=Role.USER) - content: str | list[ContentBlock] = Field(default="") - reasoning_content: str = Field(default="") - tool_calls: list[ToolCall] = Field(default_factory=list) - tool_call_id: str = Field(default="") - time_created: str = Field(default_factory=lambda: datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")) - metadata: dict = Field(default_factory=dict) - - def dump_content(self) -> str | list[dict]: - """Returns content as a raw string or a list of serialized blocks.""" - if isinstance(self.content, str): - return self.content - return [block.simple_dump() for block in self.content] - - def simple_dump( - self, - add_name: bool = False, - add_reasoning: bool = True, - add_time_created: bool = False, - add_metadata: bool = False, - enable_json_dump: bool = False, - ) -> dict | str: - """Transforms the message into a simplified dictionary for standard APIs.""" - result = {} - if add_name and self.name: - result["name"] = self.name - - result["role"] = self.role.value - result["content"] = self.dump_content() - - if add_reasoning and self.reasoning_content: - result["reasoning_content"] = self.reasoning_content - - if self.tool_calls: - result["tool_calls"] = [tc.simple_output_dump() for tc in self.tool_calls] - - if self.tool_call_id: - result["tool_call_id"] = self.tool_call_id - - if add_time_created: - result["time_created"] = self.time_created - - if add_metadata: - result["metadata"] = self.metadata - - if enable_json_dump: - return json.dumps(result, ensure_ascii=False) - else: - return result - - def format_message( - self, - index: int | None = None, - add_time: bool = False, - use_name: bool = False, - add_reasoning: bool = True, - add_tools: bool = True, - strip_markdown_headers: bool = False, - ) -> str: - """Generates a human-readable string representation of the message.""" - prefix = f"round{index} " if index is not None else "" - time_str = f"[{self.time_created}] " if add_time else "" - header = f"{self.name or self.role.value if use_name else self.role.value}:" - - lines = [f"{prefix}{time_str}{header}"] - - def strip_md_func(line): - if strip_markdown_headers: - line = re.sub(r'\n##+ +', '\n', line) - return line - - if add_reasoning and self.reasoning_content: - lines.append(self.reasoning_content) - - if isinstance(self.content, str): - lines.append(strip_md_func(self.content)) - - 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 = str(text) - lines.append(strip_md_func(text)) - - if add_tools and self.tool_calls: - for tc in self.tool_calls: - lines.append(f" - tool_call={tc.name} params={tc.arguments}") - - return " ".join(lines).strip() - - -class Trajectory(BaseModel): - """Sequence of messages representing a full conversation session and its evaluation.""" - - task_id: str = Field(default="") - messages: list[Message] = Field(default_factory=list) - score: float = Field(default=0.0) - metadata: dict = Field(default_factory=dict) diff --git a/reme_ai/core/schema/request.py b/reme_ai/core/schema/request.py deleted file mode 100644 index ece942b9..00000000 --- a/reme_ai/core/schema/request.py +++ /dev/null @@ -1,11 +0,0 @@ -"""Defines the data structure for processing incoming user requests and message history.""" - -from pydantic import Field, BaseModel, ConfigDict - - -class Request(BaseModel): - """Represents a structured request payload containing a query, message list, and metadata.""" - - model_config = ConfigDict(extra="allow") - - metadata: dict = Field(default_factory=dict) diff --git a/reme_ai/core/schema/response.py b/reme_ai/core/schema/response.py deleted file mode 100644 index 3104bc6e..00000000 --- a/reme_ai/core/schema/response.py +++ /dev/null @@ -1,11 +0,0 @@ -"""Defines the standardized data structure for model output responses.""" - -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="") - success: bool = Field(default=True) - metadata: dict = Field(default_factory=dict) diff --git a/reme_ai/core/schema/service_config.py b/reme_ai/core/schema/service_config.py deleted file mode 100644 index 4c6eb543..00000000 --- a/reme_ai/core/schema/service_config.py +++ /dev/null @@ -1,113 +0,0 @@ -"""Configuration schemas for service components using Pydantic models.""" - -import os -from typing import Dict, List - -from pydantic import BaseModel, Field, ConfigDict - -from .tool_call import ToolCall - - -class MCPConfig(BaseModel): - """Configuration for Model Context Protocol transport and network settings.""" - - model_config = ConfigDict(extra="allow") - - transport: str = Field(default="stdio") - host: str = Field(default="0.0.0.0") - port: int = Field(default=8001) - - -class HttpConfig(BaseModel): - """Configuration for the HTTP server interface and connection lifecycle.""" - - model_config = ConfigDict(extra="allow") - - host: str = Field(default="0.0.0.0") - port: int = Field(default=8001) - timeout_keep_alive: int = Field(default=3600) - limit_concurrency: int = Field(default=1000) - - -class CmdConfig(BaseModel): - """Configuration for command-line flow execution parameters.""" - - model_config = ConfigDict(extra="allow") - - flow: str = Field(default="") - - -class FlowConfig(ToolCall): - """Configuration for workflow execution, caching, and error handling.""" - - model_config = ConfigDict(extra="allow") - - flow_content: str = Field(default="") - stream: bool = Field(default=False) - raise_exception: bool = Field(default=True) - enable_cache: bool = Field(default=False) - cache_path: str = Field(default="cache/flow") - cache_expire_hours: float = Field(default=0.1) - - -class LLMConfig(BaseModel): - """Configuration for Large Language Model backend and model identification.""" - - model_config = ConfigDict(extra="allow") - - backend: str = Field(default="") - model_name: str = Field(default="") - - -class EmbeddingModelConfig(BaseModel): - """Configuration for embedding model backends and identity.""" - - model_config = ConfigDict(extra="allow") - - backend: str = Field(default="") - model_name: str = Field(default="") - - -class VectorStoreConfig(BaseModel): - """Configuration for vector database storage and associated embeddings.""" - - model_config = ConfigDict(extra="allow") - - backend: str = Field(default="local") - collection_name: str = Field(default="reme") - embedding_model: str = Field(default="default") - - -class TokenCounterConfig(BaseModel): - """Configuration for token counting services and model mapping.""" - - model_config = ConfigDict(extra="allow") - - backend: str = Field(default="base") - model_name: str = Field(default="") - - -class ServiceConfig(BaseModel): - """Root configuration schema aggregating all service-level settings and components.""" - - model_config = ConfigDict(extra="allow") - - backend: str = Field(default="") - app_name: str = Field(default=os.getenv("APP_NAME", "ReMe")) - enable_logo: bool = Field(default=True) - language: str = Field(default="") - 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") - - 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) diff --git a/reme_ai/core/schema/stream_chunk.py b/reme_ai/core/schema/stream_chunk.py deleted file mode 100644 index 764981fd..00000000 --- a/reme_ai/core/schema/stream_chunk.py +++ /dev/null @@ -1,14 +0,0 @@ -"""Defines the data structure for individual data packets in a streaming response.""" - -from pydantic import Field, BaseModel - -from ..enumeration import ChunkEnum - - -class StreamChunk(BaseModel): - """Represents a single chunk of streamed data including its type, content, and completion status.""" - - chunk_type: ChunkEnum = Field(default=ChunkEnum.ANSWER) - chunk: str | dict | list = Field(default="") - done: bool = Field(default=False) - metadata: dict = Field(default_factory=dict) diff --git a/reme_ai/core/schema/tool_call.py b/reme_ai/core/schema/tool_call.py deleted file mode 100644 index 3c01cdcd..00000000 --- a/reme_ai/core/schema/tool_call.py +++ /dev/null @@ -1,228 +0,0 @@ -""" -MCP Tool Schema definitions for recursive JSON Schema representation. -""" - -import json -from typing import Any, Dict, List, Optional, Union - -from mcp.types import Tool -from pydantic import BaseModel, ConfigDict, Field, model_validator, field_validator - -from ..enumeration.json_schema_enum import JsonSchemaEnum - - -class ToolAttr(BaseModel): - """Recursive model representing JSON Schema attributes for tool parameters.""" - - model_config = ConfigDict(extra="allow") - - type: str = Field(default=str(JsonSchemaEnum.STRING), description="The data type of the attribute") - description: Optional[str] = Field(default=None, description="Description of the attribute") - required: Optional[List[str]] = Field(default=None, description="Required property names for object types") - properties: Optional[Dict[str, "ToolAttr"]] = Field(default=None, description="Child properties for objects") - items: Optional[Union[Dict[str, Any], "ToolAttr"]] = Field(default=None, description="Schema for array items") - enum: Optional[List[str]] = Field(default=None, description="Allowed values for the attribute") - - @field_validator("type") - @classmethod - def validate_type_is_valid_enum(cls, v: str) -> str: - """Validates that the provided type string exists within JsonSchemaEnum values.""" - valid_types = [str(e) for e in JsonSchemaEnum] - - if v not in valid_types: - raise ValueError(f"Invalid type: '{v}'. Must be one of {valid_types}") - return v - - def simple_input_dump(self) -> dict: - """Serializes the attribute into a standard JSON Schema dictionary.""" - res: dict = {"type": self.type} - if self.description: - res["description"] = self.description - if self.enum: - res["enum"] = self.enum - - 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 is not None: - res["required"] = self.required - - 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 - - -# Enable recursive type resolution -ToolAttr.model_rebuild() - - -class ToolCall(BaseModel): - """ - Model representing a tool definition and its call structure. - Supports parsing from standard JSON Schema formats and converting to MCP Tool objects. - input: - { - "type": "function", - "function": { - "name": "get_current_weather", - "description": "It is very useful when you want to check the weather of a specified city.", - "parameters": { - "type": "object", - "properties": { - "location": { - "type": "string", - "description": "Cities or counties, such as Beijing, Hangzhou, Yuhang District, etc.", - } - }, - "required": ["location"] - } - } - } - output: - { - "index": 0, - "id": "call_6596dafa2a6a46f7a217da", - "function": { - "arguments": "{\"location\": \"Beijing\"}", - "name": "get_current_weather" - }, - "type": "function", - } - """ - - index: int = 0 - id: str = "" - type: str = "function" - name: str = "" - description: str = "" - - arguments: str = Field(default="", description="JSON string of tool execution arguments") - - parameters: ToolAttr = Field( - default_factory=lambda: ToolAttr(type="object", properties={}, required=[]), - 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: - """Initializes the model by parsing tool-specific body data.""" - data = data.copy() - t_type = data.get("type", "function") - body = data.get(t_type, {}) - - # Extract basic metadata - data["name"] = body.get("name", data.get("name", "")) - data["arguments"] = body.get("arguments", data.get("arguments", "")) - data["description"] = body.get("description", data.get("description", "")) - - # Handle parameters mapping - if "parameters" in body: - params = body["parameters"] - # If parameters is already a dict, ensure it matches ToolAttr structure - if isinstance(params, dict): - data["parameters"] = ToolAttr(**params) - - # Handle output mapping (if provided in source) - if "output" in body and isinstance(body["output"], dict): - data["output"] = ToolAttr(**body["output"]) - - return data - - def simple_input_dump(self) -> dict: - """Returns a standardized tool definition dictionary.""" - return { - "type": self.type, - self.type: { - "name": self.name, - "description": self.description, - "parameters": self.parameters.simple_input_dump(), - }, - } - - @classmethod - def from_mcp_tool(cls, tool: Tool) -> "ToolCall": - """Creates a ToolCall instance from an MCP Tool object.""" - # MCP Tool inputSchema maps directly to our parameters ToolAttr - return cls( - name=tool.name, - description=tool.description or "", - parameters=ToolAttr(**tool.inputSchema), - ) - - def to_mcp_tool(self) -> Tool: - """Converts the instance back into an MCP Tool object.""" - return Tool( - name=self.name, - 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 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 - - 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_ai/core/schema/vector_node.py deleted file mode 100644 index 937ef4be..00000000 --- a/reme_ai/core/schema/vector_node.py +++ /dev/null @@ -1,15 +0,0 @@ -"""Defines the data structure for individual vector embedding nodes within a retrieval system.""" - -from typing import List, Dict -from uuid import uuid4 - -from pydantic import BaseModel, Field - - -class VectorNode(BaseModel): - """Represents a discrete unit of text content paired with its corresponding vector embedding and metadata.""" - - vector_id: str = Field(default_factory=lambda: uuid4().hex) - content: str = Field(default="") - vector: List[float] | None = Field(default=None) - metadata: Dict[str, str | bool | int | float] = Field(default_factory=dict) diff --git a/reme_ai/core/utils/__init__.py b/reme_ai/core/utils/__init__.py deleted file mode 100644 index 23396f97..00000000 --- a/reme_ai/core/utils/__init__.py +++ /dev/null @@ -1,47 +0,0 @@ -"""utils""" - -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 .http_client import HttpClient -from .llm_utils import extract_content, format_messages, deduplicate_memories -from .logger_utils import init_logger -from .logo_utils import print_logo - -# Make MCPClient import optional to avoid breaking if MCP dependencies are not available -try: - from .mcp_client import MCPClient - _HAS_MCP = True -except ImportError: - MCPClient = None - _HAS_MCP = False - -from .pydantic_config_parser import PydanticConfigParser -from .pydantic_utils import create_pydantic_model -from .singleton import singleton -from .time import timer, get_now_time - -__all__ = [ - "CacheHandler", - "snake_to_camel", - "camel_to_snake", - "run_coro_safely", - "execute_stream_task", - "load_env", - "exec_code", - "run_shell_command", - "HttpClient", - "extract_content", - "format_messages", - "deduplicate_memories", - "init_logger", - "print_logo", - "MCPClient", - "PydanticConfigParser", - "create_pydantic_model", - "singleton", - "timer", - "get_now_time", -] diff --git a/reme_ai/core/reme.py b/reme_ai/reme.py similarity index 99% rename from reme_ai/core/reme.py rename to reme_ai/reme.py index b5581159..6ec678c9 100644 --- a/reme_ai/core/reme.py +++ b/reme_ai/reme.py @@ -58,7 +58,6 @@ from .mem_tool.v4 import ( ) -@singleton class ReMe(Application): """Simplified ReMe application that auto-initializes the service context.""" From 78b993f8308515a67c2738acd6166bade3133383 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 22 Jan 2026 22:07:32 +0800 Subject: [PATCH 13/19] refactor(core): restructure project modules and update base classes --- docs/work_memory/message_offload_ops.md | 2 +- docs/work_memory/message_reload_ops.md | 2 +- pyproject.toml | 2 +- reme/__init__.py | 19 + reme/agent/__init__.py | 7 + .../mem_agent => reme/agent}/chat/__init__.py | 6 +- .../agent}/chat/simple_chat.py | 8 +- .../agent}/chat/stream_chat.py | 6 +- reme/core/__init__.py | 29 ++ reme/core/context/runtime_context.py | 3 +- reme/core/vector_store/chroma_vector_store.py | 3 + reme/core/vector_store/es_vector_store.py | 10 +- reme/core/vector_store/local_vector_store.py | 9 +- reme/core/vector_store/pgvector_store.py | 9 +- reme/core/vector_store/qdrant_vector_store.py | 14 +- reme/tool/__init__.py | 9 + {reme_ai => reme}/tool/execute/__init__.py | 4 + .../tool/execute/execute_code.py | 8 +- .../tool/execute/execute_code.yaml | 0 .../tool/execute/execute_shell.py | 8 +- .../tool/execute/execute_shell.yaml | 0 {reme_ai => reme}/tool/search/__init__.py | 5 + .../tool/search/dashscope_search.py | 11 +- .../tool/search/dashscope_search.yaml | 0 {reme_ai => reme}/tool/search/mock_search.py | 8 +- .../tool/search/mock_search.yaml | 0 .../tool/search/tavily_search.py | 14 +- .../tool/search/tavily_search.yaml | 0 reme/workflow/__init__.py | 0 reme_ai/mem_agent/chat/remy_agent.py | 51 --- reme_ai/mem_agent/chat/remy_agent.yaml | 36 -- {test_op => test}/test_agentic_retrieve_op.py | 0 {test_op => test}/test_message_compact_op.py | 0 {test_op => test}/test_message_compress_op.py | 0 {test_op => test}/test_message_offload_op.py | 0 test/test_op_composition.py | 325 ------------------ test/test_reme.py | 5 +- {test => tests}/mcp_servers_demo.json | 0 {test => tests}/test_base_context.py | 3 +- {test => tests}/test_cache_handler.py | 2 +- {test => tests}/test_embedding.py | 8 +- {test => tests}/test_embedding_sync.py | 6 +- {test => tests}/test_llm.py | 10 +- {test => tests}/test_llm_sync.py | 8 +- {test => tests}/test_logo.py | 4 +- {test => tests}/test_mcp_client.py | 2 +- {test => tests}/test_mcp_server.py | 4 +- .../test_memory_vector_conversion.py | 0 {test => tests}/test_message.py | 4 +- {test => tests}/test_timer.py | 2 +- {test => tests}/test_token_counter.py | 6 +- {test => tests}/test_tool.py | 64 ++-- {test => tests}/test_tool_call.py | 2 +- {test => tests}/test_vector_store.py | 39 ++- 54 files changed, 235 insertions(+), 542 deletions(-) create mode 100644 reme/agent/__init__.py rename {reme_ai/mem_agent => reme/agent}/chat/__init__.py (57%) rename {reme_ai/mem_agent => reme/agent}/chat/simple_chat.py (93%) rename {reme_ai/mem_agent => reme/agent}/chat/stream_chat.py (95%) create mode 100644 reme/tool/__init__.py rename {reme_ai => reme}/tool/execute/__init__.py (64%) rename {reme_ai => reme}/tool/execute/execute_code.py (87%) rename {reme_ai => reme}/tool/execute/execute_code.yaml (100%) rename {reme_ai => reme}/tool/execute/execute_shell.py (90%) rename {reme_ai => reme}/tool/execute/execute_shell.yaml (100%) rename {reme_ai => reme}/tool/search/__init__.py (65%) rename {reme_ai => reme}/tool/search/dashscope_search.py (92%) rename {reme_ai => reme}/tool/search/dashscope_search.yaml (100%) rename {reme_ai => reme}/tool/search/mock_search.py (91%) rename {reme_ai => reme}/tool/search/mock_search.yaml (100%) rename {reme_ai => reme}/tool/search/tavily_search.py (90%) rename {reme_ai => reme}/tool/search/tavily_search.yaml (100%) create mode 100644 reme/workflow/__init__.py delete mode 100644 reme_ai/mem_agent/chat/remy_agent.py delete mode 100644 reme_ai/mem_agent/chat/remy_agent.yaml rename {test_op => test}/test_agentic_retrieve_op.py (100%) rename {test_op => test}/test_message_compact_op.py (100%) rename {test_op => test}/test_message_compress_op.py (100%) rename {test_op => test}/test_message_offload_op.py (100%) delete mode 100644 test/test_op_composition.py rename {test => tests}/mcp_servers_demo.json (100%) rename {test => tests}/test_base_context.py (97%) rename {test => tests}/test_cache_handler.py (98%) rename {test => tests}/test_embedding.py (98%) rename {test => tests}/test_embedding_sync.py (98%) rename {test => tests}/test_llm.py (98%) rename {test => tests}/test_llm_sync.py (98%) rename {test => tests}/test_logo.py (61%) rename {test => tests}/test_mcp_client.py (99%) rename {test => tests}/test_mcp_server.py (97%) rename {test => tests}/test_memory_vector_conversion.py (100%) rename {test => tests}/test_message.py (98%) rename {test => tests}/test_timer.py (97%) rename {test => tests}/test_token_counter.py (99%) rename {test => tests}/test_tool.py (75%) rename {test => tests}/test_tool_call.py (99%) rename {test => tests}/test_vector_store.py (98%) 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 index e69de29b..32f34911 100644 --- a/reme/__init__.py +++ 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/core/__init__.py b/reme/core/__init__.py index e69de29b..251c2433 100644 --- a/reme/core/__init__.py +++ b/reme/core/__init__.py @@ -0,0 +1,29 @@ +"""Core""" + +from . import context +from . import embedding +from . import enumeration +from . import flow +from . import llm +from . import op +from . import schema +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/core/context/runtime_context.py b/reme/core/context/runtime_context.py index 44167ca0..d4c02050 100644 --- a/reme/core/context/runtime_context.py +++ b/reme/core/context/runtime_context.py @@ -30,7 +30,8 @@ class RuntimeContext(BaseContext): if context is None: return cls(**kwargs) else: - context.update(kwargs) + if kwargs: + context.update(kwargs) return context async def _enqueue(self, chunk: StreamChunk) -> None: diff --git a/reme/core/vector_store/chroma_vector_store.py b/reme/core/vector_store/chroma_vector_store.py index eaddd0e7..710ce73e 100644 --- a/reme/core/vector_store/chroma_vector_store.py +++ b/reme/core/vector_store/chroma_vector_store.py @@ -1,5 +1,6 @@ """ChromaDB vector store implementation for the ReMe framework.""" +from concurrent.futures import ThreadPoolExecutor from typing import Any from loguru import logger @@ -26,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, @@ -44,6 +46,7 @@ class ChromaVectorStore(BaseVectorStore): super().__init__( collection_name=collection_name, embedding_model=embedding_model, + thread_pool=thread_pool, **kwargs, ) diff --git a/reme/core/vector_store/es_vector_store.py b/reme/core/vector_store/es_vector_store.py index 49054afc..9a1019b8 100644 --- a/reme/core/vector_store/es_vector_store.py +++ b/reme/core/vector_store/es_vector_store.py @@ -4,6 +4,7 @@ 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 @@ -30,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, @@ -43,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. @@ -59,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( diff --git a/reme/core/vector_store/local_vector_store.py b/reme/core/vector_store/local_vector_store.py index 24e1977c..a86e226a 100644 --- a/reme/core/vector_store/local_vector_store.py +++ b/reme/core/vector_store/local_vector_store.py @@ -1,6 +1,7 @@ """Local file system vector store implementation for ReMe.""" import json +from concurrent.futures import ThreadPoolExecutor from pathlib import Path from loguru import logger @@ -17,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) diff --git a/reme/core/vector_store/pgvector_store.py b/reme/core/vector_store/pgvector_store.py index 7d1b3401..279975d0 100644 --- a/reme/core/vector_store/pgvector_store.py +++ b/reme/core/vector_store/pgvector_store.py @@ -2,6 +2,7 @@ import json import re +from concurrent.futures import ThreadPoolExecutor from typing import Any from loguru import logger @@ -47,6 +48,7 @@ class PGVectorStore(BaseVectorStore): self, collection_name: str, embedding_model: BaseEmbeddingModel, + thread_pool: ThreadPoolExecutor, host: str = "localhost", port: int = 5432, database: str = "postgres", @@ -68,7 +70,12 @@ class PGVectorStore(BaseVectorStore): # Validate collection name to prevent SQL injection self._validate_table_name(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.dsn = dsn self.host = host diff --git a/reme/core/vector_store/qdrant_vector_store.py b/reme/core/vector_store/qdrant_vector_store.py index 5ca74cd0..93ccee70 100644 --- a/reme/core/vector_store/qdrant_vector_store.py +++ b/reme/core/vector_store/qdrant_vector_store.py @@ -1,5 +1,6 @@ """Qdrant vector store implementation for the ReMe project.""" +from concurrent.futures import ThreadPoolExecutor from typing import Any from loguru import logger @@ -42,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, @@ -59,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. @@ -76,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, @@ -87,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 diff --git a/reme/tool/__init__.py b/reme/tool/__init__.py new file mode 100644 index 00000000..60082176 --- /dev/null +++ b/reme/tool/__init__.py @@ -0,0 +1,9 @@ +"""Tool""" + +from . import execute +from . import search + +__all__ = [ + "execute", + "search", +] diff --git a/reme_ai/tool/execute/__init__.py b/reme/tool/execute/__init__.py similarity index 64% rename from reme_ai/tool/execute/__init__.py rename to reme/tool/execute/__init__.py index f78a73ac..28e3ecf3 100644 --- a/reme_ai/tool/execute/__init__.py +++ b/reme/tool/execute/__init__.py @@ -2,8 +2,12 @@ from .execute_code import ExecuteCode from .execute_shell import ExecuteShell +from ...core import R __all__ = [ "ExecuteCode", "ExecuteShell", ] + +R.op.register()(ExecuteCode) +R.op.register()(ExecuteShell) diff --git a/reme_ai/tool/execute/execute_code.py b/reme/tool/execute/execute_code.py similarity index 87% rename from reme_ai/tool/execute/execute_code.py rename to reme/tool/execute/execute_code.py index ea259487..6cb77c85 100644 --- a/reme_ai/tool/execute/execute_code.py +++ b/reme/tool/execute/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/execute/execute_code.yaml similarity index 100% rename from reme_ai/tool/execute/execute_code.yaml rename to reme/tool/execute/execute_code.yaml diff --git a/reme_ai/tool/execute/execute_shell.py b/reme/tool/execute/execute_shell.py similarity index 90% rename from reme_ai/tool/execute/execute_shell.py rename to reme/tool/execute/execute_shell.py index 6e244ddb..de862602 100644 --- a/reme_ai/tool/execute/execute_shell.py +++ b/reme/tool/execute/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/execute/execute_shell.yaml similarity index 100% rename from reme_ai/tool/execute/execute_shell.yaml rename to reme/tool/execute/execute_shell.yaml diff --git a/reme_ai/tool/search/__init__.py b/reme/tool/search/__init__.py similarity index 65% rename from reme_ai/tool/search/__init__.py rename to reme/tool/search/__init__.py index 5a73dc2e..a9e87ef5 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", ] + +R.op.register()(DashscopeSearch) +R.op.register()(MockSearch) +R.op.register()(TavilySearch) 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_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/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/test/test_op_composition.py b/test/test_op_composition.py deleted file mode 100644 index 8d32b51a..00000000 --- a/test/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/test/test_reme.py b/test/test_reme.py index c9e6843e..3279da2a 100644 --- a/test/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/mcp_servers_demo.json b/tests/mcp_servers_demo.json similarity index 100% rename from test/mcp_servers_demo.json rename to tests/mcp_servers_demo.json diff --git a/test/test_base_context.py b/tests/test_base_context.py similarity index 97% rename from test/test_base_context.py rename to tests/test_base_context.py index 316a6796..61a355c0 100644 --- a/test/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/test/test_cache_handler.py b/tests/test_cache_handler.py similarity index 98% rename from test/test_cache_handler.py rename to tests/test_cache_handler.py index ddcac86f..f1facc0f 100644 --- a/test/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/test/test_embedding.py b/tests/test_embedding.py similarity index 98% rename from test/test_embedding.py rename to tests/test_embedding.py index d1c404f5..867037f4 100644 --- a/test/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/test/test_embedding_sync.py b/tests/test_embedding_sync.py similarity index 98% rename from test/test_embedding_sync.py rename to tests/test_embedding_sync.py index 361a42b3..34ddd350 100644 --- a/test/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/test/test_llm.py b/tests/test_llm.py similarity index 98% rename from test/test_llm.py rename to tests/test_llm.py index 12c6eca3..5e1755fe 100644 --- a/test/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/test/test_llm_sync.py b/tests/test_llm_sync.py similarity index 98% rename from test/test_llm_sync.py rename to tests/test_llm_sync.py index 98751f87..0b3c8cce 100644 --- a/test/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/test/test_logo.py b/tests/test_logo.py similarity index 61% rename from test/test_logo.py rename to tests/test_logo.py index eeede81e..12b757c4 100644 --- a/test/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/test/test_mcp_client.py b/tests/test_mcp_client.py similarity index 99% rename from test/test_mcp_client.py rename to tests/test_mcp_client.py index d2fef40e..a369d7cf 100644 --- a/test/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/test/test_mcp_server.py b/tests/test_mcp_server.py similarity index 97% rename from test/test_mcp_server.py rename to tests/test_mcp_server.py index 67f9c542..f5744f99 100644 --- a/test/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/test/test_memory_vector_conversion.py b/tests/test_memory_vector_conversion.py similarity index 100% rename from test/test_memory_vector_conversion.py rename to tests/test_memory_vector_conversion.py diff --git a/test/test_message.py b/tests/test_message.py similarity index 98% rename from test/test_message.py rename to tests/test_message.py index 141174c5..7b7e44bc 100644 --- a/test/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/test/test_timer.py b/tests/test_timer.py similarity index 97% rename from test/test_timer.py rename to tests/test_timer.py index c9714e38..98b1d816 100644 --- a/test/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/test/test_token_counter.py b/tests/test_token_counter.py similarity index 99% rename from test/test_token_counter.py rename to tests/test_token_counter.py index 3c44a298..2fd8bdfd 100644 --- a/test/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/test/test_tool.py b/tests/test_tool.py similarity index 75% rename from test/test_tool.py rename to tests/test_tool.py index 9db5a3ed..02346d0a 100644 --- a/test/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.execute 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/test/test_tool_call.py b/tests/test_tool_call.py similarity index 99% rename from test/test_tool_call.py rename to tests/test_tool_call.py index 30c0a37e..7be3079e 100644 --- a/test/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/test/test_vector_store.py b/tests/test_vector_store.py similarity index 98% rename from test/test_vector_store.py rename to tests/test_vector_store.py index 00927d0c..14ed1c1d 100644 --- a/test/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, @@ -1267,7 +1278,7 @@ async def test_range_query_filters(store: BaseVectorStore, _store_name: str): for r in results_3: ts = r.metadata.get("timestamp") category = r.metadata.get("category") - assert ts >= start_time and ts <= end_time, "Timestamp should be in range" + assert 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}") @@ -1290,8 +1301,8 @@ async def test_range_query_filters(store: BaseVectorStore, _store_name: str): for r in results_4: ts = r.metadata.get("timestamp") rating = r.metadata.get("rating") - assert ts >= base_timestamp + 8000 and ts <= base_timestamp + 12000, "Timestamp out of range" - assert rating >= 65 and rating <= 75, f"Rating {rating} out of range [65, 75]" + 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 @@ -1311,7 +1322,7 @@ async def test_range_query_filters(store: BaseVectorStore, _store_name: str): # Verify rating range in list results for r in results_5: rating = r.metadata.get("rating") - assert rating >= 60 and rating <= 70, f"Rating {rating} should be in range [60, 70]" + assert 60 <= rating <= 70, f"Rating {rating} should be in range [60, 70]" logger.info("✓ Range query in list operation validated") @@ -1349,7 +1360,7 @@ async def test_range_query_filters(store: BaseVectorStore, _store_name: str): rating1 = results_7[i].metadata.get("rating") rating2 = results_7[i + 1].metadata.get("rating") assert rating1 >= rating2, f"Results not sorted: {rating1} < {rating2}" - assert rating1 >= 60 and rating1 <= 80, "Rating out of range" + assert 60 <= rating1 <= 80, "Rating out of range" logger.info("✓ Range query with sorting validated") @@ -1419,7 +1430,6 @@ async def test_string_range_queries(store: BaseVectorStore, store_name: str): logger.info(f"Test 1 - String date range ['2024-02-01', '2024-03-15']: {len(results)} results") # Verify all results are within range - expected_dates = ["2024-02-01", "2024-02-15", "2024-03-01", "2024-03-15"] for r in results: date = r.metadata.get("date") assert date >= "2024-02-01", f"Date {date} should be >= '2024-02-01'" @@ -1453,16 +1463,15 @@ async def test_sql_injection_protection(store: BaseVectorStore, store_name: str) # Test 1: Invalid collection name (SQL injection attempt) try: - from reme_ai.core.vector_store import PGVectorStore - from reme_ai.core.embedding import OpenAIEmbeddingModel - embedding_model = OpenAIEmbeddingModel() + thread_pool = ThreadPoolExecutor(max_workers=4) # This should raise ValueError due to invalid table name try: - invalid_store = PGVectorStore( + _ = 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" @@ -1471,7 +1480,7 @@ async def test_sql_injection_protection(store: BaseVectorStore, store_name: str) # Test 2: Invalid metadata key in filters try: - results = await store.search( + _ = await store.search( query="test", filters={ "normal_key": "value", From 3f0d45c51e951518526e5d6ee121b2f1e3a7e1d5 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 22 Jan 2026 22:53:11 +0800 Subject: [PATCH 14/19] refactor(tool): restructure tool modules and update base classes --- reme/core/op/base_tool.py | 5 +- reme/tool/__init__.py | 4 +- reme/tool/{execute => gallery}/__init__.py | 3 + .../tool/{execute => gallery}/execute_code.py | 0 .../{execute => gallery}/execute_code.yaml | 0 .../{execute => gallery}/execute_shell.py | 0 .../{execute => gallery}/execute_shell.yaml | 0 .../tool/gallery}/think_tool.py | 28 +--- .../tool/gallery}/think_tool.yaml | 0 reme/tool/memory/__init__.py | 0 reme/tool/memory/base_memory_tool.py | 91 +++++++++++++ reme_ai/mem_tool/base_memory_tool.py | 128 ------------------ reme_ai/tool/__init__.py | 9 -- tests/test_tool.py | 2 +- 14 files changed, 105 insertions(+), 165 deletions(-) rename reme/tool/{execute => gallery}/__init__.py (75%) rename reme/tool/{execute => gallery}/execute_code.py (100%) rename reme/tool/{execute => gallery}/execute_code.yaml (100%) rename reme/tool/{execute => gallery}/execute_shell.py (100%) rename reme/tool/{execute => gallery}/execute_shell.yaml (100%) rename {reme_ai/mem_tool => reme/tool/gallery}/think_tool.py (60%) rename {reme_ai/mem_tool => reme/tool/gallery}/think_tool.yaml (100%) create mode 100644 reme/tool/memory/__init__.py create mode 100644 reme/tool/memory/base_memory_tool.py delete mode 100644 reme_ai/mem_tool/base_memory_tool.py delete mode 100644 reme_ai/tool/__init__.py diff --git a/reme/core/op/base_tool.py b/reme/core/op/base_tool.py index 9daf033a..5b41b1fb 100644 --- a/reme/core/op/base_tool.py +++ b/reme/core/op/base_tool.py @@ -25,13 +25,10 @@ class BaseTool(BaseOp, metaclass=ABCMeta): self.context.validate_required_keys(required_keys, self.name) @property - def tool_call(self) -> ToolCall | None: + def tool_call(self) -> ToolCall: """Get the tool call schema.""" 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 return self._tool_call diff --git a/reme/tool/__init__.py b/reme/tool/__init__.py index 60082176..f39b9b04 100644 --- a/reme/tool/__init__.py +++ b/reme/tool/__init__.py @@ -1,9 +1,9 @@ """Tool""" -from . import execute +from . import gallery from . import search __all__ = [ - "execute", + "gallery", "search", ] diff --git a/reme/tool/execute/__init__.py b/reme/tool/gallery/__init__.py similarity index 75% rename from reme/tool/execute/__init__.py rename to reme/tool/gallery/__init__.py index 28e3ecf3..f7bc08e6 100644 --- a/reme/tool/execute/__init__.py +++ b/reme/tool/gallery/__init__.py @@ -2,12 +2,15 @@ from .execute_code import ExecuteCode from .execute_shell import ExecuteShell +from .think_tool import ThinkTool from ...core import R __all__ = [ "ExecuteCode", "ExecuteShell", + "ThinkTool", ] R.op.register()(ExecuteCode) R.op.register()(ExecuteShell) +R.op.register()(ThinkTool) diff --git a/reme/tool/execute/execute_code.py b/reme/tool/gallery/execute_code.py similarity index 100% rename from reme/tool/execute/execute_code.py rename to reme/tool/gallery/execute_code.py diff --git a/reme/tool/execute/execute_code.yaml b/reme/tool/gallery/execute_code.yaml similarity index 100% rename from reme/tool/execute/execute_code.yaml rename to reme/tool/gallery/execute_code.yaml diff --git a/reme/tool/execute/execute_shell.py b/reme/tool/gallery/execute_shell.py similarity index 100% rename from reme/tool/execute/execute_shell.py rename to reme/tool/gallery/execute_shell.py diff --git a/reme/tool/execute/execute_shell.yaml b/reme/tool/gallery/execute_shell.yaml similarity index 100% rename from reme/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..e69de29b diff --git a/reme/tool/memory/base_memory_tool.py b/reme/tool/memory/base_memory_tool.py new file mode 100644 index 00000000..aaef5542 --- /dev/null +++ b/reme/tool/memory/base_memory_tool.py @@ -0,0 +1,91 @@ +"""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 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/base_memory_tool.py b/reme_ai/mem_tool/base_memory_tool.py deleted file mode 100644 index 8b124496..00000000 --- a/reme_ai/mem_tool/base_memory_tool.py +++ /dev/null @@ -1,128 +0,0 @@ -"""Base class for memory tool""" - -from abc import ABCMeta -from pathlib import Path - -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 description(self) -> str: - """Get the description from context.""" - return self.context.get("description", "") - - @property - def messages_formated(self) -> str: - """Get the formated messages from context.""" - return self.context.get("messages_formated", "") - - @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", "") - - 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 {}, - ) - return node 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/tests/test_tool.py b/tests/test_tool.py index 02346d0a..8d1fb232 100644 --- a/tests/test_tool.py +++ b/tests/test_tool.py @@ -43,7 +43,7 @@ def test_execute(): including successful execution, syntax errors, runtime errors, and invalid commands to verify error handling. """ - from reme.tool.execute import ExecuteCode, ExecuteShell + from reme.tool.gallery import ExecuteCode, ExecuteShell # Test ExecuteCode print("\n" + "=" * 60) From 30990308d866fc7628cfa6d3a590d30d05112d7b Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 22 Jan 2026 23:46:11 +0800 Subject: [PATCH 15/19] refactor(memory): migrate and optimize user profile memory tools --- reme/reme_app.py | 2 +- reme/tool/memory/base_memory_tool.py | 5 + reme/tool/memory/read_user_profile.py | 60 ++++++++++++ reme/tool/memory/update_user_profile.py | 109 +++++++++++++++++++++ reme_ai/mem_tool/v4/read_user_profile.py | 81 --------------- reme_ai/mem_tool/v4/update_user_profile.py | 105 -------------------- 6 files changed, 175 insertions(+), 187 deletions(-) create mode 100644 reme/tool/memory/read_user_profile.py create mode 100644 reme/tool/memory/update_user_profile.py delete mode 100644 reme_ai/mem_tool/v4/read_user_profile.py delete mode 100644 reme_ai/mem_tool/v4/update_user_profile.py diff --git a/reme/reme_app.py b/reme/reme_app.py index 483cd5bb..e41aa8a6 100644 --- a/reme/reme_app.py +++ b/reme/reme_app.py @@ -3,11 +3,11 @@ import asyncio import sys -from reme.core.utils import execute_stream_task 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: diff --git a/reme/tool/memory/base_memory_tool.py b/reme/tool/memory/base_memory_tool.py index aaef5542..adc2f6b2 100644 --- a/reme/tool/memory/base_memory_tool.py +++ b/reme/tool/memory/base_memory_tool.py @@ -75,6 +75,11 @@ class BaseMemoryTool(BaseTool, metaclass=ABCMeta): """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.""" diff --git a/reme/tool/memory/read_user_profile.py b/reme/tool/memory/read_user_profile.py new file mode 100644 index 00000000..6b3118cf --- /dev/null +++ b/reme/tool/memory/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/update_user_profile.py b/reme/tool/memory/update_user_profile.py new file mode 100644 index 00000000..e2032d45 --- /dev/null +++ b/reme/tool/memory/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/mem_tool/v4/read_user_profile.py b/reme_ai/mem_tool/v4/read_user_profile.py deleted file mode 100644 index ff963ea2..00000000 --- a/reme_ai/mem_tool/v4/read_user_profile.py +++ /dev/null @@ -1,81 +0,0 @@ -from typing import Literal -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 = False, show_ids: Literal["both", "profile", "history", "none"] = "both", **kwargs): - kwargs["enable_multiple"] = False - super().__init__(**kwargs) - self.add_memory_type_target = add_memory_type_target - self.show_ids = show_ids - - def _build_tool_description(self) -> str: - return "Read user profile." - - 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): - # Determine which IDs to show - show_profile_id = self.show_ids in ("both", "profile") - show_history_id = self.show_ids in ("both", "history") - - cache_key = f"{self.memory_type}_{self.memory_target}".replace(" ", "_").lower() - cached_data = self.meta_memory.load(cache_key, auto_clean=False) - - if not cached_data: - self.output = "### User Profile\nNo user profile found." - logger.info(f"empty cached_data={cache_key}") - return - - 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_parts = [] - - # Add profile_id if enabled - if show_profile_id: - node_formated_parts.append(f"profile_id={node.memory_id}") - - # Always add profile_content - node_formated_parts.append(f"profile_content={node.content}") - - # Add conversation_time if available - if "conversation_time" in node.metadata and node.metadata["conversation_time"]: - node_formated_parts.append(f"conversation_time={node.metadata['conversation_time']}") - - # Add history_id if enabled and available - if show_history_id and node.ref_memory_id: - node_formated_parts.append(f"history_id={node.ref_memory_id}") - - node_formated = " ".join(node_formated_parts) - memory_formated.append(node_formated.strip()) - - self.output = "### User Profile\n" + "\n".join(memory_formated) - logger.info(f"Read {len(memory_formated)} nodes from cache key: {cache_key}") diff --git a/reme_ai/mem_tool/v4/update_user_profile.py b/reme_ai/mem_tool/v4/update_user_profile.py deleted file mode 100644 index a8fa04f5..00000000 --- a/reme_ai/mem_tool/v4/update_user_profile.py +++ /dev/null @@ -1,105 +0,0 @@ -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ...core.schema.memory_node import MemoryNode -from ...core.utils import deduplicate_memories - - -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": { - "conversation_time": { - "type": "string", - "description": "conversation_time, e.g. '2020-01-01 00:00:00'", - }, - "profile_content": { - "type": "string", - "description": "profile_content", - }, - }, - "required": ["conversation_time", "profile_content"], - }, - }, - }, - "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 profiles to remove or add. Operation has been done." - return - - cache_key = f"{self.memory_type}_{self.memory_target}".replace(" ", "_").lower() - cached_data = self.meta_memory.load(cache_key, auto_clean=False) - if cached_data: - existing_memory_nodes = [MemoryNode(**node_data) for node_data in cached_data] - else: - existing_memory_nodes = [] - - removed_count = 0 - if profile_ids_to_delete: - original_count = len(existing_memory_nodes) - existing_memory_nodes = [n for n in existing_memory_nodes if n.memory_id not in profile_ids_to_delete] - removed_count = original_count - len(existing_memory_nodes) - logger.info(f"Removed {removed_count} profiles.") - - added_count = 0 - new_memory_nodes = [] - if profiles_to_add: - for mem in profiles_to_add: - memory_node = MemoryNode( - memory_type=self.memory_type, - memory_target=self.memory_target, - when_to_use="", - content=mem.get("profile_content", ""), - ref_memory_id=self.history_node.memory_id, - author=self.author, - metadata={"conversation_time": mem.get("conversation_time", "")}, - ) - new_memory_nodes.append(memory_node) - added_count = len(new_memory_nodes) - logger.info(f"Added {added_count} new profiles.") - - updated_memory_nodes = deduplicate_memories(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 profiles") - if added_count > 0: - operations.append(f"added {added_count} new profiles") - - if operations: - self.output = f"Successfully {' and '.join(operations)} in user profile." - else: - self.output = "Operation has been done." - logger.info(self.output) From af44052a5705f98086833ff07dfb427f0e59cda4 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 22 Jan 2026 23:48:13 +0800 Subject: [PATCH 16/19] feat(memory): add memory tools module with user profile operations --- reme/tool/__init__.py | 2 ++ reme/tool/memory/__init__.py | 15 +++++++++++++++ 2 files changed, 17 insertions(+) diff --git a/reme/tool/__init__.py b/reme/tool/__init__.py index f39b9b04..5c0d019b 100644 --- a/reme/tool/__init__.py +++ b/reme/tool/__init__.py @@ -1,9 +1,11 @@ """Tool""" from . import gallery +from . import memory from . import search __all__ = [ "gallery", + "memory", "search", ] diff --git a/reme/tool/memory/__init__.py b/reme/tool/memory/__init__.py index e69de29b..2585e3d4 100644 --- a/reme/tool/memory/__init__.py +++ b/reme/tool/memory/__init__.py @@ -0,0 +1,15 @@ +"""memory tools""" + +from .base_memory_tool import BaseMemoryTool +from .read_user_profile import ReadUserProfile +from .update_user_profile import UpdateUserProfile +from ...core import R + +__all__ = [ + "BaseMemoryTool", + "ReadUserProfile", + "UpdateUserProfile", +] + +R.op.register()(ReadUserProfile) +R.op.register()(UpdateUserProfile) From 3c272ac859ba45bc94c1d5186ebf3ef6bb8e975d Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 23 Jan 2026 00:11:00 +0800 Subject: [PATCH 17/19] feat(memory): add history management tools with dynamic registration --- reme/tool/gallery/__init__.py | 6 ++-- reme/tool/memory/__init__.py | 9 ++++-- reme/tool/memory/add_history.py | 47 +++++++++++++++++++++++++++++ reme/tool/memory/read_history.py | 46 ++++++++++++++++++++++++++++ reme/tool/search/__init__.py | 6 ++-- reme_ai/mem_tool/v4/read_history.py | 38 ----------------------- 6 files changed, 106 insertions(+), 46 deletions(-) create mode 100644 reme/tool/memory/add_history.py create mode 100644 reme/tool/memory/read_history.py delete mode 100644 reme_ai/mem_tool/v4/read_history.py diff --git a/reme/tool/gallery/__init__.py b/reme/tool/gallery/__init__.py index f7bc08e6..30e1d931 100644 --- a/reme/tool/gallery/__init__.py +++ b/reme/tool/gallery/__init__.py @@ -11,6 +11,6 @@ __all__ = [ "ThinkTool", ] -R.op.register()(ExecuteCode) -R.op.register()(ExecuteShell) -R.op.register()(ThinkTool) +for name in __all__: + tool_class = globals()[name] + R.op.register()(tool_class) diff --git a/reme/tool/memory/__init__.py b/reme/tool/memory/__init__.py index 2585e3d4..2b24f341 100644 --- a/reme/tool/memory/__init__.py +++ b/reme/tool/memory/__init__.py @@ -1,15 +1,20 @@ """memory tools""" +from .add_history import AddHistory from .base_memory_tool import BaseMemoryTool +from .read_history import ReadHistory from .read_user_profile import ReadUserProfile from .update_user_profile import UpdateUserProfile from ...core import R __all__ = [ + "AddHistory", "BaseMemoryTool", + "ReadHistory", "ReadUserProfile", "UpdateUserProfile", ] -R.op.register()(ReadUserProfile) -R.op.register()(UpdateUserProfile) +for name in __all__: + tool_class = globals()[name] + R.op.register()(tool_class) diff --git a/reme/tool/memory/add_history.py b/reme/tool/memory/add_history.py new file mode 100644 index 00000000..c61a130a --- /dev/null +++ b/reme/tool/memory/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/read_history.py b/reme/tool/memory/read_history.py new file mode 100644 index 00000000..bf43819e --- /dev/null +++ b/reme/tool/memory/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/search/__init__.py b/reme/tool/search/__init__.py index a9e87ef5..6230c7a2 100644 --- a/reme/tool/search/__init__.py +++ b/reme/tool/search/__init__.py @@ -11,6 +11,6 @@ __all__ = [ "TavilySearch", ] -R.op.register()(DashscopeSearch) -R.op.register()(MockSearch) -R.op.register()(TavilySearch) +for name in __all__: + tool_class = globals()[name] + R.op.register()(tool_class) diff --git a/reme_ai/mem_tool/v4/read_history.py b/reme_ai/mem_tool/v4/read_history.py deleted file mode 100644 index 78a90eb4..00000000 --- a/reme_ai/mem_tool/v4/read_history.py +++ /dev/null @@ -1,38 +0,0 @@ -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 = f"### Historical Dialogue\n{memory.content}" - logger.info(f"Successfully read history memory: {history_id}") From c927d264e1fc2292f76c7829b84fb586fd6dc6d4 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 23 Jan 2026 00:33:16 +0800 Subject: [PATCH 18/19] refactor(memory): restructure memory tools with new identity and meta memory features --- reme/tool/memory/__init__.py | 18 ++- .../tool/memory}/history/__init__.py | 0 reme/tool/memory/{ => history}/add_history.py | 10 +- .../tool/memory/{ => history}/read_history.py | 4 +- .../tool/memory}/identity/__init__.py | 0 reme/tool/memory/identity/add_identity.py | 42 ++++++ reme/tool/memory/identity/read_identity.py | 36 ++++++ .../tool/memory}/meta/__init__.py | 0 reme/tool/memory/meta/add_meta_memory.py | 89 +++++++++++++ reme/tool/memory/meta/read_meta_memory.py | 66 ++++++++++ reme/tool/memory/user_profile/__init__.py | 0 .../{ => user_profile}/read_user_profile.py | 6 +- .../{ => user_profile}/update_user_profile.py | 8 +- .../mem_tool/history/add_history_memory.py | 50 -------- .../mem_tool/history/add_history_memory.yaml | 14 -- .../mem_tool/history/read_history_memory.py | 65 ---------- .../mem_tool/history/read_history_memory.yaml | 11 -- .../mem_tool/identity/read_identity_memory.py | 27 ---- .../identity/read_identity_memory.yaml | 3 - .../identity/update_identity_memory.py | 39 ------ .../identity/update_identity_memory.yaml | 7 - reme_ai/mem_tool/meta/add_meta_memory.py | 121 ------------------ reme_ai/mem_tool/meta/add_meta_memory.yaml | 26 ---- reme_ai/mem_tool/meta/read_meta_memory.py | 97 -------------- reme_ai/mem_tool/meta/read_meta_memory.yaml | 16 --- 25 files changed, 260 insertions(+), 495 deletions(-) rename {reme_ai/mem_tool => reme/tool/memory}/history/__init__.py (100%) rename reme/tool/memory/{ => history}/add_history.py (87%) rename reme/tool/memory/{ => history}/read_history.py (93%) rename {reme_ai/mem_tool => reme/tool/memory}/identity/__init__.py (100%) create mode 100644 reme/tool/memory/identity/add_identity.py create mode 100644 reme/tool/memory/identity/read_identity.py rename {reme_ai/mem_tool => reme/tool/memory}/meta/__init__.py (100%) create mode 100644 reme/tool/memory/meta/add_meta_memory.py create mode 100644 reme/tool/memory/meta/read_meta_memory.py create mode 100644 reme/tool/memory/user_profile/__init__.py rename reme/tool/memory/{ => user_profile}/read_user_profile.py (93%) rename reme/tool/memory/{ => user_profile}/update_user_profile.py (96%) delete mode 100644 reme_ai/mem_tool/history/add_history_memory.py delete mode 100644 reme_ai/mem_tool/history/add_history_memory.yaml delete mode 100644 reme_ai/mem_tool/history/read_history_memory.py delete mode 100644 reme_ai/mem_tool/history/read_history_memory.yaml delete mode 100644 reme_ai/mem_tool/identity/read_identity_memory.py delete mode 100644 reme_ai/mem_tool/identity/read_identity_memory.yaml delete mode 100644 reme_ai/mem_tool/identity/update_identity_memory.py delete mode 100644 reme_ai/mem_tool/identity/update_identity_memory.yaml delete mode 100644 reme_ai/mem_tool/meta/add_meta_memory.py delete mode 100644 reme_ai/mem_tool/meta/add_meta_memory.yaml delete mode 100644 reme_ai/mem_tool/meta/read_meta_memory.py delete mode 100644 reme_ai/mem_tool/meta/read_meta_memory.yaml diff --git a/reme/tool/memory/__init__.py b/reme/tool/memory/__init__.py index 2b24f341..9113f35a 100644 --- a/reme/tool/memory/__init__.py +++ b/reme/tool/memory/__init__.py @@ -1,16 +1,24 @@ """memory tools""" -from .add_history import AddHistory from .base_memory_tool import BaseMemoryTool -from .read_history import ReadHistory -from .read_user_profile import ReadUserProfile -from .update_user_profile import UpdateUserProfile +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__ = [ - "AddHistory", "BaseMemoryTool", + "AddHistory", "ReadHistory", + "AddIdentity", + "ReadIdentity", + "AddMetaMemory", + "ReadMetaMemory", "ReadUserProfile", "UpdateUserProfile", ] diff --git a/reme_ai/mem_tool/history/__init__.py b/reme/tool/memory/history/__init__.py similarity index 100% rename from reme_ai/mem_tool/history/__init__.py rename to reme/tool/memory/history/__init__.py diff --git a/reme/tool/memory/add_history.py b/reme/tool/memory/history/add_history.py similarity index 87% rename from reme/tool/memory/add_history.py rename to reme/tool/memory/history/add_history.py index c61a130a..b46259f3 100644 --- a/reme/tool/memory/add_history.py +++ b/reme/tool/memory/history/add_history.py @@ -2,10 +2,10 @@ 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 +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): @@ -25,7 +25,7 @@ class AddHistory(BaseMemoryTool): "properties": {}, "required": [], }, - } + }, ) async def execute(self): diff --git a/reme/tool/memory/read_history.py b/reme/tool/memory/history/read_history.py similarity index 93% rename from reme/tool/memory/read_history.py rename to reme/tool/memory/history/read_history.py index bf43819e..089a3309 100644 --- a/reme/tool/memory/read_history.py +++ b/reme/tool/memory/history/read_history.py @@ -2,8 +2,8 @@ from loguru import logger -from .base_memory_tool import BaseMemoryTool -from ...core.schema import MemoryNode, ToolCall +from ..base_memory_tool import BaseMemoryTool +from ....core.schema import MemoryNode, ToolCall class ReadHistory(BaseMemoryTool): diff --git a/reme_ai/mem_tool/identity/__init__.py b/reme/tool/memory/identity/__init__.py similarity index 100% rename from reme_ai/mem_tool/identity/__init__.py rename to reme/tool/memory/identity/__init__.py 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_ai/mem_tool/meta/__init__.py b/reme/tool/memory/meta/__init__.py similarity index 100% rename from reme_ai/mem_tool/meta/__init__.py rename to reme/tool/memory/meta/__init__.py 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/read_user_profile.py b/reme/tool/memory/user_profile/read_user_profile.py similarity index 93% rename from reme/tool/memory/read_user_profile.py rename to reme/tool/memory/user_profile/read_user_profile.py index 6b3118cf..dfbccd63 100644 --- a/reme/tool/memory/read_user_profile.py +++ b/reme/tool/memory/user_profile/read_user_profile.py @@ -4,9 +4,9 @@ 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 +from ..base_memory_tool import BaseMemoryTool +from ....core.schema import ToolCall +from ....core.schema.memory_node import MemoryNode class ReadUserProfile(BaseMemoryTool): diff --git a/reme/tool/memory/update_user_profile.py b/reme/tool/memory/user_profile/update_user_profile.py similarity index 96% rename from reme/tool/memory/update_user_profile.py rename to reme/tool/memory/user_profile/update_user_profile.py index e2032d45..93f3a5dc 100644 --- a/reme/tool/memory/update_user_profile.py +++ b/reme/tool/memory/user_profile/update_user_profile.py @@ -2,10 +2,10 @@ 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 +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): 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. From 83d111ae4184c5bd03324c07d1004ddd6a0a9741 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 23 Jan 2026 10:43:31 +0800 Subject: [PATCH 19/19] up init --- reme/workflow/procedural_memory/__init__.py | 0 reme/workflow/tool_memory/__init__.py | 0 2 files changed, 0 insertions(+), 0 deletions(-) create mode 100644 reme/workflow/procedural_memory/__init__.py create mode 100644 reme/workflow/tool_memory/__init__.py 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