Skip to content

Commit 48697db

Browse files
committed
🐛 fix(memory): deduplicate episodic/event_log on re-memorize and add foresight expiry cleanup
1 parent 8cb6d01 commit 48697db

8 files changed

Lines changed: 367 additions & 0 deletions

File tree

src/biz_layer/mem_cleanup.py

Lines changed: 82 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,82 @@
1+
"""
2+
Memory cleanup utilities.
3+
4+
Provides scheduled cleanup tasks for expired memory records.
5+
Currently handles foresight expiry: records whose validity window has passed
6+
are removed from all three stores (Milvus → Elasticsearch → MongoDB) in that
7+
order to minimise the window where a record is searchable but absent from the
8+
primary store.
9+
"""
10+
11+
from datetime import datetime
12+
from typing import Dict
13+
14+
from common_utils.datetime_utils import get_now_with_timezone
15+
from core.di import get_bean_by_type
16+
from core.observation.logger import get_logger
17+
from infra_layer.adapters.out.persistence.repository.foresight_record_repository import (
18+
ForesightRecordRawRepository,
19+
)
20+
from infra_layer.adapters.out.search.repository.foresight_es_repository import (
21+
ForesightEsRepository,
22+
)
23+
from infra_layer.adapters.out.search.repository.foresight_milvus_repository import (
24+
ForesightMilvusRepository,
25+
)
26+
27+
logger = get_logger(__name__)
28+
29+
30+
async def cleanup_expired_foresights(
31+
before: datetime | None = None,
32+
) -> Dict[str, int]:
33+
"""
34+
Delete foresight records that have passed their validity end time.
35+
36+
Deletion order: Milvus → Elasticsearch → MongoDB.
37+
This ensures that even if a later step fails, the record is no longer
38+
returned by vector or keyword search.
39+
40+
Args:
41+
before: Treat records with end_time < before as expired.
42+
Defaults to the current time when not provided.
43+
44+
Returns:
45+
Dict with keys ``milvus``, ``es``, ``mongo`` and the number of
46+
records deleted from each store.
47+
"""
48+
if before is None:
49+
before = get_now_with_timezone()
50+
51+
stats: Dict[str, int] = {"milvus": 0, "es": 0, "mongo": 0}
52+
53+
foresight_milvus_repo = get_bean_by_type(ForesightMilvusRepository)
54+
foresight_es_repo = get_bean_by_type(ForesightEsRepository)
55+
foresight_mongo_repo = get_bean_by_type(ForesightRecordRawRepository)
56+
57+
# Step 1: remove from Milvus (vector search)
58+
try:
59+
stats["milvus"] = await foresight_milvus_repo.delete_by_filters(end_time=before)
60+
except Exception as exc:
61+
logger.error("Failed to delete expired foresights from Milvus: %s", exc)
62+
63+
# Step 2: remove from Elasticsearch (keyword search)
64+
try:
65+
stats["es"] = await foresight_es_repo.delete_expired(before=before)
66+
except Exception as exc:
67+
logger.error("Failed to delete expired foresights from ES: %s", exc)
68+
69+
# Step 3: remove from MongoDB (primary store)
70+
try:
71+
stats["mongo"] = await foresight_mongo_repo.delete_expired(before=before)
72+
except Exception as exc:
73+
logger.error("Failed to delete expired foresights from MongoDB: %s", exc)
74+
75+
logger.info(
76+
"✅ Expired foresight cleanup complete (before=%s): milvus=%d es=%d mongo=%d",
77+
before.isoformat(),
78+
stats["milvus"],
79+
stats["es"],
80+
stats["mongo"],
81+
)
82+
return stats

src/biz_layer/mem_memorize.py

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -79,6 +79,12 @@
7979
from infra_layer.adapters.out.search.repository.episodic_memory_es_repository import (
8080
EpisodicMemoryEsRepository,
8181
)
82+
from infra_layer.adapters.out.search.repository.event_log_es_repository import (
83+
EventLogEsRepository,
84+
)
85+
from infra_layer.adapters.out.search.repository.event_log_milvus_repository import (
86+
EventLogMilvusRepository,
87+
)
8288
from biz_layer.mem_sync import MemorySyncService
8389
from core.context.context import get_current_app_info
8490

@@ -1122,6 +1128,19 @@ async def save_memory_docs(
11221128
saved_episodic: List[Any] = []
11231129

11241130
for doc in episodic_docs:
1131+
# Deduplicate: remove any existing records from the same source MemCell
1132+
# for this specific user before inserting the new one.
1133+
# Key is (parent_id, user_id) because one MemCell produces one episode
1134+
# per participant (personal) plus one group episode (user_id=None/"").
1135+
parent_id = getattr(doc, "parent_id", None)
1136+
user_id = getattr(doc, "user_id", None)
1137+
if parent_id:
1138+
await asyncio.gather(
1139+
episodic_repo.delete_by_parent_id(parent_id, user_id=user_id),
1140+
episodic_es_repo.delete_by_parent_id(parent_id, user_id=user_id),
1141+
episodic_milvus_repo.delete_by_parent_id(parent_id, user_id=user_id),
1142+
)
1143+
11251144
saved_doc = await episodic_repo.append_episodic_memory(doc)
11261145
saved_episodic.append(saved_doc)
11271146

@@ -1158,6 +1177,24 @@ async def save_memory_docs(
11581177
event_log_docs = grouped_docs.get(MemoryType.EVENT_LOG, [])
11591178
if event_log_docs:
11601179
event_log_repo = get_bean_by_type(EventLogRecordRawRepository)
1180+
event_log_es_repo = get_bean_by_type(EventLogEsRepository)
1181+
event_log_milvus_repo = get_bean_by_type(EventLogMilvusRepository)
1182+
1183+
# Deduplicate: collect unique (parent_id, user_id) pairs and delete old records
1184+
# before batch-inserting the new ones.
1185+
seen_parent_keys: set = set()
1186+
for doc in event_log_docs:
1187+
parent_id = getattr(doc, "parent_id", None)
1188+
user_id = getattr(doc, "user_id", None)
1189+
key = (parent_id, user_id)
1190+
if parent_id and key not in seen_parent_keys:
1191+
seen_parent_keys.add(key)
1192+
await asyncio.gather(
1193+
event_log_repo.delete_by_parent_id(parent_id),
1194+
event_log_es_repo.delete_by_parent_id(parent_id, user_id=user_id),
1195+
event_log_milvus_repo.delete_by_parent_id(parent_id),
1196+
)
1197+
11611198
saved_event_logs = await event_log_repo.create_batch(event_log_docs)
11621199
saved_result[MemoryType.EVENT_LOG] = saved_event_logs
11631200

src/infra_layer/adapters/out/persistence/repository/episodic_memory_raw_repository.py

Lines changed: 40 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -292,6 +292,46 @@ async def delete_by_event_id(
292292
)
293293
return False
294294

295+
async def delete_by_parent_id(
296+
self,
297+
parent_id: str,
298+
user_id: Optional[str] = None,
299+
session: Optional[AsyncClientSession] = None,
300+
) -> int:
301+
"""
302+
Delete episodic memories by parent MemCell ID, optionally scoped to a user.
303+
304+
Args:
305+
parent_id: Source MemCell ID (stored as parent_id field)
306+
user_id: If provided, only delete records belonging to this user.
307+
Pass None to delete all records (group + personal) for the parent.
308+
session: Optional MongoDB session for transaction support
309+
310+
Returns:
311+
Number of deleted records
312+
"""
313+
try:
314+
query_filter: Dict[str, Any] = {"parent_id": parent_id}
315+
if user_id is not None:
316+
query_filter["user_id"] = user_id
317+
318+
result = await self.model.find(query_filter, session=session).delete()
319+
count = result.deleted_count if result else 0
320+
logger.info(
321+
"✅ Deleted episodic memories by parent_id=%s user_id=%s: %d records",
322+
parent_id,
323+
user_id,
324+
count,
325+
)
326+
return count
327+
except Exception as e:
328+
logger.error(
329+
"❌ Failed to delete episodic memories by parent_id=%s: %s",
330+
parent_id,
331+
e,
332+
)
333+
return 0
334+
295335
async def delete_by_user_id(
296336
self, user_id: str, session: Optional[AsyncClientSession] = None
297337
) -> int:

src/infra_layer/adapters/out/persistence/repository/foresight_record_repository.py

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -266,6 +266,39 @@ async def find_by_filters(
266266
logger.error("❌ Failed to retrieve foresights: %s", e)
267267
return []
268268

269+
async def delete_expired(
270+
self,
271+
before: datetime,
272+
session: Optional[AsyncClientSession] = None,
273+
) -> int:
274+
"""
275+
Delete foresight records whose validity period ended before the given time.
276+
277+
Args:
278+
before: Delete foresights with end_time strictly before this datetime.
279+
end_time is stored as an ISO date string (YYYY-MM-DD).
280+
session: Optional MongoDB session for transaction support
281+
282+
Returns:
283+
Number of deleted records
284+
"""
285+
try:
286+
from common_utils.datetime_utils import to_date_str
287+
288+
before_str = to_date_str(before)
289+
query_filter: Dict[str, Any] = {"end_time": {"$lt": before_str}}
290+
result = await self.model.find(query_filter, session=session).delete()
291+
count = result.deleted_count if result else 0
292+
logger.info(
293+
"✅ Deleted expired foresights from MongoDB (before=%s): %d records",
294+
before_str,
295+
count,
296+
)
297+
return count
298+
except Exception as e:
299+
logger.error("❌ Failed to delete expired foresights from MongoDB: %s", e)
300+
return 0
301+
269302
async def delete_by_id(
270303
self, memory_id: str, session: Optional[AsyncClientSession] = None
271304
) -> bool:

src/infra_layer/adapters/out/search/repository/episodic_memory_es_repository.py

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -494,6 +494,50 @@ async def append_episodic_memory(
494494

495495
# ==================== Deletion functionality ====================
496496

497+
async def delete_by_parent_id(
498+
self,
499+
parent_id: str,
500+
user_id: Optional[str] = None,
501+
refresh: bool = False,
502+
) -> int:
503+
"""
504+
Delete episodic memory documents by parent MemCell ID.
505+
506+
Args:
507+
parent_id: Source MemCell ID
508+
user_id: If provided, only delete records for this user.
509+
refresh: Whether to refresh the index immediately
510+
511+
Returns:
512+
Number of deleted documents
513+
"""
514+
try:
515+
filter_queries: List[Dict[str, Any]] = [{"term": {"parent_id": parent_id}}]
516+
if user_id is not None:
517+
filter_queries.append({"term": {"user_id": user_id}})
518+
519+
delete_query = {"bool": {"must": filter_queries}}
520+
client = await self.get_client()
521+
index_name = self.get_index_name()
522+
response = await client.delete_by_query(
523+
index=index_name,
524+
body={"query": delete_query},
525+
refresh=refresh,
526+
)
527+
deleted_count = response.get("deleted", 0)
528+
logger.debug(
529+
"✅ Deleted episodic memory by parent_id=%s user_id=%s: %d records",
530+
parent_id,
531+
user_id,
532+
deleted_count,
533+
)
534+
return deleted_count
535+
except Exception as e:
536+
logger.error(
537+
"❌ Failed to delete episodic memory by parent_id=%s: %s", parent_id, e
538+
)
539+
raise
540+
497541
async def delete_by_event_id(self, event_id: str, refresh: bool = False) -> bool:
498542
"""
499543
Delete episodic memory document by event_id

src/infra_layer/adapters/out/search/repository/episodic_memory_milvus_repository.py

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -308,6 +308,45 @@ async def vector_search(
308308

309309
# ==================== Deletion Functionality ====================
310310

311+
async def delete_by_parent_id(
312+
self,
313+
parent_id: str,
314+
user_id: Optional[str] = None,
315+
) -> int:
316+
"""
317+
Delete episodic memory vectors by parent MemCell ID.
318+
319+
Args:
320+
parent_id: Source MemCell ID
321+
user_id: If provided, only delete records for this user.
322+
323+
Returns:
324+
Number of deleted records
325+
"""
326+
try:
327+
expr = f'parent_id == "{parent_id}"'
328+
if user_id is not None:
329+
expr += f' and user_id == "{user_id}"'
330+
331+
results = await self.collection.query(expr=expr, output_fields=["id"])
332+
delete_count = len(results)
333+
if delete_count > 0:
334+
await self.collection.delete(expr)
335+
logger.debug(
336+
"✅ Deleted episodic memory vectors by parent_id=%s user_id=%s: %d records",
337+
parent_id,
338+
user_id,
339+
delete_count,
340+
)
341+
return delete_count
342+
except Exception as e:
343+
logger.error(
344+
"❌ Failed to delete episodic memory vectors by parent_id=%s: %s",
345+
parent_id,
346+
e,
347+
)
348+
raise
349+
311350
async def delete_by_event_id(self, event_id: str) -> bool:
312351
"""
313352
Delete episodic memory document by event_id

src/infra_layer/adapters/out/search/repository/event_log_es_repository.py

Lines changed: 46 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -354,3 +354,49 @@ async def multi_search(
354354
e,
355355
)
356356
raise
357+
358+
# ==================== Deletion functionality ====================
359+
360+
async def delete_by_parent_id(
361+
self,
362+
parent_id: str,
363+
user_id: Optional[str] = None,
364+
refresh: bool = False,
365+
) -> int:
366+
"""
367+
Delete event log documents by parent memory ID.
368+
369+
Args:
370+
parent_id: Parent memory ID (MemCell or Episode ID)
371+
user_id: If provided, only delete records for this user.
372+
refresh: Whether to refresh the index immediately
373+
374+
Returns:
375+
Number of deleted documents
376+
"""
377+
try:
378+
filter_queries: List[Dict[str, Any]] = [{"term": {"parent_id": parent_id}}]
379+
if user_id is not None:
380+
filter_queries.append({"term": {"user_id": user_id}})
381+
382+
delete_query = {"bool": {"must": filter_queries}}
383+
client = await self.get_client()
384+
index_name = self.get_index_name()
385+
response = await client.delete_by_query(
386+
index=index_name,
387+
body={"query": delete_query},
388+
refresh=refresh,
389+
)
390+
deleted_count = response.get("deleted", 0)
391+
logger.debug(
392+
"✅ Deleted event logs by parent_id=%s user_id=%s: %d records",
393+
parent_id,
394+
user_id,
395+
deleted_count,
396+
)
397+
return deleted_count
398+
except Exception as e:
399+
logger.error(
400+
"❌ Failed to delete event logs by parent_id=%s: %s", parent_id, e
401+
)
402+
raise

0 commit comments

Comments
 (0)