Files
group_fqcd_jr/app/service/memory_service.py
T

312 lines
13 KiB
Python

import logging
from collections.abc import Iterable
from contextlib import suppress
from datetime import UTC, datetime
from math import exp
from typing import Any, Protocol
from uuid import uuid4
from sqlalchemy import select
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession
from app.model.memory import MemoryConflict, MemoryEvidence, MemoryUnit
from app.service.memory_taxonomy import BUSINESS_EVENT_TYPES, detect_memory_signals
logger = logging.getLogger(__name__)
class CacheDeleteAdapter(Protocol):
"""写入路径需要的最小缓存能力:按键删除。`MemoryCacheAdapter` 天然满足。"""
async def delete(self, *keys: str) -> int: ...
class MemoryService:
"""MySQL authoritative memory operations; vector/graph stores are projections."""
# 证据摘录入库长度上限;库列 evidence_excerpt 为 TEXT,正文可长于该值。
EXCERPT_LIMIT = 2000
def __init__(
self, session: AsyncSession, *, cache: CacheDeleteAdapter | None = None
) -> None:
self.session = session
# 召回热缓存由调用方注入:写入路径不自己造 Redis 客户端,也不反向依赖
# 召回服务(缓存只是可重建的优化层,注入失败等于没有缓存)。
self.cache = cache
@staticmethod
def detect_memory_signals(content: str) -> tuple[str, ...]:
"""识别消息里明确陈述的持久事实/偏好(受控键),供触发判定与调用点复用。"""
return detect_memory_signals(content)
@staticmethod
def should_extract_memory(
*,
conversation_content: str,
role: str,
tool_result: bool = False,
event_type: str | None = None,
signals: Iterable[str] = (),
) -> bool:
"""只记持久事实/偏好:必须存在显式信号,长度不再是门槛。
触发条件(任一满足):
1. 业务事件本身即持久事实(`BUSINESS_EVENT_TYPES`,如 risk.assessment_completed、
trade.completed);
2. 工具产生了权威业务事实(`tool_result=True`);
3. 用户消息命中受控信号(显式陈述的风险偏好、投资期限、流动性约束、职业、
家庭状况、目标等),调用方可用 `signals` 直接传入已识别的受控键。
普通问答即使很长也不触发;"只买货币基金"这类两三个字的明确陈述会触发。
"""
if not conversation_content.strip():
return False
if tool_result or event_type in BUSINESS_EVENT_TYPES:
return True
if role != "user":
return False
if signals:
return True
return bool(detect_memory_signals(conversation_content))
async def recall(self, customer_id: int, *, limit: int = 10) -> list[MemoryUnit]:
now = datetime.now(UTC).replace(tzinfo=None)
result = await self.session.scalars(
select(MemoryUnit)
.where(
MemoryUnit.customer_id == customer_id,
MemoryUnit.status == "active",
(MemoryUnit.valid_until.is_(None) | (MemoryUnit.valid_until > now)),
)
.order_by(MemoryUnit.confidence.desc(), MemoryUnit.updated_at.desc())
.limit(max(1, min(limit, 100)))
)
return list(result)
async def recall_with_decay(
self, customer_id: int, query: str | None = None, *, limit: int = 10
) -> list[MemoryUnit]:
memories = await self.recall(customer_id, limit=100)
now = datetime.now(UTC).replace(tzinfo=None)
if query:
terms = {term.lower() for term in query.split() if term}
memories = [
memory
for memory in memories
if not terms
or any(term in memory.content.lower() for term in terms)
or any(term in memory.memory_key.lower() for term in terms)
]
memories.sort(
# confidence 在库中是 DECIMAL(5,4),驱动返回 Decimal,与浮点时间衰减因子相乘
# 会抛 TypeError(该路径此前未被真实数据触发)。排序只需相对大小,转 float。
key=lambda memory: float(memory.confidence)
* exp(-max(0, (now - memory.updated_at).days) / 365),
reverse=True,
)
return memories[: max(1, min(limit, 100))]
async def upsert(
self,
customer_id: int,
memory_key: str,
content: str,
*,
memory_type: str = "fact",
confidence: float = 0.5,
source_type: str = "conversation",
structured_value: dict[str, Any] | None = None,
status: str = "active",
) -> MemoryUnit:
"""按 (customer_id, active_memory_key) 语义更新唯一有效记忆。
`content` 必须是抽取后的结构化短语,`structured_value` 保存同一份结构化结果;
内容变化时记录一条冲突:左侧为被覆盖的旧值所在记忆行,右侧为该记忆的新版本
标识(见 `_conflict_right_id`)。两侧绝不指向同一条记录,避免自引用冲突。
"""
if status not in {"active", "candidate"}:
raise ValueError("status must be active or candidate")
now = datetime.now(UTC).replace(tzinfo=None)
# 候选记录不能覆盖现有有效记忆,必须等待确认或审核后再晋升。
memory = await self._active(customer_id, memory_key) if status == "active" else None
if memory is not None:
updated = await self._update(memory, content, confidence, now, structured_value)
await self.invalidate_recall_cache(customer_id)
return updated
memory = MemoryUnit(
memory_uuid=str(uuid4()), customer_id=customer_id,
memory_key=memory_key, content=content, memory_type=memory_type,
source_type=source_type, source_confidence=confidence,
confidence=confidence, structured_value=structured_value,
evidence_count=0, conflict_count=0, recall_count=0,
status=status, valid_from=now, version=1, created_at=now, updated_at=now,
)
self.session.add(memory)
try:
async with self.session.begin_nested():
await self.session.flush()
except IntegrityError:
# 并发写入触发 uk_memory_unit_customer_active_key:改用已存在的有效记忆。
existing = await self._active(customer_id, memory_key)
if existing is None:
raise
updated = await self._update(existing, content, confidence, now, structured_value)
await self.invalidate_recall_cache(customer_id)
return updated
await self.invalidate_recall_cache(customer_id)
return memory
async def invalidate_recall_cache(self, customer_id: int) -> int:
"""写入生效后使该客户的召回热缓存失效,避免 TTL 内召回不到新记忆。
键集一律由 `MemoryRecallService.cache_keys` 枚举给出,调用方不得手写缓存前缀,
否则失效动作会打在并不存在的键上。
取舍:删除发生在写入 flush 之后。极端并发下(本事务未提交时另一读取回填了
旧结果)仍可能留下一条短命脏缓存,代价是多删一次可重建的缓存;缓存失效失败
只告警、不抛错,绝不阻塞写入主流程(缓存是可重建的加速层)。
"""
if self.cache is None:
return 0
# 延迟导入:`memory_recall_service` 反向依赖本模块,模块级导入会成环。
from app.service.memory_recall_service import MemoryRecallService
keys = MemoryRecallService.cache_keys(customer_id)
try:
removed = await self.cache.delete(*keys)
except Exception:
logger.warning("recall cache invalidation degraded customer_id=%s", customer_id)
return 0
return removed if isinstance(removed, int) else len(keys)
async def _active(self, customer_id: int, memory_key: str) -> MemoryUnit | None:
found: MemoryUnit | None = await self.session.scalar(
select(MemoryUnit).where(
MemoryUnit.customer_id == customer_id,
MemoryUnit.memory_key == memory_key,
MemoryUnit.status == "active",
)
)
return found
async def _update(
self, memory: MemoryUnit, content: str, confidence: float, now: datetime,
structured_value: dict[str, Any] | None = None,
) -> MemoryUnit:
if memory.content != content:
conflict = MemoryConflict(
left_memory_id=memory.id,
right_memory_id=await self._conflict_right_id(memory),
conflict_type="content_changed",
severity="low",
status="auto_resolved",
resolution=f"新内容覆盖旧内容:{memory.content[:200]} -> {content[:200]}",
winner_memory_id=memory.id,
created_at=now,
resolved_at=now,
)
self.session.add(conflict)
memory.conflict_count += 1
memory.content = content
memory.confidence = confidence
if structured_value is not None:
memory.structured_value = structured_value
memory.version += 1
memory.updated_at = now
await self.session.flush()
return memory
async def _conflict_right_id(self, memory: MemoryUnit) -> int:
"""冲突右侧标识:优先取同键历史版本行,否则取新版本的合成标识。
同一行原地更新时旧值与新值落在同一行,库中没有"新值行"的主键可用;
这里用左侧 id 的相反数作为新版本的稳定标识——它必然不等于左侧 id,
且与任何自增主键(恒为正)不冲突,因此不会与真实记忆行混淆。
"""
historical = await self.session.scalar(
select(MemoryUnit.id)
.where(
MemoryUnit.customer_id == memory.customer_id,
MemoryUnit.memory_key == memory.memory_key,
MemoryUnit.id != memory.id,
)
.order_by(MemoryUnit.version.desc())
.limit(1)
)
if historical is not None:
return int(historical)
return -int(memory.id)
async def record_evidence(
self,
memory: MemoryUnit,
*,
idempotency_key: str,
evidence_type: str,
excerpt: str | None,
snapshot: dict[str, Any] | None,
weight: float,
source_table: str | None = None,
source_record_id: str | None = None,
occurred_at: datetime | None = None,
) -> bool:
"""写入一条独立证据;`idempotency_key` 命中唯一键即视为已消费,返回 False。"""
now = datetime.now(UTC).replace(tzinfo=None)
existing = await self.session.scalar(
select(MemoryEvidence.id).where(MemoryEvidence.idempotency_key == idempotency_key)
)
if existing is not None:
return False
self.session.add(MemoryEvidence(
memory_id=memory.id, evidence_type=evidence_type,
source_table=source_table, source_record_id=source_record_id,
evidence_excerpt=(excerpt or "")[: self.EXCERPT_LIMIT] or None,
evidence_snapshot=snapshot, weight=weight,
idempotency_key=idempotency_key, occurred_at=occurred_at or now, created_at=now,
))
memory.evidence_count += 1
memory.last_evidenced_at = occurred_at or now
memory.updated_at = now
with suppress(IntegrityError):
async with self.session.begin_nested():
await self.session.flush()
return True
# 并发重复消费命中了 uk 唯一键:证据已由另一事务写入,不重复计数。
memory.evidence_count -= 1
return False
async def expire_stale(self, *, customer_id: int | None = None) -> int:
now = datetime.now(UTC).replace(tzinfo=None)
statement = select(MemoryUnit).where(
MemoryUnit.status == "active",
MemoryUnit.valid_until.is_not(None),
MemoryUnit.valid_until <= now,
)
if customer_id is not None:
statement = statement.where(MemoryUnit.customer_id == customer_id)
memories = list(await self.session.scalars(statement))
for memory in memories:
memory.status = "expired"
memory.updated_at = now
await self.session.flush()
return len(memories)
async def invalidate(self, memory_uuid: str, customer_id: int) -> bool:
memory = await self.session.scalar(
select(MemoryUnit).where(
MemoryUnit.memory_uuid == memory_uuid,
MemoryUnit.customer_id == customer_id,
MemoryUnit.status == "active",
)
)
if memory is None:
return False
memory.status = "invalidated"
memory.updated_at = datetime.now(UTC).replace(tzinfo=None)
await self.session.flush()
# 单条失效同样改变召回结果:不失效缓存会让已失效记忆在 TTL 内继续被召回。
await self.invalidate_recall_cache(customer_id)
return True