"""投顾 Agent 对现有统一记忆系统的适配边界。""" from __future__ import annotations from typing import Protocol class MemoryProvider(Protocol): async def recall(self, *, customer_id: int, query: str) -> list[dict]: """召回客户相关记忆,返回结构化记忆单元。""" class EmptyMemoryProvider: """记忆系统未接入时的安全默认实现。""" async def recall(self, *, customer_id: int, query: str) -> list[dict]: return [] class MemoryServiceProvider: """将现有 MemoryService 的统一上下文转换为投顾侧记忆列表。 投顾推荐场景**需要**长期记忆(偏好标签、投资目标),因此显式开启 ``include_long_term=True``。客服 Agent 走默认值 False,不重复召回。 """ def __init__( self, memory_service, *, session_prefix: str = "advisor-agent", include_long_term: bool = True, ): self.memory_service = memory_service self.session_prefix = session_prefix self.include_long_term = include_long_term self.last_warnings: list[str] = [] async def recall(self, *, customer_id: int, query: str) -> list[dict]: self.last_warnings = [] try: context = await self.memory_service.recall( customer_id=customer_id, session_id=f"{self.session_prefix}:{customer_id}", query=query, include_long_term=self.include_long_term, ) except Exception as exc: self.last_warnings = [f"advisor_memory_recall_failed:{type(exc).__name__}"] return [] self.last_warnings.extend(getattr(context, "warnings", []) or []) memories = getattr(context, "long_term_memories", context) return [self._to_dict(memory) for memory in memories] @staticmethod def _to_dict(memory) -> dict: if hasattr(memory, "model_dump"): data = memory.model_dump(mode="json") else: data = dict(memory) return { "customer_id": data.get("customer_id"), "tag": data.get("tag"), "content": data.get("content"), "info_type": data.get("info_type", "FACT"), "memory_type": data.get("memory_type"), } __all__ = ["EmptyMemoryProvider", "MemoryProvider", "MemoryServiceProvider"]