198 lines
8.0 KiB
Python
198 lines
8.0 KiB
Python
"""MySQL + Milvus + Neo4j 客户长期记忆同步服务。"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from datetime import datetime
|
|
from inspect import isawaitable
|
|
from typing import Any, Callable
|
|
|
|
from tool.llm import llm
|
|
from tool.confidence import BaseConfidenceCalcTool
|
|
|
|
from repositories.memory_unit import MemoryUnitRepo
|
|
from service.memory.milvus_memory import MilvusMemoryStore
|
|
from service.memory.neo4j_memory import Neo4jMemoryStore
|
|
from service.memory.schemas import MemoryUnitDTO
|
|
|
|
|
|
class LongTermMemoryService:
|
|
"""以 MySQL 为主体事实源,向 Milvus 和 Neo4j 同步镜像。"""
|
|
|
|
def __init__(
|
|
self,
|
|
*,
|
|
milvus_store: MilvusMemoryStore | None = None,
|
|
neo4j_store: Neo4jMemoryStore | None = None,
|
|
repository_factory=MemoryUnitRepo,
|
|
embedder: Callable[[str], Any] | None = None,
|
|
):
|
|
self.milvus_store = milvus_store or MilvusMemoryStore()
|
|
self.neo4j_store = neo4j_store or Neo4jMemoryStore()
|
|
self.repository_factory = repository_factory
|
|
self.embedder = embedder or llm.embed_one
|
|
self.confidence_tool = BaseConfidenceCalcTool()
|
|
|
|
async def save(self, db, memory: MemoryUnitDTO) -> tuple[MemoryUnitDTO, list[str]]:
|
|
"""保存主体并尽力同步两个外部索引,返回记忆和 warnings。"""
|
|
repo = self.repository_factory(db)
|
|
existing = await repo.find_exact(
|
|
memory.customer_id,
|
|
memory.memory_type.value,
|
|
memory.tag,
|
|
memory.content,
|
|
)
|
|
warnings: list[str] = []
|
|
if existing is not None:
|
|
existing = await repo.merge_evidence(existing)
|
|
entity = existing
|
|
else:
|
|
# final_score 是召回重排阶段的临时分数,不属于 MySQL 主体字段。
|
|
values = memory.model_dump(
|
|
mode="json",
|
|
exclude={"id", "milvus_id", "graph_node_id", "final_score"},
|
|
)
|
|
values["memory_type"] = memory.memory_type.value
|
|
values["source"] = memory.source.value
|
|
confidence_result, confidence_warning = self._calculate_confidence(memory)
|
|
warnings.extend(confidence_warning)
|
|
values.update(confidence_result)
|
|
entity = await repo.add_memory(values)
|
|
|
|
if existing is not None:
|
|
confidence_result, confidence_warning = self._calculate_confidence(entity)
|
|
warnings.extend(confidence_warning)
|
|
await repo.update_sync_status(entity.id, **confidence_result)
|
|
for key, value in confidence_result.items():
|
|
setattr(entity, key, value)
|
|
|
|
try:
|
|
vector = self.embedder(entity.content)
|
|
if isawaitable(vector):
|
|
vector = await vector
|
|
milvus_id = await self.milvus_store.upsert(entity, vector)
|
|
await repo.update_sync_status(
|
|
entity.id, milvus_id=milvus_id, milvus_sync_status="success"
|
|
)
|
|
entity.milvus_id = milvus_id
|
|
entity.milvus_sync_status = "success"
|
|
except Exception as exc:
|
|
warnings.append(f"milvus_sync_failed:{type(exc).__name__}")
|
|
await repo.update_sync_status(
|
|
entity.id,
|
|
milvus_sync_status="failed",
|
|
sync_retry_count=(entity.sync_retry_count or 0) + 1,
|
|
last_sync_error=str(exc)[:500],
|
|
)
|
|
|
|
try:
|
|
graph_id = await self.neo4j_store.upsert(entity)
|
|
await repo.update_sync_status(
|
|
entity.id, graph_node_id=graph_id, neo4j_sync_status="success"
|
|
)
|
|
entity.graph_node_id = graph_id
|
|
entity.neo4j_sync_status = "success"
|
|
except Exception as exc:
|
|
warnings.append(f"neo4j_sync_failed:{type(exc).__name__}")
|
|
await repo.update_sync_status(
|
|
entity.id,
|
|
neo4j_sync_status="failed",
|
|
sync_retry_count=(entity.sync_retry_count or 0) + 1,
|
|
last_sync_error=str(exc)[:500],
|
|
)
|
|
return self._to_dto(entity), warnings
|
|
|
|
def _calculate_confidence(self, memory) -> tuple[dict[str, Any], list[str]]:
|
|
"""计算记忆置信度;异常时强制降级为候选记忆。"""
|
|
source = memory.source.value if hasattr(memory.source, "value") else memory.source
|
|
memory_type = (
|
|
memory.memory_type.value
|
|
if hasattr(memory.memory_type, "value")
|
|
else memory.memory_type
|
|
)
|
|
create_time = getattr(memory, "create_time", None)
|
|
age_days = max(0, (datetime.now() - create_time).days) if create_time else 0
|
|
try:
|
|
result = self.confidence_tool.evaluate(
|
|
tag=memory.tag,
|
|
source=source,
|
|
evidence_count=memory.evidence_count or 0,
|
|
conflict_count=memory.conflict_count or 0,
|
|
age_days=age_days,
|
|
memory_type=memory_type,
|
|
)
|
|
result.pop("age_days", None)
|
|
result.pop("threshold", None)
|
|
result["confidence_update_time"] = datetime.now()
|
|
return result, []
|
|
except Exception as exc:
|
|
return {
|
|
"status": "candidate",
|
|
"confidence_reason": "置信度计算失败,降级保存为候选记忆",
|
|
"confidence_version": BaseConfidenceCalcTool.VERSION,
|
|
"confidence_update_time": datetime.now(),
|
|
}, [f"confidence_calculation_failed:{type(exc).__name__}"]
|
|
|
|
async def recall(
|
|
self,
|
|
db,
|
|
customer_id: int,
|
|
*,
|
|
memory_type: str | None = None,
|
|
tag: str | None = None,
|
|
limit: int = 100,
|
|
) -> tuple[list[MemoryUnitDTO], list[str]]:
|
|
"""按客户、类型、标签和有效期召回主体记忆。"""
|
|
entities = await self.repository_factory(db).list_for_customer(
|
|
customer_id, memory_type=memory_type, tag=tag, limit=limit
|
|
)
|
|
return [self._to_dto(entity) for entity in entities], []
|
|
|
|
async def retry_pending(self, db, *, limit: int = 100) -> dict[str, int]:
|
|
"""重试 MySQL 中缺少外部索引或同步失败的记忆。"""
|
|
entities = await self.repository_factory(db).list_pending_sync(limit)
|
|
success = 0
|
|
failed = 0
|
|
for entity in entities:
|
|
dto = self._to_dto(entity)
|
|
_, warnings = await self.save(db, dto)
|
|
if warnings:
|
|
failed += 1
|
|
else:
|
|
success += 1
|
|
return {"success": success, "failed": failed}
|
|
|
|
@staticmethod
|
|
def _to_dto(entity) -> MemoryUnitDTO:
|
|
"""将 ORM 实体转换为跨层 DTO。"""
|
|
data = {
|
|
"id": entity.id,
|
|
"customer_id": entity.customer_id,
|
|
"session_id": entity.session_id,
|
|
"agent_run_id": entity.agent_run_id,
|
|
"memory_type": entity.memory_type,
|
|
"tag": entity.tag,
|
|
"content": entity.content,
|
|
"info_type": entity.info_type,
|
|
"source": entity.source,
|
|
"evidence_ref": entity.evidence_ref or [],
|
|
"source_confidence": float(entity.source_confidence or 0),
|
|
"confidence": float(entity.confidence or 0),
|
|
"historical_accuracy": float(entity.historical_accuracy or 0),
|
|
"confidence_version": getattr(entity, "confidence_version", None),
|
|
"confidence_reason": getattr(entity, "confidence_reason", None),
|
|
"confidence_update_time": getattr(entity, "confidence_update_time", None),
|
|
"evidence_count": entity.evidence_count or 0,
|
|
"conflict_count": entity.conflict_count or 0,
|
|
"recall_count": entity.recall_count or 0,
|
|
"status": entity.status,
|
|
"valid_from": entity.valid_from,
|
|
"valid_until": entity.valid_until,
|
|
"last_verified_at": entity.last_verified_at,
|
|
"milvus_id": entity.milvus_id,
|
|
"graph_node_id": entity.graph_node_id,
|
|
}
|
|
return MemoryUnitDTO.model_validate(data)
|
|
|
|
|
|
__all__ = ["LongTermMemoryService"]
|