Files
Mutual_Fund/service/memory/archive.py
T

76 lines
2.5 KiB
Python
Raw Normal View History

"""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"]