Files
group_fqcd_jr/app/service/memory_lifecycle_service.py
T

302 lines
12 KiB
Python
Raw Normal View History

"""客户级记忆生命周期:级联失效/删除、投影清理事件与审计。
为什么需要它(对应 P3 缺口 3):`MemoryService.invalidate` 只把**单条**记忆置为
invalidated,没有客户级联、没有投影(Milvus/Neo4j)清理、没有审计。客户销户、
撤回授权或合规删除要求"一次性把该客户的记忆与派生投影全部处理掉",这里补上这条链路。
级联范围(同一事务内完成):
1. `memory_unit`:按客户把 `status='active'`(或显式指定的状态集合)置为
`invalidated`(失效)或按 id 物理删除(删除)。历史版本行不是 active,
因此不受影响——失效语义是"停止被召回",而不是抹掉审计轨迹。
2. `memory_evidence`:失效时同时删除受影响记忆的证据行。基线里 `memory_evidence`
没有状态列(AGENTS.md 禁止改已有字段),只能在"保留"和"删除"之间选择;
证据是派生的支撑材料、不是被召回的内容,客户级失效要求其不再留存,
因此这里物理删除,并把删除行数记入审计。若需保留证据,应通过新增列的新迁移实现。
3. 投影清理:不直接调用 Milvus/Neo4j,而是在 outbox 写
`memory.invalidated` / `memory.deleted` 事件(`aggregate_type='memory_unit'`),
由投影消费者按 `memory_uuid` 清理。这样即使投影侧不可用,权威库的状态也已经正确。
4. `interaction_audit`:写一条客户级审计,`detail` 含幂等键与实际计数。
幂等:`detail.idempotency_key` 由范围(客户、状态集合、记忆 uuid 集合)派生;
重复调用时若审计已存在,直接返回首次结果,不再重复写事件与审计。
"""
import logging
from dataclasses import dataclass, field
from datetime import UTC, datetime
from typing import Any, Literal
from uuid import NAMESPACE_URL, uuid5
from sqlalchemy import delete, select, update
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.contracts import DomainEvent
from app.infrastructure.memory_cache import MemoryCacheAdapter
from app.model.audit import InteractionAudit
from app.model.memory import MemoryEvidence, MemoryUnit
from app.repository.outbox_repository import OutboxRepository
from app.service.memory_recall_service import MemoryRecallService
logger = logging.getLogger(__name__)
ACTION_TYPE = "memory.customer_lifecycle"
EVENT_INVALIDATED = "memory.invalidated"
EVENT_DELETED = "memory.deleted"
AGGREGATE_TYPE = "memory_unit"
DEFAULT_STATUSES = ("active",)
Mode = Literal["invalidate", "delete"]
@dataclass
class LifecycleResult:
"""一次客户级操作的结果;`idempotent_replay` 为真表示命中幂等边界。"""
customer_id: int
mode: str
memories: int = 0
evidences: int = 0
projection_events: int = 0
memory_uuids: list[str] = field(default_factory=list)
event_ids: list[str] = field(default_factory=list)
audit_id: int | None = None
idempotent_replay: bool = False
cache_keys_removed: int = 0
class MemoryLifecycleService:
"""客户级记忆级联失效/删除;可重复调用,重复调用不产生第二次副作用。"""
def __init__(
self,
session: AsyncSession,
*,
cache: MemoryCacheAdapter | None = None,
actor_id: int | None = None,
portal: str = "admin",
) -> None:
self.session = session
self.cache = cache
self.actor_id = actor_id
self.portal = portal
async def invalidate_customer(
self,
customer_id: int,
*,
memory_uuids: list[str] | None = None,
statuses: tuple[str, ...] = DEFAULT_STATUSES,
reason: str = "customer_lifecycle",
trace_id: str = "",
) -> LifecycleResult:
"""按客户级联失效:记忆置无效 + 证据删除 + 投影清理事件 + 审计。"""
return await self.run(
customer_id,
mode="invalidate",
memory_uuids=memory_uuids,
statuses=statuses,
reason=reason,
trace_id=trace_id,
)
async def delete_customer(
self,
customer_id: int,
*,
memory_uuids: list[str] | None = None,
reason: str = "customer_lifecycle",
trace_id: str = "",
) -> LifecycleResult:
"""按客户级联删除:记忆物理删除 + 证据删除 + 投影清理事件 + 审计。"""
return await self.run(
customer_id,
mode="delete",
memory_uuids=memory_uuids,
statuses=(),
reason=reason,
trace_id=trace_id,
)
async def run(
self,
customer_id: int,
*,
mode: Mode,
memory_uuids: list[str] | None = None,
statuses: tuple[str, ...] = DEFAULT_STATUSES,
reason: str = "customer_lifecycle",
trace_id: str = "",
) -> LifecycleResult:
selection = sorted(set(memory_uuids or []))
effective_statuses = tuple(sorted(set(statuses))) if mode == "invalidate" else ()
idempotency_key = self.idempotency_key(customer_id, mode, effective_statuses, selection)
replayed = await self._find_audit(customer_id, idempotency_key)
if replayed is not None:
return await self._replayed(customer_id, mode, replayed)
now = datetime.now(UTC).replace(tzinfo=None)
targets = await self._targets(customer_id, selection, effective_statuses)
uuids = [memory.memory_uuid for memory in targets]
ids = [int(memory.id) for memory in targets]
result = LifecycleResult(
customer_id=customer_id, mode=mode, memories=len(ids), memory_uuids=uuids
)
if ids:
result.evidences = await self._drop_evidence(ids)
if mode == "delete":
await self.session.execute(delete(MemoryUnit).where(MemoryUnit.id.in_(ids)))
else:
await self.session.execute(
update(MemoryUnit)
.where(MemoryUnit.id.in_(ids))
.values(status="invalidated", updated_at=now)
)
result.projection_events = await self._request_projection_cleanup(
customer_id, mode, uuids, reason=reason, trace_id=trace_id, now=now
)
audit = await self._audit(customer_id, mode, reason, idempotency_key, result, now)
await self.session.flush()
result.audit_id = int(audit.id) if audit.id is not None else None
result.cache_keys_removed = await self._drop_cache(customer_id)
return result
async def _targets(
self, customer_id: int, memory_uuids: list[str], statuses: tuple[str, ...]
) -> list[MemoryUnit]:
conditions: list[Any] = [MemoryUnit.customer_id == customer_id]
if memory_uuids:
conditions.append(MemoryUnit.memory_uuid.in_(memory_uuids))
if statuses:
conditions.append(MemoryUnit.status.in_(statuses))
found = await self.session.scalars(select(MemoryUnit).where(*conditions))
return list(found)
async def _drop_evidence(self, memory_ids: list[int]) -> int:
# ORM 批量删除不返回行数,先计数再删,计数进审计。
evidence_ids = list(
await self.session.scalars(
select(MemoryEvidence.id).where(MemoryEvidence.memory_id.in_(memory_ids))
)
)
if not evidence_ids:
return 0
await self.session.execute(
delete(MemoryEvidence).where(MemoryEvidence.id.in_(evidence_ids))
)
return len(evidence_ids)
async def _request_projection_cleanup(
self,
customer_id: int,
mode: Mode,
memory_uuids: list[str],
*,
reason: str,
trace_id: str,
now: datetime,
) -> int:
event_type = EVENT_DELETED if mode == "delete" else EVENT_INVALIDATED
repository = OutboxRepository(self.session)
created = 0
for memory_uuid in memory_uuids:
event = DomainEvent(
# 事件 id 由幂等键派生:同一记忆的同一操作最多产生一条清理事件。
event_id=str(uuid5(NAMESPACE_URL, f"jr:{event_type}:{memory_uuid}")),
event_type=event_type,
aggregate_type=AGGREGATE_TYPE,
aggregate_id=memory_uuid,
trace_id=trace_id or f"lifecycle:{customer_id}",
payload={
"customer_id": customer_id,
"memory_uuid": memory_uuid,
"reason": reason,
"operation": mode,
"actor_id": self.actor_id,
},
occurred_at=now,
)
await repository.append(event)
created += 1
return created
async def _find_audit(self, customer_id: int, idempotency_key: str) -> InteractionAudit | None:
found: InteractionAudit | None = await self.session.scalar(
select(InteractionAudit)
.where(
InteractionAudit.action_type == ACTION_TYPE,
InteractionAudit.target_customer_id == customer_id,
InteractionAudit.detail["idempotency_key"].as_string() == idempotency_key,
)
.order_by(InteractionAudit.id.desc())
.limit(1)
)
return found
async def _replayed(
self, customer_id: int, mode: str, existing: InteractionAudit
) -> LifecycleResult:
detail = dict(existing.detail or {})
logger.info("memory lifecycle replayed customer_id=%s audit_id=%s", customer_id,
existing.id)
return LifecycleResult(
customer_id=customer_id,
mode=str(detail.get("mode", mode)),
memories=int(detail.get("memories", 0)),
evidences=int(detail.get("evidences", 0)),
projection_events=int(detail.get("projection_events", 0)),
memory_uuids=[str(uuid) for uuid in detail.get("memory_uuids", [])],
event_ids=[],
audit_id=int(existing.id),
idempotent_replay=True,
cache_keys_removed=await self._drop_cache(customer_id),
)
async def _audit(
self,
customer_id: int,
mode: Mode,
reason: str,
idempotency_key: str,
result: LifecycleResult,
now: datetime,
) -> InteractionAudit:
audit = InteractionAudit(
actor_type="admin",
actor_id=self.actor_id,
target_customer_id=customer_id,
session_id=None,
portal=self.portal,
action_type=ACTION_TYPE,
detail={
"mode": mode,
"reason": reason,
"idempotency_key": idempotency_key,
"memories": result.memories,
"evidences": result.evidences,
"projection_events": result.projection_events,
"memory_uuids": result.memory_uuids,
"cascade": ["memory_unit", "memory_evidence", "projection", "interaction_audit"],
},
created_at=now,
)
self.session.add(audit)
await self.session.flush()
return audit
async def _drop_cache(self, customer_id: int) -> int:
if self.cache is None:
return 0
removed = await self.cache.delete(*self.cache_keys(customer_id))
return int(removed) if isinstance(removed, int) else 0
@staticmethod
def cache_keys(customer_id: int) -> list[str]:
"""与 `MemoryRecallService.cache_key` 完全对齐的客户级热缓存键集合。"""
return MemoryRecallService.cache_keys(customer_id)
@staticmethod
def idempotency_key(
customer_id: int, mode: str, statuses: tuple[str, ...], memory_uuids: list[str]
) -> str:
scope = ",".join(memory_uuids) or "*"
return f"memory.customer_lifecycle:{customer_id}:{mode}:{'|'.join(statuses)}:{scope}"