Merge remote-tracking branch 'origin/main'

# Conflicts:
#	reme/reme.py
This commit is contained in:
方应 2026-01-30 16:33:32 +08:00
commit 0f90926781
12 changed files with 280 additions and 122 deletions

13
docs/REME2_README.md Normal file
View file

@ -0,0 +1,13 @@
# TODO
- [] halumem bench开发
- [] default版本开发,for cli版本体验
- [] cli开发
- [] locomo bench开发
- [] task memory迁移
- [] mcp开发
- [] reme外层接口完善
- [] reme2 readme完善
- [] 看日志,看要这个default版本怎么优化。
- [] 学习Clawdbot记忆系统

View file

@ -67,8 +67,14 @@ reme_ai = [
"**/*.json",
]
reme = [
"**/*.yaml",
"**/*.py",
"**/*.json",
]
[tool.setuptools.dynamic]
version = { attr = "reme_ai.__version__" }
version = { attr = "reme.__version__" }
[project.urls]
Homepage = "https://github.com/agentscope-ai/ReMe"
@ -77,5 +83,6 @@ Repository = "https://github.com/agentscope-ai/ReMe"
[project.scripts]
reme = "reme_ai.main:main"
reme2 = "reme.reme:main"
# python -m build && twine upload dist/*

View file

@ -32,5 +32,10 @@ __all__ = [
]
for name in __all__:
tool_class = globals()[name]
R.op.register()(tool_class)
agent_class = globals()[name]
if (
isinstance(agent_class, type)
and issubclass(agent_class, BaseMemoryAgent)
and agent_class is not BaseMemoryAgent
):
R.op.register()(agent_class)

View file

@ -106,4 +106,7 @@ class Application:
def run_service(self):
"""Run the configured service (HTTP, MCP, or CMD)."""
import warnings
warnings.filterwarnings("ignore", category=DeprecationWarning)
self.service_context.service.run()

View file

@ -121,6 +121,7 @@ class ServiceContext(BaseContext):
thread_pool=self.thread_pool,
**config.model_extra,
)
run_coro_safely(self.vector_stores[name].create_collection(config.collection_name))
# Initialize flow instances
self.flows: dict[str, BaseFlow] = {}
@ -189,6 +190,10 @@ class ServiceContext(BaseContext):
except Exception as e:
logger.exception(f"list_tool_calls: {server_name} error: {e}")
async def reset_default_collection(self, collection_name: str):
"""Reset the default vector store."""
await self.vector_stores["default"].reset_collection(collection_name)
async def close(self):
"""Close all service components asynchronously."""
for _, vector_store in self.vector_stores.items():

View file

@ -47,9 +47,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):
async def reset_collection(self, collection_name: str):
"""Change the name of the current collection."""
self.collection_name = collection_name
await self.create_collection(collection_name)
@abstractmethod
async def list_collections(self) -> list[str]:

View file

@ -69,19 +69,37 @@ class ESVectorStore(BaseVectorStore):
**kwargs,
)
# Initialize AsyncElasticsearch client
self.client = AsyncElasticsearch(
hosts=hosts,
cloud_id=cloud_id,
api_key=api_key,
basic_auth=basic_auth,
verify_certs=verify_certs,
headers=headers or {},
)
# Store connection parameters for lazy initialization
self.hosts = hosts
self.cloud_id = cloud_id
self.api_key = api_key
self.basic_auth = basic_auth
self.verify_certs = verify_certs
self.headers = headers or {}
self._client: AsyncElasticsearch | None = None
async def _get_client(self) -> AsyncElasticsearch:
"""Create or return the existing AsyncElasticsearch client.
This lazy initialization ensures the client is created in the correct event loop.
"""
if self._client is None:
self._client = AsyncElasticsearch(
hosts=self.hosts,
cloud_id=self.cloud_id,
api_key=self.api_key,
basic_auth=self.basic_auth,
verify_certs=self.verify_certs,
headers=self.headers,
)
logger.info("AsyncElasticsearch client initialized")
return self._client
async def list_collections(self) -> list[str]:
"""List all available index names in the Elasticsearch cluster."""
aliases = await self.client.indices.get_alias()
client = await self._get_client()
aliases = await client.indices.get_alias()
return list(aliases.keys())
async def create_collection(self, collection_name: str, **kwargs):
@ -92,8 +110,9 @@ class ESVectorStore(BaseVectorStore):
**kwargs: Settings like dimensions, similarity, shards, and replicas.
"""
collection_name = collection_name.lower()
client = await self._get_client()
if await self.client.indices.exists(index=collection_name):
if await client.indices.exists(index=collection_name):
return
dimensions = kwargs.get("dimensions", self.embedding_model.dimensions)
@ -125,8 +144,8 @@ class ESVectorStore(BaseVectorStore):
},
}
if not await self.client.indices.exists(index=collection_name):
await self.client.indices.create(index=collection_name, body=index_settings)
if not await client.indices.exists(index=collection_name):
await client.indices.create(index=collection_name, body=index_settings)
logger.info(f"Created index {collection_name} with dimensions={dimensions}")
else:
logger.info(f"Index {collection_name} already exists")
@ -139,9 +158,10 @@ class ESVectorStore(BaseVectorStore):
**kwargs: Additional parameters for the deletion request.
"""
collection_name = collection_name.lower()
client = await self._get_client()
if await self.client.indices.exists(index=collection_name):
await self.client.indices.delete(index=collection_name)
if await client.indices.exists(index=collection_name):
await client.indices.delete(index=collection_name)
logger.info(f"Deleted index {collection_name}")
else:
logger.warning(f"Index {collection_name} does not exist")
@ -154,8 +174,9 @@ class ESVectorStore(BaseVectorStore):
**kwargs: Additional parameters for the reindexing process.
"""
collection_name = collection_name.lower()
client = await self._get_client()
current_index = await self.client.indices.get(index=self.collection_name)
current_index = await client.indices.get(index=self.collection_name)
current_settings = current_index[self.collection_name]
settings_to_copy = current_settings.get("settings", {}).copy()
@ -174,7 +195,7 @@ class ESVectorStore(BaseVectorStore):
index_settings.pop(key, None)
settings_to_copy["index"] = index_settings
await self.client.indices.create(
await client.indices.create(
index=collection_name,
body={
"settings": settings_to_copy,
@ -182,7 +203,7 @@ class ESVectorStore(BaseVectorStore):
},
)
await self.client.reindex(
await client.reindex(
body={
"source": {"index": self.collection_name},
"dest": {"index": collection_name},
@ -224,7 +245,8 @@ class ESVectorStore(BaseVectorStore):
}
actions.append(action)
success, failed = await async_bulk(self.client, actions, raise_on_error=False)
client = await self._get_client()
success, failed = await async_bulk(client, actions, raise_on_error=False)
if failed:
logger.warning(f"Failed to insert {len(failed)} documents")
@ -232,7 +254,7 @@ class ESVectorStore(BaseVectorStore):
logger.info(f"Inserted {success} documents into {self.collection_name}")
if refresh:
await self.client.indices.refresh(index=self.collection_name)
await client.indices.refresh(index=self.collection_name)
async def search(
self,
@ -286,7 +308,8 @@ class ESVectorStore(BaseVectorStore):
filter_conditions.append({"term": {f"metadata.{key}": value}})
search_query["knn"]["filter"] = {"bool": {"must": filter_conditions}}
response = await self.client.search(index=self.collection_name, body=search_query)
client = await self._get_client()
response = await client.search(index=self.collection_name, body=search_query)
results = []
for hit in response["hits"]["hits"]:
@ -323,8 +346,9 @@ class ESVectorStore(BaseVectorStore):
},
)
client = await self._get_client()
success, failed = await async_bulk(
self.client,
client,
actions,
raise_on_error=False,
raise_on_exception=False,
@ -336,7 +360,7 @@ class ESVectorStore(BaseVectorStore):
logger.info(f"Deleted {success} documents from {self.collection_name}")
if refresh:
await self.client.indices.refresh(index=self.collection_name)
await client.indices.refresh(index=self.collection_name)
async def delete_all(self, **kwargs):
"""Remove all vectors from the collection.
@ -344,7 +368,8 @@ class ESVectorStore(BaseVectorStore):
Args:
**kwargs: Additional deletion parameters.
"""
response = await self.client.delete_by_query(
client = await self._get_client()
response = await client.delete_by_query(
index=self.collection_name,
body={"query": {"match_all": {}}},
)
@ -354,7 +379,7 @@ class ESVectorStore(BaseVectorStore):
refresh = kwargs.get("refresh", True)
if refresh:
await self.client.indices.refresh(index=self.collection_name)
await client.indices.refresh(index=self.collection_name)
async def update(self, nodes: VectorNode | list[VectorNode], refresh: bool = True, **kwargs):
"""Update existing documents with new content or metadata.
@ -394,8 +419,9 @@ class ESVectorStore(BaseVectorStore):
},
)
client = await self._get_client()
success, failed = await async_bulk(
self.client,
client,
actions,
raise_on_error=False,
raise_on_exception=False,
@ -407,7 +433,7 @@ class ESVectorStore(BaseVectorStore):
logger.info(f"Updated {success} documents in {self.collection_name}")
if refresh:
await self.client.indices.refresh(index=self.collection_name)
await client.indices.refresh(index=self.collection_name)
async def get(self, vector_ids: str | list[str]) -> VectorNode | list[VectorNode]:
"""Fetch documents by their IDs from the current index.
@ -422,7 +448,8 @@ class ESVectorStore(BaseVectorStore):
if single_result:
vector_ids = [vector_ids]
response = await self.client.mget(
client = await self._get_client()
response = await client.mget(
index=self.collection_name,
body={"ids": vector_ids},
)
@ -499,7 +526,8 @@ class ESVectorStore(BaseVectorStore):
else:
query["size"] = 10000
response = await self.client.search(index=self.collection_name, body=query)
client = await self._get_client()
response = await client.search(index=self.collection_name, body=query)
results = []
for hit in response["hits"]["hits"]:
@ -522,5 +550,7 @@ class ESVectorStore(BaseVectorStore):
async def close(self):
"""Terminate the Elasticsearch client session and release resources."""
await self.client.close()
logger.info("Elasticsearch client connection closed")
if self._client is not None:
await self._client.close()
self._client = None
logger.info("Elasticsearch client connection closed")

View file

@ -86,19 +86,17 @@ class QdrantVectorStore(BaseVectorStore):
**kwargs,
)
client_kwargs = {k: v for k, v in kwargs.items() if k != "thread_pool"}
self.client = AsyncQdrantClient(
host=host,
port=port,
path=path,
url=url,
api_key=api_key,
https=https,
grpc_port=grpc_port,
prefer_grpc=prefer_grpc,
**client_kwargs,
)
# Store connection parameters for lazy initialization
self.host = host
self.port = port
self.path = path
self.url = url
self.api_key = api_key
self.https = https
self.grpc_port = grpc_port
self.prefer_grpc = prefer_grpc
self.client_kwargs = {k: v for k, v in kwargs.items() if k != "thread_pool"}
self._client: AsyncQdrantClient | None = None
self.is_local = path is not None
distance_map = {
@ -109,9 +107,31 @@ class QdrantVectorStore(BaseVectorStore):
self.distance = distance_map.get(distance.lower(), Distance.COSINE)
self.on_disk = on_disk
async def _get_client(self) -> AsyncQdrantClient:
"""Create or return the existing AsyncQdrantClient.
This lazy initialization ensures the client is created in the correct event loop.
"""
if self._client is None:
self._client = AsyncQdrantClient(
host=self.host,
port=self.port,
path=self.path,
url=self.url,
api_key=self.api_key,
https=self.https,
grpc_port=self.grpc_port,
prefer_grpc=self.prefer_grpc,
**self.client_kwargs,
)
logger.info("AsyncQdrantClient initialized")
return self._client
async def list_collections(self) -> list[str]:
"""Retrieve names of all existing collections in the Qdrant instance."""
collections = await self.client.get_collections()
client = await self._get_client()
collections = await client.get_collections()
return [collection.name for collection in collections.collections]
async def create_collection(self, collection_name: str, **kwargs: Any):
@ -130,7 +150,8 @@ class QdrantVectorStore(BaseVectorStore):
distance = kwargs.get("distance", self.distance)
on_disk = kwargs.get("on_disk", self.on_disk)
await self.client.create_collection(
client = await self._get_client()
await client.create_collection(
collection_name=collection_name,
vectors_config=VectorParams(
size=dimensions,
@ -147,10 +168,11 @@ class QdrantVectorStore(BaseVectorStore):
async def _create_payload_indexes(self, collection_name: str):
"""Create keyword indexes for common metadata fields to optimize filtering."""
common_fields = ["user_id", "agent_id", "run_id", "actor_id", "source"]
client = await self._get_client()
for field in common_fields:
try:
await self.client.create_payload_index(
await client.create_payload_index(
collection_name=collection_name,
field_name=field,
field_schema="keyword",
@ -163,16 +185,18 @@ class QdrantVectorStore(BaseVectorStore):
"""Permanently remove a collection from the Qdrant instance."""
collections = await self.list_collections()
if collection_name in collections:
await self.client.delete_collection(collection_name=collection_name)
client = await self._get_client()
await client.delete_collection(collection_name=collection_name)
logger.info(f"Deleted collection {collection_name}")
else:
logger.warning(f"Collection {collection_name} does not exist")
async def copy_collection(self, collection_name: str, **kwargs: Any):
"""Duplicate an existing collection to a new one including all data."""
collection_info = await self.client.get_collection(collection_name=self.collection_name)
client = await self._get_client()
collection_info = await client.get_collection(collection_name=self.collection_name)
await self.client.create_collection(
await client.create_collection(
collection_name=collection_name,
vectors_config=collection_info.config.params.vectors,
)
@ -181,7 +205,7 @@ class QdrantVectorStore(BaseVectorStore):
batch_size = 100
while True:
records, next_offset = await self.client.scroll(
records, next_offset = await client.scroll(
collection_name=self.collection_name,
limit=batch_size,
offset=offset,
@ -201,7 +225,7 @@ class QdrantVectorStore(BaseVectorStore):
for record in records
]
await self.client.upsert(
await client.upsert(
collection_name=collection_name,
points=points,
)
@ -244,7 +268,8 @@ class QdrantVectorStore(BaseVectorStore):
points.append(point)
wait = kwargs.get("wait", True)
await self.client.upsert(
client = await self._get_client()
await client.upsert(
collection_name=self.collection_name,
points=points,
wait=wait,
@ -333,7 +358,8 @@ class QdrantVectorStore(BaseVectorStore):
query_filter = self._create_filter(filters) if filters else None
score_threshold = kwargs.get("score_threshold", None)
results = await self.client.query_points(
client = await self._get_client()
results = await client.query_points(
collection_name=self.collection_name,
query=query_vector,
query_filter=query_filter,
@ -369,7 +395,8 @@ class QdrantVectorStore(BaseVectorStore):
point_ids.append(point_id)
wait = kwargs.get("wait", True)
await self.client.delete(
client = await self._get_client()
await client.delete(
collection_name=self.collection_name,
points_selector=PointIdsList(points=point_ids),
wait=wait,
@ -384,7 +411,8 @@ class QdrantVectorStore(BaseVectorStore):
# Delete all points by using an empty filter (matches all)
from qdrant_client.models import FilterSelector
await self.client.delete(
client = await self._get_client()
await client.delete(
collection_name=self.collection_name,
points_selector=FilterSelector(filter=Filter(must=[])),
wait=wait,
@ -424,7 +452,8 @@ class QdrantVectorStore(BaseVectorStore):
points.append(point)
wait = kwargs.get("wait", True)
await self.client.upsert(
client = await self._get_client()
await client.upsert(
collection_name=self.collection_name,
points=points,
wait=wait,
@ -446,7 +475,8 @@ class QdrantVectorStore(BaseVectorStore):
point_id = abs(hash(vector_id)) % (10**18)
point_ids.append(point_id)
points = await self.client.retrieve(
client = await self._get_client()
points = await client.retrieve(
collection_name=self.collection_name,
ids=point_ids,
with_payload=True,
@ -489,7 +519,8 @@ class QdrantVectorStore(BaseVectorStore):
# If sorting is needed, fetch more records than the limit to ensure correct sorting
fetch_limit = 10000 if sort_key else (limit or 10000)
records, _ = await self.client.scroll(
client = await self._get_client()
records, _ = await client.scroll(
collection_name=self.collection_name,
scroll_filter=scroll_filter,
limit=fetch_limit,
@ -528,5 +559,7 @@ class QdrantVectorStore(BaseVectorStore):
async def close(self):
"""Close the AsyncQdrantClient connection and release resources."""
await self.client.close()
logger.info("Qdrant client connection closed")
if self._client is not None:
await self._client.close()
self._client = None
logger.info("Qdrant client connection closed")

View file

@ -37,6 +37,7 @@ from .tool.memory import (
AddProfile,
AddHistory,
ReadAllProfiles,
AddMemory,
)
@ -55,9 +56,9 @@ class ReMe(Application):
embedding_model: dict | None = None,
vector_store: dict | None = None,
token_counter: dict | None = None,
personal_memory_target: list[str] | None = None,
procedural_memory_target: list[str] | None = None,
tool_memory_target: list[str] | None = None,
target_user_names: list[str] | None = None,
target_task_names: list[str] | None = None,
target_tool_names: list[str] | None = None,
profile_dir: str = "reme_profile",
**kwargs,
):
@ -76,25 +77,25 @@ class ReMe(Application):
**kwargs,
)
memory_target_type_mapping: dict[str, MemoryType] = {}
if personal_memory_target:
for name in personal_memory_target:
assert name not in memory_target_type_mapping, f"Memory target name {name} is already used."
if target_user_names:
for name in target_user_names:
assert name not in memory_target_type_mapping, f"target_user_names={name} is already used."
memory_target_type_mapping[name] = MemoryType.PERSONAL
if procedural_memory_target:
for name in procedural_memory_target:
assert name not in memory_target_type_mapping, f"Memory target name {name} is already used."
if target_task_names:
for name in target_task_names:
assert name not in memory_target_type_mapping, f"target_task_names={name} is already used."
memory_target_type_mapping[name] = MemoryType.PROCEDURAL
if tool_memory_target:
for name in tool_memory_target:
assert name not in memory_target_type_mapping, f"Memory target name {name} is already used."
if target_tool_names:
for name in target_tool_names:
assert name not in memory_target_type_mapping, f"target_tool_names={name} is already used."
memory_target_type_mapping[name] = MemoryType.TOOL
self.service_context.memory_target_type_mapping = memory_target_type_mapping
self.profile_dir: str = profile_dir
def add_meta_memory(self, memory_type: str | MemoryType, memory_target: str):
def _add_meta_memory(self, memory_type: str | MemoryType, memory_target: str):
"""Register or validate a memory target with the given memory type."""
if memory_target in self.service_context.memory_target_type_mapping:
assert self.service_context.memory_target_type_mapping[memory_target] is memory_type
@ -148,11 +149,26 @@ class ReMe(Application):
tools=[
AddAndRetrieveSimilarMemory(
enable_thinking_params=enable_thinking_params,
top_k=retrieve_top_k,
enable_memory_target=False,
enable_when_to_use=False,
enable_multiple=True,
),
AddMemory(
enable_thinking_params=enable_thinking_params,
enable_memory_target=False,
enable_when_to_use=False,
enable_multiple=True,
),
UpdateMemoryV2(enable_thinking_params=enable_thinking_params),
AddDraftAndReadAllProfiles(
enable_thinking_params=enable_thinking_params,
enable_memory_target=False,
enable_multiple=True,
profile_dir=self.profile_dir,
),
UpdateProfile(
enable_thinking_params=enable_thinking_params,
enable_memory_target=False,
enable_multiple=True,
profile_dir=self.profile_dir,
),
],
@ -217,30 +233,30 @@ class ReMe(Application):
if isinstance(user_name, str):
for message in format_messages:
message.name = user_name
self.add_meta_memory(MemoryType.PERSONAL, user_name)
self._add_meta_memory(MemoryType.PERSONAL, user_name)
elif isinstance(user_name, list):
for name in user_name:
self.add_meta_memory(MemoryType.PERSONAL, name)
self._add_meta_memory(MemoryType.PERSONAL, name)
else:
raise RuntimeError("user_name must be str or list[str]")
memory_agents.append(personal_summarizer)
if task_name:
if isinstance(task_name, str):
self.add_meta_memory(MemoryType.PROCEDURAL, task_name)
self._add_meta_memory(MemoryType.PROCEDURAL, task_name)
elif isinstance(task_name, list):
for name in task_name:
self.add_meta_memory(MemoryType.PROCEDURAL, name)
self._add_meta_memory(MemoryType.PROCEDURAL, name)
else:
raise RuntimeError("task_name must be str or list[str]")
memory_agents.append(procedural_summarizer)
if tool_name:
if isinstance(tool_name, str):
self.add_meta_memory(MemoryType.TOOL, tool_name)
self._add_meta_memory(MemoryType.TOOL, tool_name)
elif isinstance(tool_name, list):
for name in tool_name:
self.add_meta_memory(MemoryType.TOOL, name)
self._add_meta_memory(MemoryType.TOOL, name)
else:
raise RuntimeError("tool_name must be str or list[str]")
memory_agents.append(tool_summarizer)
@ -305,12 +321,14 @@ class ReMe(Application):
tools=[
ReadAllProfiles(
enable_thinking_params=enable_thinking_params,
enable_memory_target=False,
profile_dir=self.profile_dir,
),
RetrieveMemory(
enable_thinking_params=enable_thinking_params,
top_k=retrieve_top_k,
enable_thinking_params=enable_thinking_params,
enable_time_filter=enable_time_filter,
enable_multiple=True
),
ReadHistory(enable_thinking_params=enable_thinking_params),
],
@ -348,30 +366,30 @@ class ReMe(Application):
memory_agents = []
if user_name:
if isinstance(user_name, str):
self.add_meta_memory(MemoryType.PERSONAL, user_name)
self._add_meta_memory(MemoryType.PERSONAL, user_name)
elif isinstance(user_name, list):
for name in user_name:
self.add_meta_memory(MemoryType.PERSONAL, name)
self._add_meta_memory(MemoryType.PERSONAL, name)
else:
raise RuntimeError("user_name must be str or list[str]")
memory_agents.append(personal_retriever)
if task_name:
if isinstance(task_name, str):
self.add_meta_memory(MemoryType.PROCEDURAL, task_name)
self._add_meta_memory(MemoryType.PROCEDURAL, task_name)
elif isinstance(task_name, list):
for name in task_name:
self.add_meta_memory(MemoryType.PROCEDURAL, name)
self._add_meta_memory(MemoryType.PROCEDURAL, name)
else:
raise RuntimeError("task_name must be str or list[str]")
memory_agents.append(procedural_retriever)
if tool_name:
if isinstance(tool_name, str):
self.add_meta_memory(MemoryType.TOOL, tool_name)
self._add_meta_memory(MemoryType.TOOL, tool_name)
elif isinstance(tool_name, list):
for name in tool_name:
self.add_meta_memory(MemoryType.TOOL, name)
self._add_meta_memory(MemoryType.TOOL, name)
else:
raise RuntimeError("tool_name must be str or list[str]")
memory_agents.append(tool_retriever)

View file

@ -48,6 +48,5 @@ __all__ = [
for name in __all__:
tool_class = globals()[name]
# Only register classes that inherit from BaseMemoryTool
if isinstance(tool_class, type) and issubclass(tool_class, BaseMemoryTool) and tool_class is not BaseMemoryTool:
R.op.register()(tool_class)

View file

@ -10,25 +10,41 @@ from ....core.schema import ToolCall
class ReadAllProfiles(BaseMemoryTool):
"""Tool to read all user profiles"""
def __init__(self, **kwargs):
def __init__(self, enable_memory_target: bool = False, **kwargs):
kwargs["enable_multiple"] = False
super().__init__(**kwargs)
self.enable_memory_target: bool = enable_memory_target
def _build_tool_call(self) -> ToolCall:
"""Build and return the tool call schema"""
properties = {}
required = []
if self.enable_memory_target:
properties["memory_target"] = {
"type": "string",
"description": "memory_target",
}
required.append("memory_target")
return ToolCall(
**{
"description": "Read all user profiles.",
"parameters": {
"type": "object",
"properties": {},
"required": [],
"properties": properties,
"required": required,
},
},
)
async def execute(self):
profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=self.memory_target)
if self.enable_memory_target:
target = self.context.get("memory_target")
else:
target = self.memory_target
profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=target)
profiles_str = profile_handler.read_all(add_profile_id=True)
if not profiles_str:
output = "No profiles found."

View file

@ -10,12 +10,36 @@ from ....core.schema import ToolCall
class UpdateProfile(BaseMemoryTool):
"""Tool to update user profile by adding or removing profile entries"""
def __init__(self, **kwargs):
def __init__(self, enable_memory_target: bool = False, **kwargs):
kwargs["enable_multiple"] = True
super().__init__(**kwargs)
self.enable_memory_target: bool = enable_memory_target
def _build_multiple_tool_call(self) -> ToolCall:
"""Build and return the multiple tool call schema"""
profile_properties = {
"message_time": {
"type": "string",
"description": "Message time, e.g. '2020-01-01 00:00:00'",
},
"profile_key": {
"type": "string",
"description": "Profile key or category, e.g. 'name'",
},
"profile_value": {
"type": "string",
"description": "Profile value or content, e.g. 'John Smith'",
},
}
profile_required = ["message_time", "profile_key", "profile_value"]
if self.enable_memory_target:
profile_properties["memory_target"] = {
"type": "string",
"description": "memory_target",
}
profile_required.append("memory_target")
return ToolCall(
**{
"description": "update user profile by removing and adding profile entries.",
@ -34,21 +58,8 @@ class UpdateProfile(BaseMemoryTool):
"description": "List of profiles to add",
"items": {
"type": "object",
"properties": {
"message_time": {
"type": "string",
"description": "Message time, e.g. '2020-01-01 00:00:00'",
},
"profile_key": {
"type": "string",
"description": "Profile key or category, e.g. 'name'",
},
"profile_value": {
"type": "string",
"description": "Profile value or content, e.g. 'John Smith'",
},
},
"required": ["message_time", "profile_key", "profile_value"],
"properties": profile_properties,
"required": profile_required,
},
},
},
@ -58,8 +69,6 @@ class UpdateProfile(BaseMemoryTool):
)
async def execute(self):
profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=self.memory_target)
# Get parameters
profile_ids_to_delete = self.context.get("profile_ids_to_delete", [])
profile_ids_to_delete = sorted({pid for pid in profile_ids_to_delete if pid})
@ -68,17 +77,36 @@ class UpdateProfile(BaseMemoryTool):
if not profile_ids_to_delete and not profiles_to_add:
return "No profiles to remove or add, operation completed."
# Delete profiles using ProfileHandler (batch mode)
removed_count = 0
added_count = 0
# Delete profiles (using self.memory_target)
if profile_ids_to_delete:
profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=self.memory_target)
removed_count = profile_handler.delete(profile_ids_to_delete)
# Add new profiles using ProfileHandler (batch mode)
added_count = 0
# Add new profiles
if profiles_to_add:
new_nodes = profile_handler.add_batch(profiles=profiles_to_add, ref_memory_id=self.history_id)
self.memory_nodes.extend(new_nodes)
added_count = len(new_nodes)
if self.enable_memory_target:
# Group profiles by memory_target
from collections import defaultdict
profiles_by_target = defaultdict(list)
for profile in profiles_to_add:
target = profile.get("memory_target", self.memory_target)
profiles_by_target[target].append(profile)
# Add profiles for each target
for target, target_profiles in profiles_by_target.items():
profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=target)
new_nodes = profile_handler.add_batch(profiles=target_profiles, ref_memory_id=self.history_id)
self.memory_nodes.extend(new_nodes)
added_count += len(new_nodes)
else:
# Use self.memory_target for all profiles
profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=self.memory_target)
new_nodes = profile_handler.add_batch(profiles=profiles_to_add, ref_memory_id=self.history_id)
self.memory_nodes.extend(new_nodes)
added_count = len(new_nodes)
# Build output message
operations = []