mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-11 03:40:03 +00:00
feat(core): implement memory node tracking and embedding text truncation
This commit is contained in:
parent
d91fc6a14c
commit
0b7843f557
31 changed files with 714 additions and 208 deletions
5
.gitignore
vendored
5
.gitignore
vendored
|
|
@ -32,4 +32,7 @@ site/*
|
|||
docs/_build/*
|
||||
test_compact_storage/*
|
||||
test_working_memory/*
|
||||
*.code-workspace
|
||||
*.code-workspace
|
||||
local_vector_store/*
|
||||
bench_results/*
|
||||
meta_memory/*
|
||||
329
bench/eval_reme.py
Normal file
329
bench/eval_reme.py
Normal file
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
|
|
@ -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 `<NO_RELEVANT_MEMORY>`.
|
||||
|
||||
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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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'("- <memory_type>\(<memory_target>\): <description>"\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,
|
||||
|
|
|
|||
|
|
@ -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 `<memory_type>(<memory_target>)` entry:
|
||||
- If the required `<memory_type>(<memory_target>)` 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=<person's name>`.
|
||||
- For procedural memories: specify `memory_type="procedural"` and `memory_target=<topic or domain>`.
|
||||
- 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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.")
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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 <memory_type>(<memory_target>) 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 <memory_type>(<memory_target>) 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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
155
reme_ai/reme.py
155
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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue