From 6e72bcac2ea00350c3ec3a0f33451fa12ca65c63 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 29 Jan 2026 14:18:53 +0800 Subject: [PATCH 1/3] refactor(memory): update memory registration and configuration handling --- docs/REME2_README.md | 11 ++++++++ pyproject.toml | 9 ++++++- reme/agent/memory/__init__.py | 9 +++++-- reme/core/application.py | 3 +++ reme/reme.py | 50 +++++++++++++++++------------------ reme/tool/memory/__init__.py | 1 - 6 files changed, 54 insertions(+), 29 deletions(-) create mode 100644 docs/REME2_README.md diff --git a/docs/REME2_README.md b/docs/REME2_README.md new file mode 100644 index 00000000..9f19c0e8 --- /dev/null +++ b/docs/REME2_README.md @@ -0,0 +1,11 @@ + + +# TODO +- [] halumem bench开发 +- [] default版本开发,for cli版本体验 +- [] cli开发 +- [] locomo bench开发 +- [] task memory迁移 +- [] mcp开发 +- [] reme外层接口完善 +- [] reme2 readme完善 \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml index 47aa478d..f7b62c4a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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/* diff --git a/reme/agent/memory/__init__.py b/reme/agent/memory/__init__.py index 024b68c4..73640f71 100644 --- a/reme/agent/memory/__init__.py +++ b/reme/agent/memory/__init__.py @@ -28,5 +28,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) diff --git a/reme/core/application.py b/reme/core/application.py index f96b1b03..2ff5d05e 100644 --- a/reme/core/application.py +++ b/reme/core/application.py @@ -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() diff --git a/reme/reme.py b/reme/reme.py index 5f587ea5..5620fc1e 100644 --- a/reme/reme.py +++ b/reme/reme.py @@ -50,9 +50,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, ): @@ -71,25 +71,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 @@ -172,30 +172,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) @@ -289,30 +289,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) diff --git a/reme/tool/memory/__init__.py b/reme/tool/memory/__init__.py index ff58d44f..2c99998a 100644 --- a/reme/tool/memory/__init__.py +++ b/reme/tool/memory/__init__.py @@ -43,6 +43,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) From c581dd38b7867154e2dcaf1253c27d8f49027a51 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 29 Jan 2026 15:10:58 +0800 Subject: [PATCH 2/3] feat(vector-store): add async collection reset functionality --- reme/core/context/service_context.py | 5 +++++ reme/core/vector_store/base_vector_store.py | 3 ++- 2 files changed, 7 insertions(+), 1 deletion(-) diff --git a/reme/core/context/service_context.py b/reme/core/context/service_context.py index 5be7b449..ac13d9d8 100644 --- a/reme/core/context/service_context.py +++ b/reme/core/context/service_context.py @@ -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(): diff --git a/reme/core/vector_store/base_vector_store.py b/reme/core/vector_store/base_vector_store.py index cc15101b..62a73a3e 100644 --- a/reme/core/vector_store/base_vector_store.py +++ b/reme/core/vector_store/base_vector_store.py @@ -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]: From 0ea51a9378033c77ce19d420686756d53943f403 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 29 Jan 2026 16:27:02 +0800 Subject: [PATCH 3/3] refactor(core): implement lazy initialization for Elasticsearch and Qdrant clients --- docs/REME2_README.md | 4 +- reme/core/vector_store/es_vector_store.py | 92 +++++++++++------- reme/core/vector_store/qdrant_vector_store.py | 93 +++++++++++++------ reme/reme.py | 25 ++++- .../tool/memory/profiles/read_all_profiles.py | 24 ++++- reme/tool/memory/profiles/update_profile.py | 76 ++++++++++----- 6 files changed, 221 insertions(+), 93 deletions(-) diff --git a/docs/REME2_README.md b/docs/REME2_README.md index 9f19c0e8..fd6174a6 100644 --- a/docs/REME2_README.md +++ b/docs/REME2_README.md @@ -8,4 +8,6 @@ - [] task memory迁移 - [] mcp开发 - [] reme外层接口完善 -- [] reme2 readme完善 \ No newline at end of file +- [] reme2 readme完善 +- [] 看日志,看要这个default版本怎么优化。 +- [] 学习Clawdbot记忆系统 \ No newline at end of file diff --git a/reme/core/vector_store/es_vector_store.py b/reme/core/vector_store/es_vector_store.py index 9a1019b8..1090d4b6 100644 --- a/reme/core/vector_store/es_vector_store.py +++ b/reme/core/vector_store/es_vector_store.py @@ -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") diff --git a/reme/core/vector_store/qdrant_vector_store.py b/reme/core/vector_store/qdrant_vector_store.py index 93ccee70..0afb37d7 100644 --- a/reme/core/vector_store/qdrant_vector_store.py +++ b/reme/core/vector_store/qdrant_vector_store.py @@ -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") diff --git a/reme/reme.py b/reme/reme.py index 5620fc1e..102b98dc 100644 --- a/reme/reme.py +++ b/reme/reme.py @@ -32,6 +32,7 @@ from .tool.memory import ( UpdateProfile, AddHistory, ReadAllProfiles, + AddMemory, ) @@ -141,12 +142,28 @@ class ReMe(Application): personal_summarizer = PersonalV1Summarizer( tools=[ AddDraftAndRetrieveSimilarMemory( - enable_thinking_params=enable_thinking_params, top_k=retrieve_top_k, + enable_thinking_params=enable_thinking_params, + 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, ), ], @@ -260,12 +277,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), ], diff --git a/reme/tool/memory/profiles/read_all_profiles.py b/reme/tool/memory/profiles/read_all_profiles.py index 892bbb58..1d7ddb1a 100644 --- a/reme/tool/memory/profiles/read_all_profiles.py +++ b/reme/tool/memory/profiles/read_all_profiles.py @@ -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." diff --git a/reme/tool/memory/profiles/update_profile.py b/reme/tool/memory/profiles/update_profile.py index 5edd990b..94e653aa 100644 --- a/reme/tool/memory/profiles/update_profile.py +++ b/reme/tool/memory/profiles/update_profile.py @@ -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 = []