59 lines
2.0 KiB
Python
59 lines
2.0 KiB
Python
"""投顾 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 的统一上下文转换为投顾侧记忆列表。"""
|
|
|
|
def __init__(self, memory_service, *, session_prefix: str = "advisor-agent"):
|
|
self.memory_service = memory_service
|
|
self.session_prefix = session_prefix
|
|
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,
|
|
)
|
|
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"]
|