Files
Mutual_Fund/service/memory/schemas.py
T

106 lines
3.7 KiB
Python

"""记忆模块跨层共享的数据契约。
这些模型只描述记忆模块与客服 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)
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)