diff --git a/.gitignore b/.gitignore index b3537023..603ccfcf 100644 --- a/.gitignore +++ b/.gitignore @@ -32,4 +32,7 @@ site/* docs/_build/* test_compact_storage/* test_working_memory/* -*.code-workspace \ No newline at end of file +*.code-workspace +local_vector_store/* +bench_results/* +meta_memory/* \ No newline at end of file diff --git a/bench/eval_reme.py b/bench/eval_reme.py new file mode 100644 index 00000000..8e6c5e9e --- /dev/null +++ b/bench/eval_reme.py @@ -0,0 +1,329 @@ +"""ReMe evaluation script for HaluMem-like benchmarks.""" + +import asyncio +import copy +import json +import os +import re +import time +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.reme import ReMe + +TEMPLATE_REME = """Memories for user {user_id}: + + {memories} +""" + +RETRY_TIMES = 3 +WAIT_TIME = 2 + +# Default prompt for answering questions with memory context +PROMPT_REME = """You are a helpful AI assistant with access to the user's memories. +Use the following context to answer the user's question accurately. + +Context: +{context} + +Question: {question} + +Please provide a detailed and accurate answer based on the available context. +If the context doesn't contain enough information to answer the question, say so clearly.""" + + +async def add_memory_async( + reme: ReMe, + user_id: str, + messages: list[dict], + description: str = "", +): + """Add memory to ReMe system asynchronously.""" + start = time.time() + + result = await reme.summary( + messages=messages, + user_id=user_id, + description=description, + memory_mode="personal", + ) + + duration_ms = (time.time() - start) * 1000 + return result, duration_ms + + +async def search_memory_async( + reme: ReMe, + query: str, + user_id: str, + top_k: int = 20, +): + """Search memory from ReMe system asynchronously.""" + start = time.time() + + result = await reme.retrieve( + query=query, + user_id=user_id, + memory_mode="personal", + top_k=top_k, + ) + + # Format the context + context = TEMPLATE_REME.format( + user_id=user_id, + memories=result if isinstance(result, str) else json.dumps(result, indent=4, ensure_ascii=False), + ) + + duration_ms = (time.time() - start) * 1000 + + return context, result, duration_ms + + +async def llm_request_async(reme: ReMe, prompt: str): + """Make LLM request using ReMe's llm.""" + messages = [ + Message(role=Role.SYSTEM, content="You are a helpful assistant."), + Message(role=Role.USER, content=prompt), + ] + + response = await reme.llm.chat(messages=messages) + return response.content + + +def extract_user_name(persona_info: str): + """Extract user name from persona info.""" + match = re.search(r"Name:\s*(.*?); Gender:", persona_info) + + if match: + username = match.group(1).strip() + return username + else: + raise ValueError("No name found.") + + +async def _process_session_questions( + session: dict, + new_session: dict, + reme: ReMe, + user_name: str, + top_k_value: int, +) -> None: + """Process questions for a session.""" + if "questions" not in session: + return + + new_session["questions"] = [] + + for qa in session["questions"]: + context, _, duration_ms = await search_memory_async( + reme=reme, + query=qa["question"], + user_id=user_name, + top_k=top_k_value, + ) + + new_qa = copy.deepcopy(qa) + new_qa["context"] = context + new_qa["search_duration_ms"] = duration_ms + + prompt = PROMPT_REME.format( + context=context, + question=qa["question"], + ) + + start_time = time.time() + response = await llm_request_async(reme, prompt) + new_qa["system_response"] = response + new_qa["response_duration_ms"] = (time.time() - start_time) * 1000 + + new_session["questions"].append(new_qa) + + +async def process_user_async( + user_data: dict, + top_k_value: int, + save_path: str, + reme: ReMe, +): + """Process a single user's data asynchronously.""" + user_name = extract_user_name(user_data["persona_info"]) + sessions = user_data["sessions"] + + tmp_dir = os.path.join(save_path, "tmp") + os.makedirs(tmp_dir, exist_ok=True) + + tmp_file = os.path.join(tmp_dir, f"{user_data['uuid']}.json") + + # Clear existing memories for this user + await reme.vector_store.delete_collection(f"reme_eval_{user_name}") + + # Update collection name for this user + reme.vector_store.set_collection_name(f"reme_eval_{user_name}") + + new_user_data = { + "uuid": user_data["uuid"], + "user_name": user_name, + "sessions": [], + } + + for session in tqdm(sessions, total=len(sessions), desc=f"Processing user {user_name}"): + new_session = { + "memory_points": session["memory_points"], + "dialogue": session["dialogue"], + } + + # Add messages to ReMe + dialogue = session["dialogue"] + # Parse timestamp and format as "YYYY-MM-DD HH:MM:SS" + date_format = "%b %d, %Y, %H:%M:%S" + # dt = datetime.strptime(session["start_time"], date_format).replace(tzinfo=timezone.utc) + # time_created = dt.strftime("%Y-%m-%d %H:%M:%S") + + formatted_dialogue = [ + { + "role": turn["role"], + "content": turn["content"], + "time_created": datetime.strptime(turn["timestamp"], date_format) + .replace(tzinfo=timezone.utc) + .strftime("%Y-%m-%d %H:%M:%S"), + } + for turn in dialogue + ] + + # Add memory + result, duration_ms = await add_memory_async( + reme=reme, + user_id=user_name, + messages=formatted_dialogue, + ) + memories = [] + for memory_modes in result: + for memory_mode in memory_modes: + if not isinstance(memory_mode, MemoryNode): + continue + + memories.append(memory_mode.content) + + print(memories) + + if session.get("is_generated_qa_session", False): + new_session["add_dialogue_duration_ms"] = duration_ms + new_session["is_generated_qa_session"] = True + del new_session["dialogue"] + del new_session["memory_points"] + new_user_data["sessions"].append(new_session) + continue + + # Store the result from summary + new_session["extracted_memories"] = memories + new_session["add_dialogue_duration_ms"] = duration_ms + + # Search updated memories for memory points + # for memory in new_session["memory_points"]: + # if memory["is_update"] == "False" or not memory["original_memories"]: + # continue + # + # _, memories_from_system, duration_ms = await search_memory_async( + # reme=reme, + # query=memory["memory_content"], + # user_id=user_name, + # top_k=10, + # ) + # + # memory["memories_from_system"] = str(memories_from_system) + + # Process questions + await _process_session_questions(session, new_session, reme, user_name, top_k_value) + + new_user_data["sessions"].append(new_session) + with open(tmp_file, "w", encoding="utf-8") as f: + json.dump(new_user_data, f, ensure_ascii=False, indent=2) + # raise NotImplementedError + + # Save results + with open(tmp_file, "w", encoding="utf-8") as f: + json.dump(new_user_data, f, ensure_ascii=False, indent=2) + + print(f"✅ Saved user {user_name} to {tmp_file}") + return {"uuid": user_data["uuid"], "status": "ok", "path": tmp_file} + + +def iter_jsonl(file_path: str): + """Iterate over lines in a JSONL file.""" + with open(file_path, "r", encoding="utf-8") as f: + for line in f: + line = line.strip() + if line: + yield json.loads(line) + + +async def main_async( + data_path_arg: str, + version_arg: str = "default", + top_k_arg: int = 20, +): + """Main evaluation function.""" + frame = "reme" + save_path = f"bench_results/{frame}-{version_arg}/" + os.makedirs(save_path, exist_ok=True) + + output_file = os.path.join(save_path, f"{frame}_eval_results.jsonl") + tmp_dir = os.path.join(save_path, "tmp") + os.makedirs(tmp_dir, exist_ok=True) + + start_time = time.time() + + # Initialize ReMe instance (will reuse for all users) + reme = ReMe() + + # Load all user data + user_data_list = list(iter_jsonl(data_path_arg)) + total_users = len(user_data_list) + + print(f"Processing {total_users} users sequentially...") + + # Sequential processing + for idx, user_data in enumerate(user_data_list, 1): + result = await process_user_async(user_data, top_k_arg, save_path, reme) + print(f"[{idx}/{total_users}] ✅ Finished {user_data['uuid']} ({result['status']})") + + # Combine all results into final output + with open(output_file, "w", encoding="utf-8") as f_out: + for file in os.listdir(tmp_dir): + if file.endswith(".json"): + file_path = os.path.join(tmp_dir, file) + with open(file_path, "r", encoding="utf-8") as f_in: + data = json.load(f_in) + f_out.write(json.dumps(data, ensure_ascii=False) + "\n") + + elapsed = time.time() - start_time + print(f"✅ All done in {elapsed:.2f}s") + print(f"✅ Final results saved to: {output_file}") + + +def main( + data_path_arg: str, + version_arg: str = "default", + top_k_arg: int = 20, +): + """Synchronous entry point for main evaluation.""" + asyncio.run(main_async(data_path_arg, version_arg, top_k_arg)) + + +if __name__ == "__main__": + # Example usage - update these paths as needed + # Note: Don't use HaluMem-long.jsonl directly as each line is too large + # Instead, create a smaller test dataset or use a different data file + + DEFAULT_DATA_PATH = "/Users/yuli/workspace/HaluMem/data/HaluMem-Long.jsonl" + DEFAULT_VERSION = "test" + DEFAULT_TOP_K = 20 + + main( + data_path_arg=DEFAULT_DATA_PATH, + version_arg=DEFAULT_VERSION, + top_k_arg=DEFAULT_TOP_K, + ) diff --git a/reme_ai/bench/__init__.py b/reme_ai/bench/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/reme_ai/core/embedding/base_embedding_model.py b/reme_ai/core/embedding/base_embedding_model.py index fff59c83..83246f7b 100644 --- a/reme_ai/core/embedding/base_embedding_model.py +++ b/reme_ai/core/embedding/base_embedding_model.py @@ -26,6 +26,7 @@ class BaseEmbeddingModel(ABC): max_batch_size: int = 10, max_retries: int = 3, raise_exception: bool = True, + max_input_length: int = 8192, **kwargs, ): """Initialize model configuration and parameters.""" @@ -34,8 +35,22 @@ class BaseEmbeddingModel(ABC): self.max_batch_size = max_batch_size self.max_retries = max_retries self.raise_exception = raise_exception + self.max_input_length = max_input_length self.kwargs = kwargs + def _truncate_text(self, text: str) -> str: + """Truncate text to max_input_length if it exceeds the limit.""" + if len(text) > self.max_input_length: + logger.warning( + f"Text length {len(text)} exceeds max_input_length {self.max_input_length}, truncating" + ) + return text[: self.max_input_length] + return text + + def _truncate_texts(self, texts: list[str]) -> list[str]: + """Truncate a list of texts to max_input_length.""" + return [self._truncate_text(text) for text in texts] + async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]: """Internal async implementation for calling the embedding API with batch input.""" @@ -44,9 +59,10 @@ class BaseEmbeddingModel(ABC): async def get_embedding(self, input_text: str, **kwargs) -> list[float]: """Async get embedding for a single text with exponential backoff retries.""" + truncated_text = self._truncate_text(input_text) for i in range(self.max_retries): try: - result = await self._get_embeddings([input_text], **kwargs) + result = await self._get_embeddings([truncated_text], **kwargs) return result[0] except Exception as e: logger.error(f"Model {self.model_name} failed: {e}") @@ -59,10 +75,13 @@ class BaseEmbeddingModel(ABC): async def get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]: """Async get embeddings with automatic batching and exponential backoff retries.""" + # Truncate all input texts first + truncated_texts = self._truncate_texts(input_text) + # Split into batches and process sequentially to respect rate limits results = [] - for i in range(0, len(input_text), self.max_batch_size): - batch = input_text[i : i + self.max_batch_size] + for i in range(0, len(truncated_texts), self.max_batch_size): + batch = truncated_texts[i : i + self.max_batch_size] # Process each batch with retry logic for retry in range(self.max_retries): try: @@ -81,9 +100,10 @@ class BaseEmbeddingModel(ABC): def get_embedding_sync(self, input_text: str, **kwargs) -> list[float]: """Synchronous get embedding for a single text with retry logic.""" + truncated_text = self._truncate_text(input_text) for i in range(self.max_retries): try: - result = self._get_embeddings_sync([input_text], **kwargs) + result = self._get_embeddings_sync([truncated_text], **kwargs) return result[0] except Exception as exc: logger.error(f"Model {self.model_name} failed: {exc}") @@ -96,9 +116,12 @@ class BaseEmbeddingModel(ABC): def get_embeddings_sync(self, input_text: list[str], **kwargs) -> list[list[float]]: """Synchronous get embeddings with automatic batching and retry logic.""" + # Truncate all input texts first + truncated_texts = self._truncate_texts(input_text) + results = [] - for i in range(0, len(input_text), self.max_batch_size): - batch = input_text[i : i + self.max_batch_size] + for i in range(0, len(truncated_texts), self.max_batch_size): + batch = truncated_texts[i : i + self.max_batch_size] # Process each batch with retry logic for retry in range(self.max_retries): try: diff --git a/reme_ai/core/flow/base_flow.py b/reme_ai/core/flow/base_flow.py index 5fb3d00e..7decce48 100644 --- a/reme_ai/core/flow/base_flow.py +++ b/reme_ai/core/flow/base_flow.py @@ -62,7 +62,7 @@ class BaseFlow(ABC): payload = json.dumps(params, sort_keys=True, ensure_ascii=False, default=str) return hashlib.sha256(payload.encode("utf-8")).hexdigest() except Exception as e: - logger.exception(f"{self.name} cache key serialization failed: {e}") + logger.exception(f"[{self.__class__.__name__}] {self.name} cache key serialization failed: {e}") return None def _maybe_load_cached(self, params: dict) -> Response | None: @@ -72,7 +72,7 @@ class BaseFlow(ABC): if key := self._compute_cache_key(params): if cached := self.cache.load(key): - logger.info(f"Loaded {self.name} response from cache.") + logger.info(f"[{self.__class__.__name__}] Loaded {self.name} response from cache.") return Response(**cached) return None @@ -92,7 +92,7 @@ class BaseFlow(ABC): """Recursively log the hierarchy of the flow's operation tree.""" prefix = " " * indent op_type = "sequential" if isinstance(op, SequentialOp) else "parallel" if isinstance(op, ParallelOp) else name - logger.info(f"{prefix}{op_type} execution") + logger.info(f"[{self.__class__.__name__}] {prefix}{op_type} execution") for sub_op in op.sub_ops or []: self._print_operation_tree(sub_op.name, sub_op, indent + 2) @@ -160,15 +160,15 @@ class BaseFlow(ABC): def print_flow(self): """Log the visual structure of the flow once.""" if not self._flow_printed: - logger.info(f"---------- [Flow Structure] {self.name} ----------") + logger.info(f"[{self.__class__.__name__}] ---------- [Flow Structure] {self.name} ----------") self._print_operation_tree(self.name, self.flow_op, 0) - logger.info("-" * 50) + logger.info(f"[{self.__class__.__name__}] " + "-" * 50) self._flow_printed = True async def call(self, **kwargs) -> Response | asyncio.Queue: """Execute the flow asynchronously with parameter caching.""" kwargs["stream"] = self.stream - logger.info(f"{self.name} incoming params: {kwargs}") + logger.info(f"[{self.__class__.__name__}] {self.name} incoming params: {kwargs}") if cached := self._maybe_load_cached(kwargs): return cached @@ -187,7 +187,7 @@ class BaseFlow(ABC): self._maybe_save_cache(kwargs, result) return result except Exception as e: - logger.exception(f"{self.name} async call failed: {e}") + logger.exception(f"[{self.__class__.__name__}] {self.name} async call failed: {e}") if self.raise_exception: raise e if self.stream: @@ -199,7 +199,7 @@ class BaseFlow(ABC): def call_sync(self, **kwargs) -> Response: """Execute the flow synchronously with parameter caching.""" - logger.info(f"{self.name} incoming sync params: {kwargs}") + logger.info(f"[{self.__class__.__name__}] {self.name} incoming sync params: {kwargs}") assert not self.stream, "Synchronous call cannot be used in stream mode." if cached := self._maybe_load_cached(kwargs): return cached @@ -214,7 +214,7 @@ class BaseFlow(ABC): self._maybe_save_cache(kwargs, context.response) return context.response except Exception as e: - logger.exception(f"{self.name} sync call failed: {e}") + logger.exception(f"[{self.__class__.__name__}] {self.name} sync call failed: {e}") if self.raise_exception: raise e context.add_response_error(e) diff --git a/reme_ai/core/op/base_op.py b/reme_ai/core/op/base_op.py index 24c416d4..357db44e 100644 --- a/reme_ai/core/op/base_op.py +++ b/reme_ai/core/op/base_op.py @@ -100,7 +100,7 @@ class BaseOp: def _handle_failure(self, e: Exception, attempt: int): """Log failures and handle final retry logic.""" - message = f"{self.name} failed (attempt {attempt + 1}): {e}" + 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: @@ -305,7 +305,7 @@ class BaseOp: results = [] for res in raw_results: if isinstance(res, Exception): - logger.error(f"Async task failed: {res}") + logger.error(f"[{self.__class__.__name__}] Async task failed: {res}") continue if res: results.extend(res if isinstance(res, list) else [res]) diff --git a/reme_ai/core/schema/message.py b/reme_ai/core/schema/message.py index ca53f2fc..1748f961 100644 --- a/reme_ai/core/schema/message.py +++ b/reme_ai/core/schema/message.py @@ -124,12 +124,12 @@ class Message(BaseModel): """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}:\n" + header = f"{self.name or self.role.value if use_name else self.role.value}:" lines = [f"{prefix}{time_str}{header}"] if add_reasoning and self.reasoning_content: - lines.append(f"{self.reasoning_content}\n") + lines.append(self.reasoning_content) if isinstance(self.content, str): lines.append(self.content) @@ -144,7 +144,7 @@ class Message(BaseModel): for tc in self.tool_calls: lines.append(f" - tool_call={tc.name} params={tc.arguments}") - return "\n".join(lines).strip() + return " ".join(lines).strip() class Trajectory(BaseModel): diff --git a/reme_ai/core/vector_store/base_vector_store.py b/reme_ai/core/vector_store/base_vector_store.py index e64b6dbf..5815dca5 100644 --- a/reme_ai/core/vector_store/base_vector_store.py +++ b/reme_ai/core/vector_store/base_vector_store.py @@ -48,6 +48,10 @@ class BaseVectorStore(ABC): """Convert multiple text queries into vector embeddings using the configured model.""" return await self.embedding_model.get_embeddings(queries) + def set_collection_name(self, collection_name: str): + """Change the name of the current collection.""" + self.collection_name = collection_name + @abstractmethod async def list_collections(self) -> list[str]: """Retrieve a list of all existing collection names in the store.""" diff --git a/reme_ai/core/vector_store/chroma_vector_store.py b/reme_ai/core/vector_store/chroma_vector_store.py index 24b88d36..231ca248 100644 --- a/reme_ai/core/vector_store/chroma_vector_store.py +++ b/reme_ai/core/vector_store/chroma_vector_store.py @@ -103,7 +103,7 @@ class ChromaVectorStore(BaseVectorStore): metadata = metadatas[i] if i < len(metadatas) and metadatas[i] else {} if include_score and distances and i < len(distances): - metadata["_score"] = 1.0 - distances[i] + metadata["score"] = 1.0 - distances[i] node = VectorNode( vector_id=vector_id, @@ -300,7 +300,7 @@ class ChromaVectorStore(BaseVectorStore): score_threshold = kwargs.get("score_threshold") if score_threshold is not None: - nodes = [n for n in nodes if n.metadata.get("_score", 0) >= score_threshold] + nodes = [n for n in nodes if n.metadata.get("score", 0) >= score_threshold] return nodes async def delete(self, vector_ids: str | list[str], **kwargs): @@ -392,6 +392,15 @@ class ChromaVectorStore(BaseVectorStore): await self._run_sync_in_executor(_recreate) logger.info(f"Collection {self.collection_name} has been reset") + def set_collection_name(self, collection_name: str): + """Set the collection name and reinitialize the collection object.""" + super().set_collection_name(collection_name) + self.collection = self.client.get_or_create_collection( + name=collection_name, + metadata={"hnsw:space": "cosine"}, + ) + logger.info(f"Collection name set to {collection_name}, collection object reinitialized") + async def close(self): """Close the vector store and log the shutdown process.""" logger.info(f"ChromaDB vector store for collection {self.collection_name} closed") diff --git a/reme_ai/core/vector_store/es_vector_store.py b/reme_ai/core/vector_store/es_vector_store.py index d68f283e..0f1fa19f 100644 --- a/reme_ai/core/vector_store/es_vector_store.py +++ b/reme_ai/core/vector_store/es_vector_store.py @@ -279,7 +279,7 @@ class ESVectorStore(BaseVectorStore): vector=source.get("vector"), metadata=source.get("metadata", {}), ) - node.metadata["_score"] = hit["_score"] + node.metadata["score"] = hit["_score"] results.append(node) return results @@ -452,6 +452,12 @@ class ESVectorStore(BaseVectorStore): return results + def set_collection_name(self, collection_name: str): + """Set the collection name and ensure it's lowercase for Elasticsearch compatibility.""" + collection_name = collection_name.lower() + super().set_collection_name(collection_name) + logger.info(f"Collection name set to {collection_name} (converted to lowercase)") + async def close(self): """Terminate the Elasticsearch client session and release resources.""" await self.client.close() diff --git a/reme_ai/core/vector_store/local_vector_store.py b/reme_ai/core/vector_store/local_vector_store.py index 25beef8d..04649841 100644 --- a/reme_ai/core/vector_store/local_vector_store.py +++ b/reme_ai/core/vector_store/local_vector_store.py @@ -202,7 +202,7 @@ class LocalVectorStore(BaseVectorStore): scored_nodes = scored_nodes[:limit] results = [] for node, score in scored_nodes: - node.metadata["_score"] = score + node.metadata["score"] = score results.append(node) return results @@ -276,6 +276,12 @@ class LocalVectorStore(BaseVectorStore): return filtered_nodes + def set_collection_name(self, collection_name: str): + """Set the collection name and reinitialize the collection path.""" + super().set_collection_name(collection_name) + self.collection_path = self.root_path / collection_name + logger.info(f"Collection name set to {collection_name}, path updated to {self.collection_path}") + async def close(self): """Close the vector store (no-op for local file system).""" logger.info("Local vector store closed") diff --git a/reme_ai/core/vector_store/pgvector_store.py b/reme_ai/core/vector_store/pgvector_store.py index 38c5667a..695cee70 100644 --- a/reme_ai/core/vector_store/pgvector_store.py +++ b/reme_ai/core/vector_store/pgvector_store.py @@ -324,7 +324,7 @@ class PGVectorStore(BaseVectorStore): if isinstance(metadata, str): metadata = json.loads(metadata) - metadata["_score"] = 1 - distance + metadata["score"] = 1 - distance metadata["_distance"] = distance node = VectorNode( diff --git a/reme_ai/core/vector_store/qdrant_vector_store.py b/reme_ai/core/vector_store/qdrant_vector_store.py index 97eb5b61..8227c69b 100644 --- a/reme_ai/core/vector_store/qdrant_vector_store.py +++ b/reme_ai/core/vector_store/qdrant_vector_store.py @@ -304,7 +304,7 @@ class QdrantVectorStore(BaseVectorStore): vector=point.vector if hasattr(point, "vector") else None, metadata=payload.get("metadata", {}), ) - node.metadata["_score"] = point.score + node.metadata["score"] = point.score nodes.append(node) return nodes diff --git a/reme_ai/mem_agent/base_memory_agent.py b/reme_ai/mem_agent/base_memory_agent.py index f5b33329..c7db4231 100644 --- a/reme_ai/mem_agent/base_memory_agent.py +++ b/reme_ai/mem_agent/base_memory_agent.py @@ -1,13 +1,14 @@ """Base memory agent for handling memory operations with tool-based reasoning.""" import asyncio +import json 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 +from ..core.schema import Message, ToolCall, MemoryNode from ..mem_tool import BaseMemoryTool, ThinkTool @@ -35,6 +36,7 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): self.messages: list[Message] = [] self.success: bool = True + self.memory_nodes: list[MemoryNode | str] = [] def _build_tool_call(self) -> ToolCall: return ToolCall( @@ -71,7 +73,7 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): ) @property - def tools(self): + def tools(self) -> list[BaseMemoryTool]: """Returns the list of memory tools available to this agent.""" return self.sub_ops @@ -79,7 +81,7 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): def tools(self, tools: list[BaseMemoryTool]): self.sub_ops = tools - def get_messages(self) -> list[Message]: + def get_messages(self) -> list[Message] | str: """Extracts and returns messages from the context query or messages.""" if self.context.get("query"): messages = [Message(role=Role.USER, content=self.context.query)] @@ -100,7 +102,7 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): **kwargs, ) messages.append(assistant_message) - logger.info(f"step{step + 1}.assistant={assistant_message.simple_dump(enable_json_dump=True)}") + logger.info(f"[{self.__class__.__name__}] step{step + 1}.assistant={assistant_message.simple_dump(enable_json_dump=True)}") should_act = bool(assistant_message.tool_calls) return assistant_message, should_act @@ -114,10 +116,10 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): for j, tool_call in enumerate(assistant_message.tool_calls): if tool_call.name not in tool_dict: - logger.warning(f"unknown tool_call.name={tool_call.name}") + logger.warning(f"[{self.__class__.__name__}] unknown tool_call.name={tool_call.name}") continue - logger.info(f"step{step + 1}.{j} submit tool_calls={tool_call.name} argument={tool_call.arguments}") + logger.info(f"[{self.__class__.__name__}] step{step + 1}.{j} submit tool_calls={tool_call.name} argument={tool_call.arguments}") tool_copy: BaseMemoryTool = tool_dict[tool_call.name].copy() tool_copy.tool_call.id = tool_call.id tool_list.append(tool_copy) @@ -129,6 +131,9 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): await self.join_async_tasks() for j, op in enumerate(tool_list): + if op.memory_nodes: + self.memory_nodes.extend(op.memory_nodes) + tool_result = str(op.output) tool_message = Message( role=Role.TOOL, @@ -136,7 +141,7 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): tool_call_id=op.tool_call.id, ) tool_result_messages.append(tool_message) - logger.info(f"step{step + 1}.{j} join tool_result={tool_result[:200]}...\n\n") + logger.info(f"[{self.__class__.__name__}] step{step + 1}.{j} join tool_result={tool_result[:500]}...\n\n") return tool_result_messages async def react(self, messages: list[Message]): @@ -157,13 +162,16 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): async def execute(self): messages = await self.build_messages() for i, message in enumerate(messages): - logger.info(f"step0.{i} {message.role} {message.name or ''} {message.simple_dump(enable_json_dump=True)}") + logger.info(f"[{self.__class__.__name__}] step0.{i} {message.role} {message.name or ''} " + f"{message.simple_dump(enable_json_dump=True)}") + for i, tool in enumerate(self.tools): + logger.info(f"[{self.__class__.__name__}] step0.{i} tool_call={json.dumps(tool.tool_call.simple_input_dump(), ensure_ascii=False)}") self.messages, self.success = await self.react(messages) if self.success and self.messages: self.output = self.messages[-1].content else: - self.output = "" + self.output = "No relevant memories found." @property def memory_target(self) -> str: diff --git a/reme_ai/mem_agent/retriever/reme_retriever.py b/reme_ai/mem_agent/retriever/reme_retriever.py index 9ba15d66..6841432a 100644 --- a/reme_ai/mem_agent/retriever/reme_retriever.py +++ b/reme_ai/mem_agent/retriever/reme_retriever.py @@ -13,33 +13,36 @@ from ...core.utils import get_now_time, format_messages class ReMeRetriever(BaseMemoryAgent): """Memory agent that retrieves and builds messages with meta memory context.""" - def __init__(self, meta_memories: list[dict] = None, **kwargs): + def __init__(self, meta_memories: list[dict] | None = None, **kwargs): super().__init__(**kwargs) - self.meta_memories: list[dict] = meta_memories + self.meta_memories: list[dict] = meta_memories or [] - @staticmethod - async def _read_meta_memories() -> str: - """Read and return meta memories as string.""" + 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) - await op.call() - return str(op.output) + if self.meta_memories: + return op.format_memory_metadata(self.meta_memories) + else: + await op.call() + return str(op.output) async def build_messages(self) -> List[Message]: """Build messages with system prompt and user message.""" - from ...mem_tool import ReadMetaMemory - - if self.meta_memories: - meta_memory_info = ReadMetaMemory().format_memory_metadata(self.meta_memories) + meta_memory_info = await self._read_meta_memories() + if self.context.get("query"): + context = self.context.query + elif self.context.get("messages"): + messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages] + context = self.description + format_messages(messages) else: - meta_memory_info = await self._read_meta_memories() - + raise ValueError("input must have either `query` or `messages`") system_prompt = self.prompt_format( prompt_name="system_prompt", now_time=get_now_time(), meta_memory_info=meta_memory_info, - context=format_messages(self.get_messages()), + context=context, ) messages = [ diff --git a/reme_ai/mem_agent/retriever/reme_retriever.yaml b/reme_ai/mem_agent/retriever/reme_retriever.yaml index c818edd3..056b5138 100644 --- a/reme_ai/mem_agent/retriever/reme_retriever.yaml +++ b/reme_ai/mem_agent/retriever/reme_retriever.yaml @@ -1,12 +1,12 @@ tool: | - Retrieve relevant memories from the memory bank to assist in answering questions. + Retrieve relevant memories to assist in answering questions. Use this tool when you need to search for historical information, user preferences, procedural knowledge, or any other stored memories that may help answer the current query. The agent will analyze the context, determine what information is needed, and perform semantic searches across different memory types to find the most relevant memories. system_prompt: | - You are a memory agent. Please analyze the context, retrieve relevant information from the memory bank when needed, and return a summary of the retrieved memories to assist in answering the user's question. + You are a memory agent. Please analyze the context, retrieve relevant memories when needed, and return a summary of the retrieved memories to assist in answering the user's question. ## Context {context} @@ -47,4 +47,4 @@ system_prompt: | - If multiple attempts still yield no relevant memory, output ``. user_message: | - Please analyze the context, retrieve relevant information from the memory bank when needed, and return a summary of the retrieved memories to assist in answering the user's question. + Please analyze the context, retrieve relevant memories when needed, and return a summary of the retrieved memories to assist in answering the user's question. diff --git a/reme_ai/mem_agent/summarizer/personal_summarizer.py b/reme_ai/mem_agent/summarizer/personal_summarizer.py index 424eec39..696ff256 100644 --- a/reme_ai/mem_agent/summarizer/personal_summarizer.py +++ b/reme_ai/mem_agent/summarizer/personal_summarizer.py @@ -3,7 +3,7 @@ from ..base_memory_agent import BaseMemoryAgent from ...core.context import C from ...core.enumeration import Role, MemoryType -from ...core.schema import Message +from ...core.schema import Message, ToolCall from ...core.utils import get_now_time, format_messages @@ -13,6 +13,36 @@ class PersonalSummarizer(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( @@ -34,9 +64,9 @@ class PersonalSummarizer(BaseMemoryAgent): return await super()._acting_step( assistant_message, step, - memory_target=self.memory_target, memory_type=self.memory_type.value, - author=self.author, + memory_target=self.memory_target, ref_memory_id=self.ref_memory_id, + author=self.author, **kwargs, ) diff --git a/reme_ai/mem_agent/summarizer/reme_summarizer.py b/reme_ai/mem_agent/summarizer/reme_summarizer.py index e5222e7b..02bd5b73 100644 --- a/reme_ai/mem_agent/summarizer/reme_summarizer.py +++ b/reme_ai/mem_agent/summarizer/reme_summarizer.py @@ -8,7 +8,7 @@ from loguru import logger from ..base_memory_agent import BaseMemoryAgent from ...core.context import C from ...core.enumeration import Role -from ...core.schema import Message, MemoryNode +from ...core.schema import Message, MemoryNode, ToolCall from ...core.utils import get_now_time, format_messages @@ -16,11 +16,41 @@ from ...core.utils import get_now_time, format_messages class ReMeSummarizer(BaseMemoryAgent): """Coordinates memory updates by delegating to specialized memory agents.""" - def __init__(self, enable_tool_memory: bool = True, enable_identity_memory: bool = True, **kwargs): - """Initialize with flags to enable/disable tool and identity memory processing.""" + def __init__(self, meta_memories: list[dict] | None = None, enable_identity_memory: bool = False, **kwargs): + """Initialize with flags to enable/disable identity memory processing.""" super().__init__(**kwargs) - self.enable_tool_memory = enable_tool_memory self.enable_identity_memory = enable_identity_memory + 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 _add_history_memory(self) -> MemoryNode: """Store conversation history and return the memory node.""" @@ -28,7 +58,7 @@ class ReMeSummarizer(BaseMemoryAgent): op = AddHistoryMemory() await op.call(messages=self.get_messages()) - return op.output + return op.memory_nodes[0] @staticmethod async def _read_identity_memory() -> str: @@ -43,27 +73,28 @@ class ReMeSummarizer(BaseMemoryAgent): """Fetch all meta-memory entries that define specialized memory agents.""" 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) + op = ReadMetaMemory(enable_identity_memory=self.enable_identity_memory) + if self.meta_memories: + return op.format_memory_metadata(self.meta_memories) + else: + await op.call() + return str(op.output) async def build_messages(self) -> List[Message]: """Construct initial messages with context, identity, and meta-memory information.""" memory_node: MemoryNode = await self._add_history_memory() self.context["ref_memory_id"] = memory_node.memory_id + now_time = get_now_time() identity_memory = await self._read_identity_memory() meta_memory_info = await self._read_meta_memories() - context = format_messages(self.get_messages()) + context = self.description + "\n" + format_messages(self.get_messages()) logger.info( f"now_time={now_time} " - f"memory_node={memory_node} " + f"memory_node={memory_node.content[:100]}... " f"identity_memory={identity_memory} " f"meta_memory_info={meta_memory_info} " - f"context={context}", + f"context={context[:100]}", ) system_prompt = self.prompt_format( @@ -84,13 +115,21 @@ class ReMeSummarizer(BaseMemoryAgent): async def _reasoning_step(self, messages: list[Message], step: int, **kwargs) -> tuple[Message, bool]: """Refresh meta-memory info in system prompt before each reasoning step.""" - meta_memory_info = await self._read_meta_memories() system_messages = [message for message in messages if message.role is Role.SYSTEM] + if system_messages: system_message = system_messages[0] - pattern = r'("- \(\): "\n)(.*?)(\n\n)' - replacement = rf"\g<1>{meta_memory_info}\g<3>" - system_message.content = re.sub(pattern, replacement, system_message.content, flags=re.DOTALL) + now_time = get_now_time() + identity_memory = await self._read_identity_memory() + meta_memory_info = await self._read_meta_memories() + context = self.description + "\n" + format_messages(self.get_messages()) + system_message.content = self.prompt_format( + prompt_name="system_prompt", + now_time=now_time, + identity_memory=identity_memory, + meta_memory_info=meta_memory_info, + context=context, + ) return await super()._reasoning_step(messages, step, **kwargs) @@ -99,6 +138,8 @@ class ReMeSummarizer(BaseMemoryAgent): 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"], author=self.author, **kwargs, diff --git a/reme_ai/mem_agent/summarizer/reme_summarizer.yaml b/reme_ai/mem_agent/summarizer/reme_summarizer.yaml index baa6a7d6..ad254342 100644 --- a/reme_ai/mem_agent/summarizer/reme_summarizer.yaml +++ b/reme_ai/mem_agent/summarizer/reme_summarizer.yaml @@ -6,11 +6,11 @@ tool: | 3. Delegating to specialized memory agents for detailed memory extraction and update 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} - You are a Memory Agent responsible for performing necessary updates and summaries of the main Agent's memories based on the **context**. - ## Current Time {now_time} @@ -25,11 +25,12 @@ system_prompt: | ## Your Tasks ### 1. Create New Meta Memory (if needed) - When the context contains significant personal or procedural information not yet covered by existing meta memories: - - Use `add_meta_memory` to create one or more new meta memory entries. + When the context contains significant new valuable information, first check if the Main Agent's Meta Memory already contains a corresponding `()` entry: + - If the required `()` does NOT exist in the Meta Memory, use `add_meta_memory` to create a new meta memory entry. - For personal memories: specify `memory_type="personal"` and `memory_target=`. - For procedural memories: specify `memory_type="procedural"` and `memory_target=`. - Each meta memory entry will instantiate a dedicated specialized Memory Agent for that dimension. + - Only create new meta memory entries when necessary; avoid duplicating existing ones. ### 2. Add Summary Memory (if valuable) When the context includes information worth remembering for quick future recall: diff --git a/reme_ai/mem_tool/base_memory_tool.py b/reme_ai/mem_tool/base_memory_tool.py index 7ec1a133..72a94b2e 100644 --- a/reme_ai/mem_tool/base_memory_tool.py +++ b/reme_ai/mem_tool/base_memory_tool.py @@ -3,6 +3,8 @@ 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 @@ -23,7 +25,7 @@ class BaseMemoryTool(BaseOp, metaclass=ABCMeta): self.enable_multiple: bool = enable_multiple self.enable_thinking_params: bool = enable_thinking_params self.meta_memory_path: str = meta_memory_path - self._meta_memory: CacheHandler | None = None + self.memory_nodes: list[MemoryNode | str] = [] def _build_parameters(self) -> dict: return {} @@ -58,10 +60,8 @@ class BaseMemoryTool(BaseOp, metaclass=ABCMeta): @property def meta_memory(self) -> CacheHandler: - """Get or create the meta memory cache handler.""" - if self._meta_memory is None: - self._meta_memory = CacheHandler(Path(self.meta_memory_path) / self.vector_store.collection_name) - return self._meta_memory + """Create the meta memory cache handler.""" + return CacheHandler(Path(self.meta_memory_path) / self.vector_store.collection_name) @property def memory_type(self) -> MemoryType: @@ -94,7 +94,7 @@ class BaseMemoryTool(BaseOp, metaclass=ABCMeta): metadata: dict | None = None, ) -> MemoryNode: """Build MemoryNode from content, when_to_use, and metadata.""" - return MemoryNode( + 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 "", @@ -103,3 +103,6 @@ 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 \ No newline at end of file diff --git a/reme_ai/mem_tool/hands_off_tool.py b/reme_ai/mem_tool/hands_off_tool.py index a3924ba1..2cb28d60 100644 --- a/reme_ai/mem_tool/hands_off_tool.py +++ b/reme_ai/mem_tool/hands_off_tool.py @@ -35,12 +35,7 @@ class HandsOffTool(BaseMemoryTool): "memory_type": { "type": "string", "description": self.get_prompt("memory_type"), - "enum": [ - MemoryType.IDENTITY.value, - MemoryType.PERSONAL.value, - MemoryType.PROCEDURAL.value, - MemoryType.TOOL.value, - ], + "enum": [k.value for k in self.memory_agent_dict], }, "memory_target": { "type": "string", @@ -82,10 +77,7 @@ class HandsOffTool(BaseMemoryTool): def _parse_memory_type_target(task: dict): memory_type = task.get("memory_type", "") memory_target = task.get("memory_target", "") - return { - "memory_type": MemoryType(memory_type), - "memory_target": memory_target, - } + return {"memory_type": MemoryType(memory_type), "memory_target": memory_target} def _collect_tasks(self) -> list[dict]: """Collect memory tasks from context based on enable_multiple flag.""" @@ -116,21 +108,17 @@ class HandsOffTool(BaseMemoryTool): logger.warning(f"No agent found for memory_type={memory_type}") continue - agent_copy = self.memory_agent_dict[memory_type].copy() - agent_list.append( - { - "agent": agent_copy, - "memory_type": memory_type, - "memory_target": 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_copy.call, + 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", ""), ) @@ -140,6 +128,9 @@ class HandsOffTool(BaseMemoryTool): 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) + results.append( { "memory_type": memory_type.value, diff --git a/reme_ai/mem_tool/history/add_history_memory.py b/reme_ai/mem_tool/history/add_history_memory.py index 65065a2c..a92deca5 100644 --- a/reme_ai/mem_tool/history/add_history_memory.py +++ b/reme_ai/mem_tool/history/add_history_memory.py @@ -40,10 +40,11 @@ class AddHistoryMemory(BaseMemoryTool): 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/read_history_memory.py b/reme_ai/mem_tool/history/read_history_memory.py index 056a7ca7..d01cb4e8 100644 --- a/reme_ai/mem_tool/history/read_history_memory.py +++ b/reme_ai/mem_tool/history/read_history_memory.py @@ -15,48 +15,49 @@ class ReadHistoryMemory(BaseMemoryTool): return { "type": "object", "properties": { - "memory_id": { + "ref_memory_id": { "type": "string", - "description": self.get_prompt("memory_id"), + "description": self.get_prompt("ref_memory_id"), }, }, - "required": ["memory_id"], + "required": ["ref_memory_id"], } def _build_multiple_parameters(self) -> dict: return { "type": "object", "properties": { - "memory_ids": { + "ref_memory_ids": { "type": "array", - "description": self.get_prompt("memory_ids"), + "description": self.get_prompt("ref_memory_ids"), "items": {"type": "string"}, }, }, - "required": ["memory_ids"], + "required": ["ref_memory_ids"], } async def execute(self): if self.enable_multiple: - memory_ids: list[str] = self.context.get("memory_ids", []) + ref_memory_ids: list[str] = self.context.get("ref_memory_ids", []) else: - memory_id = self.context.get("memory_id", "") - memory_ids: list[str] = [memory_id] if memory_id else [] + ref_memory_id = self.context.get("ref_memory_id", "") + ref_memory_ids: list[str] = [ref_memory_id] if ref_memory_id else [] - memory_ids = [mid for mid in memory_ids if mid] + ref_memory_ids = [mid for mid in ref_memory_ids if mid] - if not memory_ids: - self.output = "No valid history memory IDs provided for reading." + if not ref_memory_ids: + self.output = "No valid reference memory IDs provided for reading." logger.warning(self.output) return - nodes = await self.vector_store.get(vector_ids=memory_ids) + # 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 IDs." + 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.") + 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 index f50878a7..564af50f 100644 --- a/reme_ai/mem_tool/history/read_history_memory.yaml +++ b/reme_ai/mem_tool/history/read_history_memory.yaml @@ -1,11 +1,11 @@ tool: | - Read history memory by ID. + Read original history dialogue by reference memory ID. tool_multiple: | - Read multiple history memories by IDs. + Read multiple original history dialogues by reference memory IDs. -memory_id: | - Unique identifier of the history memory. +ref_memory_id: | + Reference memory ID to query the original history dialogue. -memory_ids: | - List of unique identifiers of history memories. +ref_memory_ids: | + List of reference memory IDs to query the original history dialogues. diff --git a/reme_ai/mem_tool/meta/add_meta_memory.yaml b/reme_ai/mem_tool/meta/add_meta_memory.yaml index d74628ca..d54c3345 100644 --- a/reme_ai/mem_tool/meta/add_meta_memory.yaml +++ b/reme_ai/mem_tool/meta/add_meta_memory.yaml @@ -1,11 +1,13 @@ 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. diff --git a/reme_ai/mem_tool/vector/add_memory.py b/reme_ai/mem_tool/vector/add_memory.py index c4a1a9c2..1f87df4b 100644 --- a/reme_ai/mem_tool/vector/add_memory.py +++ b/reme_ai/mem_tool/vector/add_memory.py @@ -71,6 +71,7 @@ class AddMemory(BaseMemoryTool): "description": metadata_description, "properties": metadata_properties, } + required.append("metadata") return properties, required @@ -158,6 +159,7 @@ class AddMemory(BaseMemoryTool): # Delete existing IDs (upsert behavior), then insert 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/vector/add_summary_memory.py b/reme_ai/mem_tool/vector/add_summary_memory.py index 4eb63d34..f3667d84 100644 --- a/reme_ai/mem_tool/vector/add_summary_memory.py +++ b/reme_ai/mem_tool/vector/add_summary_memory.py @@ -4,6 +4,8 @@ from loguru import logger from .add_memory import AddMemory from ...core.context import C +from ...core.enumeration import MemoryType +from ...core.schema import MemoryNode @C.register_op() @@ -53,6 +55,7 @@ class AddSummaryMemory(AddMemory): "description": metadata_description, "properties": metadata_properties, } + required.append("metadata") return { "type": "object", @@ -60,6 +63,30 @@ class AddSummaryMemory(AddMemory): "required": required, } + 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=MemoryType.SUMMARY, + memory_target="", + when_to_use="", + content=memory_content, + ref_memory_id=self.ref_memory_id, + author=self.author, + metadata=metadata or {}, + ) + logger.info(f"Adding summary memory: {node.model_dump_json(indent=2, exclude_none=True)}") + + return node + async def execute(self): """Execute addition: map summary_memory to memory_content and call parent.""" # Map summary_memory to memory_content diff --git a/reme_ai/mem_tool/vector/delete_memory.py b/reme_ai/mem_tool/vector/delete_memory.py index a40f7509..45b28632 100644 --- a/reme_ai/mem_tool/vector/delete_memory.py +++ b/reme_ai/mem_tool/vector/delete_memory.py @@ -56,5 +56,6 @@ class DeleteMemory(BaseMemoryTool): return await self.vector_store.delete(vector_ids=memory_ids) + self.memory_nodes = memory_ids self.output = f"Successfully deleted {len(memory_ids)} memories from vector_store." logger.info(self.output) diff --git a/reme_ai/mem_tool/vector/update_memory.py b/reme_ai/mem_tool/vector/update_memory.py index 3a92722a..4873ce24 100644 --- a/reme_ai/mem_tool/vector/update_memory.py +++ b/reme_ai/mem_tool/vector/update_memory.py @@ -169,6 +169,7 @@ class UpdateMemory(BaseMemoryTool): all_ids_to_delete = list(set(old_memory_ids + new_vector_ids)) await self.vector_store.delete(vector_ids=all_ids_to_delete) await self.vector_store.insert(nodes=vector_nodes) + self.memory_nodes = new_memory_nodes self.output = f"Update: deleted {len(old_memory_ids)} old memories, added {len(new_memory_nodes)} new memories." logger.info(self.output) diff --git a/reme_ai/mem_tool/vector/vector_retrieve_memory.py b/reme_ai/mem_tool/vector/vector_retrieve_memory.py index 06350b46..2d758994 100644 --- a/reme_ai/mem_tool/vector/vector_retrieve_memory.py +++ b/reme_ai/mem_tool/vector/vector_retrieve_memory.py @@ -23,7 +23,7 @@ class VectorRetrieveMemory(BaseMemoryTool): enable_summary_memory: bool = False, add_memory_type_target: bool = False, metadata_desc: dict[str, str] | None = None, - top_k: int = 10, + top_k: int = 20, **kwargs, ): """Initialize VectorRetrieveMemory. @@ -169,11 +169,7 @@ class VectorRetrieveMemory(BaseMemoryTool): value = str(value).strip() filter_dict[key] = [value] if not isinstance(value, list) else value - nodes: list[VectorNode] = await self.vector_store.search( - query=query, - top_k=self.top_k, - filter_dict=filter_dict, - ) + nodes: list[VectorNode] = await self.vector_store.search(query=query, limit=self.top_k, filters=filter_dict) memory_nodes: list[MemoryNode] = [MemoryNode.from_vector_node(n) for n in nodes] @@ -220,8 +216,8 @@ class VectorRetrieveMemory(BaseMemoryTool): self.output = "No valid query texts provided for retrieval." return - # Retrieve memories for all queries - memories: list[MemoryNode] = [] + # Retrieve memory_nodes for all queries + memory_nodes: list[MemoryNode] = [] for item in query_items: memory_type = item.get("memory_type") or default_memory_type memory_target = item.get("memory_target") or default_memory_target @@ -237,14 +233,15 @@ class VectorRetrieveMemory(BaseMemoryTool): query=item["query"], metadata_filters=metadata_filters, ) - memories.extend(retrieved) + memory_nodes.extend(retrieved) # Deduplicate and format output - memories = deduplicate_memories(memories) + memory_nodes = deduplicate_memories(memory_nodes) + self.memory_nodes = memory_nodes - if not memories: - self.output = "No memories found matching the query." + if not memory_nodes: + self.output = "No memory_nodes found matching the query." else: - self.output = "\n".join([m.format_memory() for m in memories]) + self.output = "\n".join([m.format_memory() for m in memory_nodes]) - logger.info(f"Retrieved {len(memories)} memories") + logger.info(f"Retrieved {len(memory_nodes)} memory_nodes") diff --git a/reme_ai/reme.py b/reme_ai/reme.py index 93b00ede..a4f9de4c 100644 --- a/reme_ai/reme.py +++ b/reme_ai/reme.py @@ -1,27 +1,18 @@ """ReMe classes for simplified configuration and execution.""" -from typing import Literal - 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.vector_store import BaseVectorStore -from .mem_agent.summarizer import ( - ReMeSummarizer, - # ToolSummarizer, - PersonalSummarizer, - ProceduralSummarizer, - # IdentitySummarizer, -) from .mem_agent.retriever import ReMeRetriever - -# from .mem_agent.chat import ReMyAgent +from .mem_agent.summarizer import ReMeSummarizer, PersonalSummarizer from .mem_tool import ( HandsOffTool, ReadHistoryMemory, - # ReadIdentityMemory, - # UpdateIdentityMemory, AddMetaMemory, AddMemory, AddSummaryMemory, @@ -29,7 +20,6 @@ from .mem_tool import ( UpdateMemory, VectorRetrieveMemory, ) -from .core.schema import Message class ReMe(Application): @@ -47,10 +37,6 @@ class ReMe(Application): embedding_model: dict | None = None, vector_store: dict | None = None, token_counter: dict | None = None, - enable_identity_memory: bool = True, - enable_tool_memory: bool = True, - force_tool_language: bool = True, - add_think_tool: bool = False, **kwargs, ): super().__init__( @@ -71,30 +57,10 @@ class ReMe(Application): ) C.initialize_service_context() - self.enable_identity_memory = enable_identity_memory - self.enable_tool_memory = enable_tool_memory - self.force_tool_language = force_tool_language - self.add_think_tool = add_think_tool - - self._personal_summarizer = PersonalSummarizer( - tools=[VectorRetrieveMemory(), AddMemory(), DeleteMemory(), UpdateMemory()], - ) - self._procedural_summarizer = ProceduralSummarizer( - tools=[VectorRetrieveMemory(), AddMemory(), DeleteMemory(), UpdateMemory()], - ) - hands_off_tool = HandsOffTool(memory_agents=[self._personal_summarizer, self._procedural_summarizer]) - self._reme_summarizer = ReMeSummarizer( - tools=[AddMetaMemory(), AddSummaryMemory(), hands_off_tool], - enable_identity_memory=self.enable_identity_memory, - enable_tool_memory=self.enable_tool_memory, - force_tool_language=self.force_tool_language, - add_think_tool=self.add_think_tool, - ) - self._reme_retriever = ReMeRetriever( - tools=[VectorRetrieveMemory(add_memory_type_target=True), ReadHistoryMemory()], - ) + self.llm: BaseLLM = C.get_llm("default") self.vector_store: BaseVectorStore = C.get_vector_store("default") + self.embedding_model: BaseEmbeddingModel = C.get_embedding_model("default") @staticmethod def _prepare_messages(messages: list[dict | Message], user_id: str, assistant_id: str): @@ -115,33 +81,61 @@ class ReMe(Application): description: str = "", user_id: str = "", assistant_id: str = "", - memory_mode: Literal["personal", "procedural", "auto"] = "personal", **kwargs, ): """Summarizes messages and stores them as memory based on the specified memory mode.""" - messages = self._prepare_messages(messages, user_id, assistant_id) - if memory_mode == "personal": - return await self._personal_summarizer.call( - messages=messages, - description=description, - memory_target=user_id, - **kwargs, + if user_id: + # halumem: user_id -> message.name + # locomo: add description + metadata_summary = { + "year": "The `year` information associated with the memory(Optional)", + "month": "The `month` information associated with the memory(Optional)", + "day": "The `day` information associated with the memory(Optional)", + "hour": "The `hour` information associated with the memory(Optional)", + # "year": "The year when the memory content occurred(Optional)", + # "month": "The month when the memory content occurred(Optional)", + # "day": "The day when the memory content occurred(Optional)", + # "hour": "The hour when the memory content occurred(Optional)", + } + meta_memories = [ + { + "memory_type": "personal", + "memory_target": user_id, + }, + ] + + messages = self._prepare_messages(messages, user_id, assistant_id) + + personal_summarizer = PersonalSummarizer( + tools=[ + VectorRetrieveMemory( + enable_summary_memory=False, + add_memory_type_target=False, + metadata_desc=None, + top_k=15, + ), + AddMemory(add_when_to_use=False, metadata_desc=metadata_summary), + DeleteMemory(), + UpdateMemory(add_when_to_use=False, metadata_desc=metadata_summary), + ], ) - elif memory_mode == "procedural": - return await self._procedural_summarizer.call( - messages=messages, - description=description, - memory_target=user_id, - **kwargs, + + reme_summarizer = ReMeSummarizer( + meta_memories=meta_memories, + enable_identity_memory=False, + tools=[ + AddMetaMemory(), + AddSummaryMemory(metadata_desc=metadata_summary), + HandsOffTool(memory_agents=[personal_summarizer]), + ], ) + + await reme_summarizer.call(messages=messages, description=description, **kwargs) + return reme_summarizer.memory_nodes + else: - return await self._reme_summarizer.call( - messages=messages, - description=description, - memory_target=user_id, - **kwargs, - ) + raise NotImplementedError async def retrieve( self, @@ -150,19 +144,42 @@ class ReMe(Application): description: str = "", user_id: str = "", assistant_id: str = "", - memory_mode: Literal["personal", "procedural", "auto"] = "personal", **kwargs, ): """Retrieves relevant memories based on the query and specified memory mode.""" - messages = self._prepare_messages(messages, user_id, assistant_id) - if memory_mode == "personal": - self._reme_retriever.meta_memories = [{"memory_type": "personal", "memory_target": user_id}] - return await self._reme_retriever.call(query=query, messages=messages, description=description, **kwargs) + if user_id: + messages = self._prepare_messages(messages, user_id, assistant_id) - elif memory_mode == "procedural": - self._reme_retriever.meta_memories = [{"memory_type": "procedural", "memory_target": user_id}] - return await self._reme_retriever.call(query=query, messages=messages, description=description, **kwargs) + metadata_retrieve = { + "year": "The year to filter memories(Optional)", + "month": "The month to filter memories(Optional)", + "day": "The day to filter memories(Optional)", + "hour": "The hour to filter memories(Optional)", + } + meta_memories = [ + { + "memory_type": "personal", + "memory_target": user_id, + }, + ] + + reme_retriever = ReMeRetriever( + meta_memories=meta_memories, + tools=[ + VectorRetrieveMemory( + enable_summary_memory=True, + add_memory_type_target=True, + metadata_desc=metadata_retrieve, + top_k=20, + ), + ReadHistoryMemory(), + ], + ) + + await reme_retriever.call(query=query, messages=messages, description=description, **kwargs) + + return reme_retriever.output else: - return await self._reme_retriever.call(query=query, messages=messages, description=description, **kwargs) + raise NotImplementedError