Files
ai-video-fullstack/backend/services/conversation_history.py
2026-08-01 11:21:31 +08:00

183 lines
6.8 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""对话历史持久化。
只依赖管线已经发给客户端的最终文本事件,不侵入 Pipecat。媒体历史以后写入
conversation_artifacts并把对象存储地址关联到会话或消息。
"""
import asyncio
from copy import deepcopy
from datetime import UTC, datetime
from uuid import uuid4
from db.models import ConversationMessage, ConversationSession
from db.session import SessionLocal
from loguru import logger
MAX_WORKFLOW_TRACE_EVENTS = 1000
def _parse_timestamp(value: object) -> datetime:
if not isinstance(value, str) or not value:
return datetime.now(UTC)
try:
parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
return parsed if parsed.tzinfo else parsed.replace(tzinfo=UTC)
except ValueError:
return datetime.now(UTC)
class ConversationRecorder:
"""按事件顺序写入一通会话;写库失败不应中断实时通话。"""
def __init__(self, session_id: str):
self.session_id = session_id
self._sequence = 0
self._trace_sequence = 0
self._lock = asyncio.Lock()
self._seen_events: set[str] = set()
@classmethod
async def start(
cls,
*,
assistant_id: str | None,
assistant_name: str,
channel: str,
runtime_mode: str,
session_id: str | None = None,
extra: dict | None = None,
) -> "ConversationRecorder | None":
session_id = session_id or f"conv_{uuid4().hex[:20]}"
try:
async with SessionLocal() as db:
db.add(
ConversationSession(
id=session_id,
assistant_id=assistant_id,
assistant_name=assistant_name,
channel=channel,
runtime_mode=runtime_mode,
status="active",
message_count=0,
extra=deepcopy(extra or {}),
)
)
await db.commit()
return cls(session_id)
except Exception as exc:
logger.error(f"创建对话历史会话失败,不影响本次通话: {exc}")
return None
async def record_transport_message(self, message: object) -> None:
if not isinstance(message, dict):
return
event_type = message.get("type")
if event_type == "workflow-event":
await self._append_workflow_event(message)
return
role = ""
content = ""
extra: dict = {}
timestamp = message.get("timestamp")
event_key = ""
if event_type == "transcript":
role = str(message.get("role") or "")
content = str(message.get("content") or "").strip()
event_key = f"transcript:{role}:{timestamp}:{content}"
if message.get("source"):
extra["source"] = str(message["source"])
if message.get("nodeId"):
extra["node_id"] = str(message["nodeId"])
elif event_type == "assistant-text-end":
role = "assistant"
content = str(message.get("content") or "").strip()
turn_id = str(message.get("turn_id") or "")
event_key = f"assistant:{turn_id}"
extra = {
"turn_id": turn_id,
"interrupted": bool(message.get("interrupted", False)),
}
else:
return
if role not in {"user", "assistant"} or not content or event_key in self._seen_events:
return
self._seen_events.add(event_key)
await self._append(role, content, timestamp, extra)
async def _append_workflow_event(self, message: dict) -> None:
"""Persist bounded control-plane history without creating chat rows."""
event_id = str(message.get("eventId") or "")
event_name = str(message.get("event") or "")
event_key = f"workflow:{event_id}"
if not event_id or not event_name or event_key in self._seen_events:
return
self._seen_events.add(event_key)
async with self._lock:
next_sequence = self._trace_sequence + 1
payload = deepcopy(message)
payload["sessionId"] = self.session_id
payload["sequence"] = next_sequence
try:
async with SessionLocal() as db:
conversation = await db.get(ConversationSession, self.session_id)
if not conversation:
return
extra = dict(conversation.extra or {})
trace = list(extra.get("workflowTrace") or [])
trace.append(payload)
extra["workflowTrace"] = trace[-MAX_WORKFLOW_TRACE_EVENTS:]
conversation.extra = extra
await db.commit()
self._trace_sequence = next_sequence
except Exception as exc:
logger.error(f"保存 Workflow 轨迹失败,不影响本次通话: {exc}")
async def _append(
self,
role: str,
content: str,
timestamp: object,
extra: dict,
) -> None:
async with self._lock:
next_sequence = self._sequence + 1
try:
async with SessionLocal() as db:
db.add(
ConversationMessage(
id=f"msg_{uuid4().hex[:20]}",
session_id=self.session_id,
sequence=next_sequence,
role=role,
content_type="text",
content=content,
occurred_at=_parse_timestamp(timestamp),
extra=extra,
)
)
conversation = await db.get(ConversationSession, self.session_id)
if conversation:
conversation.message_count = next_sequence
await db.commit()
self._sequence = next_sequence
except Exception as exc:
logger.error(f"保存对话文本失败,不影响本次通话: {exc}")
async def finish(self, *, status: str = "completed") -> None:
async with self._lock:
try:
async with SessionLocal() as db:
conversation = await db.get(ConversationSession, self.session_id)
if conversation:
conversation.status = status
conversation.ended_at = datetime.now(UTC)
conversation.message_count = self._sequence
await db.commit()
except Exception as exc:
logger.error(f"结束对话历史会话失败: {exc}")