mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-07 03:00:27 +00:00
Merge remote-tracking branch 'origin/main'
# Conflicts: # reme/reme.py
This commit is contained in:
commit
0f90926781
12 changed files with 280 additions and 122 deletions
13
docs/REME2_README.md
Normal file
13
docs/REME2_README.md
Normal file
|
|
@ -0,0 +1,13 @@
|
|||
|
||||
|
||||
# TODO
|
||||
- [] halumem bench开发
|
||||
- [] default版本开发,for cli版本体验
|
||||
- [] cli开发
|
||||
- [] locomo bench开发
|
||||
- [] task memory迁移
|
||||
- [] mcp开发
|
||||
- [] reme外层接口完善
|
||||
- [] reme2 readme完善
|
||||
- [] 看日志,看要这个default版本怎么优化。
|
||||
- [] 学习Clawdbot记忆系统
|
||||
|
|
@ -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/*
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
74
reme/reme.py
74
reme/reme.py
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue