feat:客户agent以及记忆模块功能开发
This commit is contained in:
@@ -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",
|
||||
]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
"""记忆模块依赖健康检查。"""
|
||||
|
||||
from config.database import check_ready_detail
|
||||
|
||||
|
||||
async def check_memory_dependencies() -> dict[str, dict]:
|
||||
"""复用项目统一数据库健康检查,避免记忆模块重复管理连接。"""
|
||||
return await check_ready_detail()
|
||||
@@ -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)
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
Reference in New Issue
Block a user