feat:客户agent以及记忆模块功能开发

This commit is contained in:
2026-09-11 22:38:15 +08:00
parent 0197ac2ef2
commit 4da4775e8c
39 changed files with 2508 additions and 14 deletions
+29
View File
@@ -0,0 +1,29 @@
"""客服 Agent 记忆模块。"""
from .facade import MemoryService
from .customer_product import CustomerProductMemory
from .customer_product import CustomerProductMemory
from .short_term import SessionExpiredError, ShortTermMemory, ShortTermMemoryError
from .schemas import (
CustomerMemoryContext,
MemorySource,
MemoryStatus,
MemoryType,
MemoryUnitDTO,
ShortTermMessage,
)
__all__ = [
"CustomerMemoryContext",
"CustomerProductMemory",
"CustomerProductMemory",
"MemoryService",
"MemorySource",
"MemoryStatus",
"MemoryType",
"MemoryUnitDTO",
"SessionExpiredError",
"ShortTermMemory",
"ShortTermMessage",
"ShortTermMemoryError",
]
+75
View File
@@ -0,0 +1,75 @@
"""Redis 短期消息到 MySQL conversation_archive 的归档服务。"""
from __future__ import annotations
import re
from typing import Any
from repositories.conversation_archive import ConversationArchiveRepo
from service.memory.short_term import ShortTermMemory
class ConversationArchiver:
"""处理会话读取、脱敏、批量写入和成功后清理。"""
def __init__(
self,
*,
short_term: ShortTermMemory,
repository_factory=ConversationArchiveRepo,
):
self.short_term = short_term
self.repository_factory = repository_factory
async def archive_session(
self,
db,
*,
session_id: str,
user_id: int,
customer_id: int | None,
agent_type: str = "customer",
trace_id: str | None = None,
agent_run_id: str | None = None,
) -> int:
"""归档完整会话;只有数据库成功后才清理 Redis。"""
messages = await self.short_term.load_messages(session_id)
rows = [
{
"session_id": session_id,
"customer_id": customer_id,
"user_id": user_id,
"agent_type": agent_type,
"role": message.role,
"content": self.redact(message.content),
"tool_calls": self.redact(message.tool_calls),
"message_id": message.message_id,
"agent_run_id": message.agent_run_id or agent_run_id,
"trace_id": trace_id,
}
for message in messages
]
archived_count = await self.repository_factory(db).archive_batch(rows)
await self.short_term.clear_session(session_id)
return archived_count
@staticmethod
def redact(value: Any) -> Any:
"""脱敏手机号、邮箱和常见身份证号,保留消息可读性。"""
if isinstance(value, str):
value = re.sub(r"(?<!\d)(1[3-9]\d)\d{4}(\d{4})(?!\d)", r"\1****\2", value)
value = re.sub(
r"([A-Za-z0-9._%+-])[A-Za-z0-9._%+-]*(@[A-Za-z0-9.-]+)",
r"\1***\2",
value,
)
value = re.sub(r"(?<!\d)(\d{3})\d{11}(\d{2})(?!\d)", r"\1***********\2", value)
return value
if isinstance(value, list):
return [ConversationArchiver.redact(item) for item in value]
if isinstance(value, dict):
return {key: ConversationArchiver.redact(item) for key, item in value.items()}
return value
__all__ = ["ConversationArchiver"]
+34
View File
@@ -0,0 +1,34 @@
"""客服 Agent 记忆上下文组装。"""
from __future__ import annotations
from service.memory.schemas import CustomerMemoryContext, MemoryUnitDTO, ShortTermMessage
def build_customer_memory_context(
*,
customer_id: int,
session_id: str,
short_term_messages: list[ShortTermMessage],
customer_profile: dict | None,
work_orders: list[dict],
long_term_memories: list[MemoryUnitDTO],
customer_relations: list[dict] | None = None,
customer_products: list[dict] | None = None,
warnings: list[str] | None = None,
) -> CustomerMemoryContext:
"""构造稳定的客服记忆上下文结构。"""
return CustomerMemoryContext(
customer_id=customer_id,
session_id=session_id,
short_term_messages=short_term_messages,
customer_profile=customer_profile,
work_orders=work_orders,
long_term_memories=long_term_memories,
customer_relations=customer_relations or [],
customer_products=customer_products or [],
warnings=warnings or [],
)
__all__ = ["build_customer_memory_context"]
+52
View File
@@ -0,0 +1,52 @@
"""客户-产品关系中期记忆读取。"""
from __future__ import annotations
from decimal import Decimal
from typing import Any
from repositories.fin_holdings import FinHoldingsRepo
class CustomerProductMemory:
"""读取客户持仓及其关联基金产品,供客服上下文使用。"""
def __init__(self, *, repository_factory=FinHoldingsRepo):
self.repository_factory = repository_factory
async def list(
self, db, customer_id: int, *, include_closed: bool = True
) -> list[dict[str, Any]]:
"""按客户 ID 返回持仓与产品信息,避免跨客户读取。"""
holdings = await self.repository_factory(db).list_with_products(
customer_id, include_closed=include_closed
)
return [self._to_dict(holding, product) for holding, product in holdings]
@staticmethod
def _to_dict(holding, product) -> dict[str, Any]:
"""将持仓和产品 ORM 对象转换为上下文可用的 JSON 友好结构。"""
def scalar(value):
return str(value) if isinstance(value, Decimal) else value
return {
"holding_id": holding.id,
"customer_id": holding.customer_id,
"product_id": holding.product_id,
"shares": scalar(holding.shares),
"cost_amount": scalar(holding.cost_amount),
"current_value": scalar(holding.current_value),
"profit_loss": scalar(holding.profit_loss),
"profit_ratio": scalar(holding.profit_ratio),
"holding_status": holding.status,
"holding_create_time": holding.create_time,
"holding_update_time": holding.update_time,
"product_code": product.product_code if product else None,
"product_name": product.product_name if product else None,
"product_type": product.product_type if product else None,
"risk_level": product.risk_level if product else None,
"product_status": product.status if product else None,
}
__all__ = ["CustomerProductMemory"]
+30
View File
@@ -0,0 +1,30 @@
"""客户关系中期记忆读取。"""
from repositories.customer_relation import CustomerRelationRepo
class CustomerRelationMemory:
"""读取客服上下文需要的客户-投顾关系。"""
def __init__(self, *, repository_factory=CustomerRelationRepo):
self.repository_factory = repository_factory
async def list(self, db, customer_id: int) -> list[dict]:
"""按客户 ID 返回有效关系。"""
relations = await self.repository_factory(db).list_by_customer(customer_id)
return [
{
"id": relation.id,
"customer_id": relation.customer_id,
"advisor_id": relation.advisor_id,
"assign_time": relation.assign_time,
"signed_time": relation.signed_time,
"end_time": relation.end_time,
"status": relation.status,
"reason": relation.reason,
}
for relation in relations
]
__all__ = ["CustomerRelationMemory"]
+175
View File
@@ -0,0 +1,175 @@
"""MemoryService Facade:客服 Agent 的唯一记忆调用入口。"""
from __future__ import annotations
from contextlib import asynccontextmanager
from typing import Any
from config.database.mysql import get_session_factory
from tool.confidence_rank import FinalConfidenceRankTool
from .archive import ConversationArchiver
from .customer_relation import CustomerRelationMemory
from .customer_product import CustomerProductMemory
from .context_builder import build_customer_memory_context
from .long_term import LongTermMemoryService
from .profile import CustomerProfileMemory
from .schemas import CustomerMemoryContext, MemoryUnitDTO, ShortTermMessage
from .short_term import ShortTermMemory
from .work_order import WorkOrderMemory
class MemoryService:
"""协调短期、中期和长期记忆,屏蔽底层数据库细节。"""
def __init__(
self,
*,
session_factory=None,
short_term: ShortTermMemory | None = None,
profile: CustomerProfileMemory | None = None,
work_orders: WorkOrderMemory | None = None,
relations: CustomerRelationMemory | None = None,
products: CustomerProductMemory | None = None,
long_term: LongTermMemoryService | None = None,
archiver: ConversationArchiver | None = None,
rank_tool: FinalConfidenceRankTool | None = None,
):
self.session_factory = session_factory or get_session_factory()
self.short_term = short_term or ShortTermMemory()
self.profile = profile or CustomerProfileMemory(redis=self.short_term.redis)
self.work_orders = work_orders or WorkOrderMemory()
self.relations = relations or CustomerRelationMemory()
self.products = products or CustomerProductMemory()
self.long_term = long_term or LongTermMemoryService()
self.archiver = archiver or ConversationArchiver(short_term=self.short_term)
self.rank_tool = rank_tool or FinalConfidenceRankTool()
@asynccontextmanager
async def _db(self):
"""按请求获取并释放 MySQL 会话。"""
async with self.session_factory() as session:
yield session
async def recall(
self,
*,
customer_id: int,
session_id: str,
query: str | None = None,
limit: int = 10,
) -> CustomerMemoryContext:
"""召回并组装客服 Agent 当前请求的全部可用记忆。"""
if limit < 0:
raise ValueError("limit 必须是非负整数")
warnings: list[str] = []
short_term_messages = await self.short_term.load_messages(session_id)
warnings.extend(self.short_term.last_warnings)
async with self._db() as db:
profile, profile_warnings = await self.profile.get(db, customer_id)
warnings.extend(profile_warnings)
try:
work_orders = await self.work_orders.list(db, customer_id)
except Exception as exc:
work_orders = []
warnings.append(f"work_order_recall_failed:{type(exc).__name__}")
try:
customer_relations = await self.relations.list(db, customer_id)
except Exception as exc:
customer_relations = []
warnings.append(f"customer_relation_recall_failed:{type(exc).__name__}")
try:
customer_products = await self.products.list(db, customer_id)
except Exception as exc:
customer_products = []
warnings.append(f"customer_product_recall_failed:{type(exc).__name__}")
try:
memories, memory_warnings = await self.long_term.recall(
db, customer_id, limit=max(limit, 10)
)
warnings.extend(memory_warnings)
except Exception as exc:
memories = []
warnings.append(f"long_term_recall_failed:{type(exc).__name__}")
ranked = self.rank_tool.rank(
[memory.model_dump(mode="json") for memory in memories],
top_k=limit,
)
ranked_memories = [MemoryUnitDTO.model_validate(item) for item in ranked]
return build_customer_memory_context(
customer_id=customer_id,
session_id=session_id,
short_term_messages=short_term_messages,
customer_profile=profile,
work_orders=work_orders,
customer_relations=customer_relations,
customer_products=customer_products,
long_term_memories=ranked_memories,
warnings=warnings,
)
async def append_message(
self,
*,
customer_id: int,
session_id: str,
message: ShortTermMessage,
) -> list[str]:
"""写入短期消息,并统一返回降级 warnings。"""
await self.short_term.append_message(
session_id,
message.role,
message.content,
message_id=message.message_id,
agent_run_id=message.agent_run_id,
tool_calls=message.tool_calls,
)
return list(self.short_term.last_warnings)
async def save_memory(
self,
*,
customer_id: int,
memory: MemoryUnitDTO,
) -> tuple[MemoryUnitDTO, list[str]]:
"""保存长期记忆,并返回记忆结果和底层 warnings。"""
if memory.customer_id != customer_id:
raise ValueError("memory.customer_id 与当前客户不一致")
async with self._db() as db:
return await self.long_term.save(db, memory)
async def close_session(
self,
*,
customer_id: int,
session_id: str,
agent_run_id: str | None = None,
) -> list[str]:
"""归档并关闭客户会话;归档失败时保留 Redis 消息。"""
try:
async with self._db() as db:
await self.archiver.archive_session(
db,
session_id=session_id,
user_id=customer_id,
customer_id=customer_id,
agent_type="customer",
agent_run_id=agent_run_id,
)
return []
except Exception as exc:
return [f"session_close_failed:{type(exc).__name__}"]
def normalize_warnings(value: Any) -> list[str]:
"""将底层异常或组件返回的提示统一为字符串列表。"""
if value is None:
return []
if isinstance(value, str):
return [value]
return [str(item) for item in value]
__all__ = ["MemoryService", "normalize_warnings"]
+193
View File
@@ -0,0 +1,193 @@
"""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:
values = memory.model_dump(mode="json", exclude={"id", "milvus_id", "graph_node_id"})
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"]
+114
View File
@@ -0,0 +1,114 @@
"""客户长期记忆的 Milvus 向量镜像。"""
from __future__ import annotations
from pymilvus import AsyncMilvusClient, DataType
from config.database.milvus import client as configured_client
from rag.embedding import EMBEDDING_DIMENSION
CUSTOMER_MEMORY_COLLECTION = "customer_memory"
def build_memory_schema():
"""构造客户记忆向量集合结构,维度来自 LLM_EMBED_DIMENSIONS。"""
schema = AsyncMilvusClient.create_schema(auto_id=False, enable_dynamic_field=False)
schema.add_field("memory_id", DataType.VARCHAR, is_primary=True, max_length=128)
schema.add_field("customer_id", DataType.INT64)
schema.add_field("memory_type", DataType.VARCHAR, max_length=32)
schema.add_field("tag", DataType.VARCHAR, max_length=64)
schema.add_field("content", DataType.VARCHAR, max_length=2048)
schema.add_field("status", DataType.VARCHAR, max_length=16)
schema.add_field("vector", DataType.FLOAT_VECTOR, dim=EMBEDDING_DIMENSION)
return schema
def build_memory_index_params():
"""构造客户记忆向量索引。"""
params = AsyncMilvusClient.prepare_index_params()
params.add_index(
field_name="vector",
index_type="HNSW",
metric_type="COSINE",
params={"M": 16, "efConstruction": 200},
)
return params
class MilvusMemoryStore:
"""封装客户记忆向量写入、查询和删除。"""
def __init__(self, client=None, *, collection_name=CUSTOMER_MEMORY_COLLECTION):
self.client = client or configured_client()
self.collection_name = collection_name
async def ensure_collection(self) -> None:
"""创建集合或校验已有集合的向量维度。"""
if not await self.client.has_collection(self.collection_name):
await self.client.create_collection(
collection_name=self.collection_name,
schema=build_memory_schema(),
index_params=build_memory_index_params(),
)
return
desc = await self.client.describe_collection(self.collection_name)
for field in desc.get("fields", []):
if field.get("name") == "vector":
dim = field.get("params", {}).get("dim")
if dim is not None and int(dim) != EMBEDDING_DIMENSION:
raise RuntimeError(
f"Milvus collection {self.collection_name!r} vector dim={dim}, "
f"expected {EMBEDDING_DIMENSION}"
)
async def upsert(self, memory, vector: list[float]) -> str:
"""写入一条客户记忆向量并返回 Milvus 主键。"""
await self.ensure_collection()
memory_id = str(memory.id)
await self.client.delete(
collection_name=self.collection_name,
filter=f'memory_id == "{memory_id}"',
)
await self.client.insert(
collection_name=self.collection_name,
data=[
{
"memory_id": memory_id,
"customer_id": memory.customer_id,
"memory_type": memory.memory_type,
"tag": memory.tag,
"content": memory.content,
"status": memory.status,
"vector": vector,
}
],
)
return memory_id
async def search(self, vector: list[float], customer_id: int, *, limit: int = 10) -> list[dict]:
"""按客户 ID 过滤向量查询结果。"""
await self.ensure_collection()
return await self.client.search(
collection_name=self.collection_name,
data=[vector],
limit=limit,
filter=f"customer_id == {int(customer_id)}",
output_fields=["memory_id", "customer_id", "memory_type", "tag", "content", "status"],
)
async def delete(self, memory_id: int | str) -> None:
"""删除一条客户记忆向量。"""
await self.client.delete(
collection_name=self.collection_name,
filter=f'memory_id == "{memory_id}"',
)
__all__ = [
"CUSTOMER_MEMORY_COLLECTION",
"MilvusMemoryStore",
"build_memory_index_params",
"build_memory_schema",
]
+60
View File
@@ -0,0 +1,60 @@
"""客户长期记忆的 Neo4j 关系镜像。"""
from __future__ import annotations
from config.database.neo4j import client as configured_client
class Neo4jMemoryStore:
"""保存客户节点、记忆节点及其 HAS_MEMORY 关系。"""
def __init__(self, driver=None):
self.driver = driver or configured_client()
async def upsert(self, memory) -> str:
"""写入客户和记忆节点,返回图谱记忆节点 ID。"""
memory_id = str(memory.id)
query = """
MERGE (c:Customer {customer_id: $customer_id})
MERGE (m:CustomerMemory {memory_id: $memory_id})
SET m.memory_type = $memory_type,
m.tag = $tag,
m.content = $content,
m.status = $status
MERGE (c)-[:HAS_MEMORY]->(m)
RETURN m.memory_id AS memory_id
"""
async with self.driver.session() as session:
record = await session.run(
query,
customer_id=int(memory.customer_id),
memory_id=memory_id,
memory_type=memory.memory_type,
tag=memory.tag,
content=memory.content,
status=memory.status,
)
row = await record.single()
return row["memory_id"] if row else memory_id
async def list_by_customer(self, customer_id: int, *, limit: int = 100) -> list[dict]:
"""按客户查询图谱记忆关系。"""
query = """
MATCH (c:Customer {customer_id: $customer_id})-[:HAS_MEMORY]->(m:CustomerMemory)
RETURN m.memory_id AS memory_id, m.memory_type AS memory_type,
m.tag AS tag, m.content AS content, m.status AS status
LIMIT $limit
"""
async with self.driver.session() as session:
result = await session.run(query, customer_id=int(customer_id), limit=limit)
return [dict(record) async for record in result]
async def delete(self, memory_id: int | str) -> None:
"""删除记忆节点及其关系。"""
query = "MATCH (m:CustomerMemory {memory_id: $memory_id}) DETACH DELETE m"
async with self.driver.session() as session:
await session.run(query, memory_id=str(memory_id))
__all__ = ["Neo4jMemoryStore"]
+88
View File
@@ -0,0 +1,88 @@
"""客户画像中期记忆:MySQL 事实源 + Redis Cache-Aside。"""
from __future__ import annotations
import json
from decimal import Decimal
from typing import Any
from config.database.redis import client as redis_client
from repositories.fin_customer_profile import FinCustomerProfileRepo
class CustomerProfileMemory:
"""提供客户画像读取、缓存失效和刷新能力。"""
CACHE_TTL = 7 * 24 * 60 * 60
def __init__(self, *, redis=None, repository_factory=FinCustomerProfileRepo):
self.redis = redis or redis_client()
self.repository_factory = repository_factory
@staticmethod
def cache_key(customer_id: int) -> str:
"""生成客户画像缓存 Key。"""
return f"profile:{customer_id}"
async def get(self, db, customer_id: int) -> tuple[dict[str, Any] | None, list[str]]:
"""优先读取缓存,未命中后回源 MySQL,并返回 warnings。"""
warnings: list[str] = []
key = self.cache_key(customer_id)
try:
cached = await self.redis.get(key)
if cached:
return json.loads(cached), warnings
except Exception as exc:
warnings.append(f"profile_cache_read_failed:{type(exc).__name__}")
try:
profile = await self.repository_factory(db).get_by_customer_id(customer_id)
except Exception as exc:
warnings.append(f"profile_mysql_read_failed:{type(exc).__name__}")
return None, warnings
if profile is None:
return None, warnings
payload = self._to_dict(profile)
try:
await self.redis.set(key, json.dumps(payload, ensure_ascii=False), ex=self.CACHE_TTL)
except Exception as exc:
warnings.append(f"profile_cache_write_failed:{type(exc).__name__}")
return payload, warnings
async def invalidate(self, customer_id: int) -> list[str]:
"""删除客户画像缓存,确保更新后不会长期读取旧值。"""
try:
await self.redis.delete(self.cache_key(customer_id))
return []
except Exception as exc:
return [f"profile_cache_invalidate_failed:{type(exc).__name__}"]
async def refresh(self, db, customer_id: int) -> tuple[dict[str, Any] | None, list[str]]:
"""先删除缓存,再从 MySQL 读取并重新缓存画像。"""
warnings = await self.invalidate(customer_id)
profile, read_warnings = await self.get(db, customer_id)
return profile, warnings + read_warnings
@staticmethod
def _to_dict(profile) -> dict[str, Any]:
"""将 ORM 画像转换为可安全写入 Redis 的字典。"""
data = {
"customer_id": profile.customer_id,
"risk_level": profile.risk_level,
"risk_score": profile.risk_score,
"investment_experience": profile.investment_experience,
"annual_income_range": profile.annual_income_range,
"total_assets": profile.total_assets,
"asset_allocation": profile.asset_allocation,
"product_preference": profile.product_preference,
"customer_level": profile.customer_level,
"confidence_score": profile.confidence_score,
"profile_version": profile.profile_version,
"update_time": profile.update_time.isoformat() if profile.update_time else None,
}
return json.loads(json.dumps(data, default=lambda value: str(value), ensure_ascii=False))
__all__ = ["CustomerProfileMemory"]
+8
View File
@@ -0,0 +1,8 @@
"""记忆模块依赖健康检查。"""
from config.database import check_ready_detail
async def check_memory_dependencies() -> dict[str, dict]:
"""复用项目统一数据库健康检查,避免记忆模块重复管理连接。"""
return await check_ready_detail()
+106
View File
@@ -0,0 +1,106 @@
"""记忆模块跨层共享的数据契约。
这些模型只描述记忆模块与客服 Agent 之间的输入输出,不绑定 Redis、MySQL、
Milvus 或 Neo4j 的实现细节。
"""
from __future__ import annotations
from datetime import datetime
from enum import StrEnum
from typing import Any
from pydantic import BaseModel, ConfigDict, Field
class MemoryType(StrEnum):
"""客服记忆可以保存的业务类型。"""
PROFILE_FACT = "PROFILE_FACT"
PROFILE_CANDIDATE = "PROFILE_CANDIDATE"
CUSTOMER_PREFERENCE = "CUSTOMER_PREFERENCE"
INVESTMENT_GOAL = "INVESTMENT_GOAL"
SERVICE_FACT = "SERVICE_FACT"
CUSTOMER_RELATION = "CUSTOMER_RELATION"
class MemoryStatus(StrEnum):
"""记忆生命周期状态。"""
CANDIDATE = "candidate"
CONFIRMED = "confirmed"
EXPIRED = "expired"
REJECTED = "rejected"
ARCHIVED = "archived"
class MemorySource(StrEnum):
"""客服对话形成记忆的证据来源。"""
DIALOGUE_CONFIRMED = "dialogue_confirmed"
DIALOGUE_STATED = "dialogue_stated"
DIALOGUE_INFERRED = "dialogue_inferred"
class ShortTermMessage(BaseModel):
"""Redis 短期会话中的一条消息。"""
model_config = ConfigDict(extra="forbid")
message_id: str = Field(min_length=1, max_length=64)
session_id: str = Field(min_length=1, max_length=64)
role: str = Field(min_length=1, max_length=16)
content: str = Field(min_length=1)
token_count: int = Field(default=0, ge=0)
agent_run_id: str | None = Field(default=None, max_length=64)
tool_calls: list[dict[str, Any]] = Field(default_factory=list)
create_time: datetime | None = None
class MemoryUnitDTO(BaseModel):
"""统一表示一条客户记忆,供写入、召回和重排使用。"""
model_config = ConfigDict(extra="forbid")
id: int | None = None
customer_id: int
session_id: str | None = Field(default=None, max_length=64)
agent_run_id: str | None = Field(default=None, max_length=64)
memory_type: MemoryType
tag: str = Field(min_length=1, max_length=64)
content: str = Field(min_length=1, max_length=512)
info_type: str = Field(default="FACT", min_length=1, max_length=8)
source: MemorySource
evidence_ref: list[dict[str, Any]] = Field(default_factory=list)
source_confidence: float = Field(default=0.2, ge=0.0, le=1.0)
confidence: float = Field(default=0.2, ge=0.0, le=1.0)
historical_accuracy: float = Field(default=0.5, ge=0.0, le=1.0)
confidence_version: str | None = Field(default=None, max_length=32)
confidence_reason: str | None = Field(default=None, max_length=255)
confidence_update_time: datetime | None = None
final_score: float | None = Field(default=None, ge=0.0, le=1.0)
evidence_count: int = Field(default=0, ge=0)
conflict_count: int = Field(default=0, ge=0)
recall_count: int = Field(default=0, ge=0)
status: MemoryStatus = MemoryStatus.CANDIDATE
valid_from: datetime | None = None
valid_until: datetime | None = None
last_verified_at: datetime | None = None
milvus_id: str | None = Field(default=None, max_length=128)
graph_node_id: str | None = Field(default=None, max_length=128)
class CustomerMemoryContext(BaseModel):
"""客服 Agent 每轮请求可使用的统一记忆上下文。"""
model_config = ConfigDict(extra="forbid")
customer_id: int
session_id: str
short_term_messages: list[ShortTermMessage] = Field(default_factory=list)
customer_profile: dict[str, Any] | None = None
work_orders: list[dict[str, Any]] = Field(default_factory=list)
long_term_memories: list[MemoryUnitDTO] = Field(default_factory=list)
customer_relations: list[dict[str, Any]] = Field(default_factory=list)
customer_products: list[dict[str, Any]] = Field(default_factory=list)
warnings: list[str] = Field(default_factory=list)
+232
View File
@@ -0,0 +1,232 @@
"""Redis 短期会话记忆。"""
from __future__ import annotations
import asyncio
import json
import time
import uuid
from datetime import datetime, timezone
from inspect import isawaitable
from typing import Any, Callable
from config.database.redis import client as redis_client
from .schemas import ShortTermMessage
class ShortTermMemoryError(RuntimeError):
"""短期记忆操作失败。"""
class SessionExpiredError(ShortTermMemoryError):
"""会话已超过最长生命周期。"""
async def _config(config_getter, key: str, default):
"""读取项目配置,并将数据库配置值转换为默认值类型。"""
value = config_getter(key, str(default))
if isawaitable(value):
value = await value
return type(default)(value)
class ShortTermMemory:
"""管理当前客服会话的 Redis 消息、Token 预算和生命周期。"""
MESSAGE_KEY = "session:{session_id}:messages"
META_KEY = "session:{session_id}:meta"
ALLOWED_ROLES = frozenset({"user", "assistant", "system"})
def __init__(
self,
redis=None,
*,
config_getter=None,
token_counter: Callable[[str], int] | None = None,
clock=time.time,
fail_soft: bool = True,
):
"""创建短期记忆服务,默认复用项目级 Redis 客户端。"""
self.redis = redis or redis_client()
self.config_getter = config_getter or (lambda _key, default: default)
self.token_counter = token_counter or (lambda text: max(1, len(text) // 4))
self.clock = clock
self.fail_soft = fail_soft
self.last_warnings: list[str] = []
self._locks: dict[str, asyncio.Lock] = {}
@classmethod
def message_key(cls, session_id: str) -> str:
"""生成会话消息列表 Key。"""
return cls.MESSAGE_KEY.format(session_id=session_id)
@classmethod
def meta_key(cls, session_id: str) -> str:
"""生成会话元数据 Hash Key。"""
return cls.META_KEY.format(session_id=session_id)
def _lock(self, session_id: str) -> asyncio.Lock:
"""获取进程内会话锁,避免并发截断互相覆盖。"""
return self._locks.setdefault(session_id, asyncio.Lock())
async def append_message(
self,
session_id: str,
role: str,
content: str,
*,
message_id: str | None = None,
agent_run_id: str | None = None,
tool_calls: list[dict[str, Any]] | None = None,
) -> ShortTermMessage | None:
"""追加消息、刷新 TTL,并按 Token 预算从旧到新截断。"""
self.last_warnings = []
if role not in self.ALLOWED_ROLES:
raise ValueError(f"非法消息角色: {role}")
if not isinstance(content, str) or not content.strip():
raise ValueError("消息内容不能为空")
message = ShortTermMessage(
message_id=message_id or uuid.uuid4().hex,
session_id=session_id,
role=role,
content=content,
token_count=self.token_counter(content),
agent_run_id=agent_run_id,
tool_calls=tool_calls or [],
create_time=datetime.fromtimestamp(self.clock(), tz=timezone.utc),
)
try:
async with self._lock(session_id):
await self._ensure_meta(session_id)
await self.redis.rpush(
self.message_key(session_id), self._serialize(message)
)
await self._touch(session_id)
await self._trim(session_id)
return message
except SessionExpiredError:
raise
except Exception as exc:
return self._degrade("append_message", exc)
async def load_messages(self, session_id: str) -> list[ShortTermMessage]:
"""按写入顺序读取当前会话消息,并刷新空闲 TTL。"""
self.last_warnings = []
try:
if not await self._session_is_active(session_id):
return []
raw_messages = await self.redis.lrange(self.message_key(session_id), 0, -1)
await self._touch(session_id)
return [self._deserialize(raw) for raw in raw_messages if raw]
except Exception as exc:
return self._degrade("load_messages", exc) or []
async def get_message_count(self, session_id: str) -> int:
"""返回当前会话消息数量。"""
return len(await self.load_messages(session_id))
async def get_token_count(self, session_id: str) -> int:
"""返回当前会话消息的 Token 估算总数。"""
return sum(message.token_count for message in await self.load_messages(session_id))
async def clear_session(self, session_id: str) -> None:
"""清理会话消息和元数据。"""
self.last_warnings = []
try:
await self.redis.delete(self.message_key(session_id), self.meta_key(session_id))
except Exception as exc:
self._degrade("clear_session", exc)
async def _ensure_meta(self, session_id: str) -> None:
"""初始化会话元数据,并固定 24 小时绝对过期时间。"""
meta_key = self.meta_key(session_id)
meta = await self.redis.hgetall(meta_key)
now = self.clock()
if meta:
absolute_expire_at = float(meta.get("absolute_expire_at", now))
if absolute_expire_at <= now:
await self.clear_session(session_id)
raise SessionExpiredError("会话已超过最长生命周期")
return
max_lifetime = await _config(
self.config_getter, "agent.customer.session.max_lifetime", 86400
)
await self.redis.hset(
meta_key,
mapping={
"created_at": str(now),
"absolute_expire_at": str(now + max_lifetime),
},
)
await self.redis.expire(meta_key, max_lifetime)
async def _session_is_active(self, session_id: str) -> bool:
"""检查会话元数据是否存在且未达到绝对过期时间。"""
meta = await self.redis.hgetall(self.meta_key(session_id))
if not meta:
return False
if float(meta.get("absolute_expire_at", 0)) <= self.clock():
await self.clear_session(session_id)
return False
return True
async def _touch(self, session_id: str) -> None:
"""刷新空闲 TTL,但不超过绝对过期时间。"""
meta = await self.redis.hgetall(self.meta_key(session_id))
if not meta:
return
remaining = int(float(meta["absolute_expire_at"]) - self.clock())
if remaining <= 0:
await self.clear_session(session_id)
raise SessionExpiredError("会话已超过最长生命周期")
idle_ttl = await _config(
self.config_getter, "agent.customer.session.ttl", 1800
)
ttl = min(idle_ttl, remaining)
await self.redis.expire(self.message_key(session_id), ttl)
await self.redis.expire(self.meta_key(session_id), remaining)
async def _trim(self, session_id: str) -> None:
"""保留最新消息,确保最新一条消息不会因超预算被删除。"""
limit = await _config(
self.config_getter, "agent.customer.session_max_token", 4096
)
key = self.message_key(session_id)
raw_messages = [raw for raw in await self.redis.lrange(key, 0, -1) if raw]
total = 0
keep_from = len(raw_messages)
for index in range(len(raw_messages) - 1, -1, -1):
message = self._deserialize(raw_messages[index])
total += message.token_count
keep_from = index
if total > limit:
keep_from = index + 1
break
if raw_messages and keep_from == len(raw_messages):
keep_from = len(raw_messages) - 1
await self.redis.ltrim(key, keep_from, -1)
@staticmethod
def _serialize(message: ShortTermMessage) -> str:
"""将消息转换为 Redis List 中的 JSON 字符串。"""
return json.dumps(message.model_dump(mode="json"), ensure_ascii=False)
@staticmethod
def _deserialize(raw: str | bytes) -> ShortTermMessage:
"""将 Redis JSON 字符串恢复为消息 DTO。"""
if isinstance(raw, bytes):
raw = raw.decode("utf-8")
return ShortTermMessage.model_validate(json.loads(raw))
def _degrade(self, operation: str, exc: Exception):
"""记录降级原因;fail_soft 模式下不让 Redis 故障击穿客服请求。"""
warning = f"short_term_{operation}_degraded:{type(exc).__name__}"
self.last_warnings = [warning]
if not self.fail_soft:
raise ShortTermMemoryError(warning) from exc
return None
__all__ = ["SessionExpiredError", "ShortTermMemory", "ShortTermMemoryError"]
+40
View File
@@ -0,0 +1,40 @@
"""客户工单中期记忆。"""
from __future__ import annotations
from typing import Any
from repositories.work_order import WorkOrderRepo
class WorkOrderMemory:
"""以 MySQL 为事实来源读取客户工单状态。"""
def __init__(self, *, repository_factory=WorkOrderRepo):
self.repository_factory = repository_factory
async def list(self, db, customer_id: int, *, active_only: bool = True) -> list[dict[str, Any]]:
"""按客户 ID 查询工单并转换为客服上下文结构。"""
orders = await self.repository_factory(db).list_by_customer(
customer_id, active_only=active_only
)
return [
{
"id": order.id,
"work_order_no": order.work_order_no,
"order_type": order.order_type,
"sub_type": order.sub_type,
"customer_id": order.customer_id,
"handler_id": order.handler_id,
"current_node": order.current_node,
"priority": order.priority,
"status": order.status,
"biz_content": order.biz_content,
"create_time": order.create_time,
"update_time": order.update_time,
}
for order in orders
]
__all__ = ["WorkOrderMemory"]