Add conversation history management and API endpoints
- Introduce new database models for conversation sessions, messages, and artifacts to support conversation history tracking. - Implement API routes for listing conversations and retrieving detailed conversation data, enhancing user interaction with historical records. - Add a conversation recorder service to persist conversation messages in real-time without disrupting ongoing calls. - Update the frontend to display conversation history, including filtering and sorting options, improving user experience. - Enhance the pipeline to integrate conversation history recording seamlessly during interactions.
This commit is contained in:
@@ -23,6 +23,7 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||
from routes import (
|
||||
assistants,
|
||||
auth,
|
||||
conversations,
|
||||
health,
|
||||
knowledge_bases,
|
||||
model_registry,
|
||||
@@ -52,6 +53,7 @@ app.add_middleware(
|
||||
|
||||
app.include_router(health.router)
|
||||
app.include_router(auth.router)
|
||||
app.include_router(conversations.router)
|
||||
app.include_router(assistants.router)
|
||||
app.include_router(knowledge_bases.router)
|
||||
app.include_router(model_registry.router)
|
||||
|
||||
@@ -9,7 +9,17 @@
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import JSON, Boolean, DateTime, ForeignKey, String, func
|
||||
from sqlalchemy import (
|
||||
JSON,
|
||||
Boolean,
|
||||
DateTime,
|
||||
ForeignKey,
|
||||
Integer,
|
||||
String,
|
||||
Text,
|
||||
UniqueConstraint,
|
||||
func,
|
||||
)
|
||||
from sqlalchemy.dialects.postgresql import JSONB
|
||||
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
|
||||
|
||||
@@ -187,3 +197,83 @@ class AssistantToolBinding(Base):
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), server_default=func.now()
|
||||
)
|
||||
|
||||
|
||||
class ConversationSession(Base):
|
||||
"""一次完整连接对应的对话会话;媒体文件以后作为 artifact 关联进来。"""
|
||||
|
||||
__tablename__ = "conversation_sessions"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(40), primary_key=True)
|
||||
assistant_id: Mapped[str | None] = mapped_column(
|
||||
String(40),
|
||||
ForeignKey("assistants.id", ondelete="SET NULL"),
|
||||
nullable=True,
|
||||
index=True,
|
||||
)
|
||||
assistant_name: Mapped[str] = mapped_column(String(128), default="")
|
||||
channel: Mapped[str] = mapped_column(String(24), index=True)
|
||||
runtime_mode: Mapped[str] = mapped_column(String(16), default="pipeline")
|
||||
status: Mapped[str] = mapped_column(String(16), index=True, default="active")
|
||||
message_count: Mapped[int] = mapped_column(Integer, default=0)
|
||||
started_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), server_default=func.now(), index=True
|
||||
)
|
||||
ended_at: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True), nullable=True
|
||||
)
|
||||
extra: Mapped[dict] = mapped_column(JSONB, default=dict)
|
||||
|
||||
|
||||
class ConversationMessage(Base):
|
||||
"""会话中的有序消息;当前只写 text,结构预留其他内容类型。"""
|
||||
|
||||
__tablename__ = "conversation_messages"
|
||||
__table_args__ = (
|
||||
UniqueConstraint(
|
||||
"session_id", "sequence", name="uq_conversation_message_sequence"
|
||||
),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(40), primary_key=True)
|
||||
session_id: Mapped[str] = mapped_column(
|
||||
String(40),
|
||||
ForeignKey("conversation_sessions.id", ondelete="CASCADE"),
|
||||
index=True,
|
||||
)
|
||||
sequence: Mapped[int] = mapped_column(Integer)
|
||||
role: Mapped[str] = mapped_column(String(16))
|
||||
content_type: Mapped[str] = mapped_column(String(16), default="text")
|
||||
content: Mapped[str] = mapped_column(Text, default="")
|
||||
occurred_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), server_default=func.now()
|
||||
)
|
||||
extra: Mapped[dict] = mapped_column(JSONB, default=dict)
|
||||
|
||||
|
||||
class ConversationArtifact(Base):
|
||||
"""未来的音频、视频、图片等大对象索引;二进制本体不进入数据库。"""
|
||||
|
||||
__tablename__ = "conversation_artifacts"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(40), primary_key=True)
|
||||
session_id: Mapped[str] = mapped_column(
|
||||
String(40),
|
||||
ForeignKey("conversation_sessions.id", ondelete="CASCADE"),
|
||||
index=True,
|
||||
)
|
||||
message_id: Mapped[str | None] = mapped_column(
|
||||
String(40),
|
||||
ForeignKey("conversation_messages.id", ondelete="SET NULL"),
|
||||
nullable=True,
|
||||
index=True,
|
||||
)
|
||||
kind: Mapped[str] = mapped_column(String(16), index=True)
|
||||
storage_uri: Mapped[str] = mapped_column(String(1024))
|
||||
mime_type: Mapped[str] = mapped_column(String(128), default="")
|
||||
size_bytes: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
duration_ms: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), server_default=func.now()
|
||||
)
|
||||
extra: Mapped[dict] = mapped_column(JSONB, default=dict)
|
||||
|
||||
@@ -0,0 +1,81 @@
|
||||
"""add conversation history
|
||||
|
||||
Revision ID: 20260710_0003
|
||||
Revises: 20260710_0002
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
|
||||
revision: str = "20260710_0003"
|
||||
down_revision: str | Sequence[str] | None = "20260710_0002"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"conversation_sessions",
|
||||
sa.Column("id", sa.String(length=40), nullable=False),
|
||||
sa.Column("assistant_id", sa.String(length=40), nullable=True),
|
||||
sa.Column("assistant_name", sa.String(length=128), server_default="", nullable=False),
|
||||
sa.Column("channel", sa.String(length=24), nullable=False),
|
||||
sa.Column("runtime_mode", sa.String(length=16), server_default="pipeline", nullable=False),
|
||||
sa.Column("status", sa.String(length=16), server_default="active", nullable=False),
|
||||
sa.Column("message_count", sa.Integer(), server_default="0", nullable=False),
|
||||
sa.Column("started_at", sa.DateTime(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.Column("ended_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("extra", postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
|
||||
sa.ForeignKeyConstraint(["assistant_id"], ["assistants.id"], ondelete="SET NULL"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_conversation_sessions_assistant_id", "conversation_sessions", ["assistant_id"])
|
||||
op.create_index("ix_conversation_sessions_channel", "conversation_sessions", ["channel"])
|
||||
op.create_index("ix_conversation_sessions_status", "conversation_sessions", ["status"])
|
||||
op.create_index("ix_conversation_sessions_started_at", "conversation_sessions", ["started_at"])
|
||||
|
||||
op.create_table(
|
||||
"conversation_messages",
|
||||
sa.Column("id", sa.String(length=40), nullable=False),
|
||||
sa.Column("session_id", sa.String(length=40), nullable=False),
|
||||
sa.Column("sequence", sa.Integer(), nullable=False),
|
||||
sa.Column("role", sa.String(length=16), nullable=False),
|
||||
sa.Column("content_type", sa.String(length=16), server_default="text", nullable=False),
|
||||
sa.Column("content", sa.Text(), server_default="", nullable=False),
|
||||
sa.Column("occurred_at", sa.DateTime(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.Column("extra", postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
|
||||
sa.ForeignKeyConstraint(["session_id"], ["conversation_sessions.id"], ondelete="CASCADE"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("session_id", "sequence", name="uq_conversation_message_sequence"),
|
||||
)
|
||||
op.create_index("ix_conversation_messages_session_id", "conversation_messages", ["session_id"])
|
||||
|
||||
op.create_table(
|
||||
"conversation_artifacts",
|
||||
sa.Column("id", sa.String(length=40), nullable=False),
|
||||
sa.Column("session_id", sa.String(length=40), nullable=False),
|
||||
sa.Column("message_id", sa.String(length=40), nullable=True),
|
||||
sa.Column("kind", sa.String(length=16), nullable=False),
|
||||
sa.Column("storage_uri", sa.String(length=1024), nullable=False),
|
||||
sa.Column("mime_type", sa.String(length=128), server_default="", nullable=False),
|
||||
sa.Column("size_bytes", sa.Integer(), nullable=True),
|
||||
sa.Column("duration_ms", sa.Integer(), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.Column("extra", postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
|
||||
sa.ForeignKeyConstraint(["message_id"], ["conversation_messages.id"], ondelete="SET NULL"),
|
||||
sa.ForeignKeyConstraint(["session_id"], ["conversation_sessions.id"], ondelete="CASCADE"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_conversation_artifacts_session_id", "conversation_artifacts", ["session_id"])
|
||||
op.create_index("ix_conversation_artifacts_message_id", "conversation_artifacts", ["message_id"])
|
||||
op.create_index("ix_conversation_artifacts_kind", "conversation_artifacts", ["kind"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("conversation_artifacts")
|
||||
op.drop_table("conversation_messages")
|
||||
op.drop_table("conversation_sessions")
|
||||
117
backend/routes/conversations.py
Normal file
117
backend/routes/conversations.py
Normal file
@@ -0,0 +1,117 @@
|
||||
"""对话历史查询 API。"""
|
||||
|
||||
from db.models import ConversationMessage, ConversationSession
|
||||
from db.session import get_session
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from schemas import (
|
||||
ConversationDetailOut,
|
||||
ConversationListOut,
|
||||
ConversationMessageOut,
|
||||
ConversationOut,
|
||||
)
|
||||
from services.auth import require_admin
|
||||
from sqlalchemy import func, or_, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
|
||||
router = APIRouter(
|
||||
prefix="/api/conversations",
|
||||
tags=["conversations"],
|
||||
dependencies=[Depends(require_admin)],
|
||||
)
|
||||
|
||||
|
||||
def _session_out(row: ConversationSession) -> ConversationOut:
|
||||
return ConversationOut(
|
||||
id=row.id,
|
||||
assistant_id=row.assistant_id,
|
||||
assistant_name=row.assistant_name,
|
||||
channel=row.channel,
|
||||
runtime_mode=row.runtime_mode,
|
||||
status=row.status,
|
||||
message_count=row.message_count,
|
||||
started_at=row.started_at,
|
||||
ended_at=row.ended_at,
|
||||
)
|
||||
|
||||
|
||||
@router.get("", response_model=ConversationListOut)
|
||||
async def list_conversations(
|
||||
page: int = Query(1, ge=1),
|
||||
page_size: int = Query(20, ge=1, le=100),
|
||||
assistant_id: str | None = None,
|
||||
search: str | None = Query(None, max_length=128),
|
||||
channel: str | None = None,
|
||||
status: str | None = None,
|
||||
sort_order: str = Query("newest", pattern="^(newest|oldest)$"),
|
||||
session: AsyncSession = Depends(get_session),
|
||||
):
|
||||
filters = []
|
||||
if assistant_id:
|
||||
filters.append(ConversationSession.assistant_id == assistant_id)
|
||||
if search and search.strip():
|
||||
keyword = f"%{search.strip()}%"
|
||||
filters.append(
|
||||
or_(
|
||||
ConversationSession.assistant_name.ilike(keyword),
|
||||
ConversationSession.id.ilike(keyword),
|
||||
)
|
||||
)
|
||||
if channel:
|
||||
filters.append(ConversationSession.channel == channel)
|
||||
if status:
|
||||
filters.append(ConversationSession.status == status)
|
||||
total = await session.scalar(
|
||||
select(func.count()).select_from(ConversationSession).where(*filters)
|
||||
)
|
||||
rows = (
|
||||
await session.execute(
|
||||
select(ConversationSession)
|
||||
.where(*filters)
|
||||
.order_by(
|
||||
ConversationSession.started_at.asc()
|
||||
if sort_order == "oldest"
|
||||
else ConversationSession.started_at.desc()
|
||||
)
|
||||
.offset((page - 1) * page_size)
|
||||
.limit(page_size)
|
||||
)
|
||||
).scalars().all()
|
||||
return ConversationListOut(
|
||||
items=[_session_out(row) for row in rows],
|
||||
total=int(total or 0),
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/{conversation_id}", response_model=ConversationDetailOut)
|
||||
async def get_conversation(
|
||||
conversation_id: str,
|
||||
session: AsyncSession = Depends(get_session),
|
||||
):
|
||||
conversation = await session.get(ConversationSession, conversation_id)
|
||||
if not conversation:
|
||||
raise HTTPException(404, "对话记录不存在")
|
||||
messages = (
|
||||
await session.execute(
|
||||
select(ConversationMessage)
|
||||
.where(ConversationMessage.session_id == conversation_id)
|
||||
.order_by(ConversationMessage.sequence)
|
||||
)
|
||||
).scalars().all()
|
||||
return ConversationDetailOut(
|
||||
**_session_out(conversation).model_dump(),
|
||||
messages=[
|
||||
ConversationMessageOut(
|
||||
id=message.id,
|
||||
sequence=message.sequence,
|
||||
role=message.role,
|
||||
content_type=message.content_type,
|
||||
content=message.content,
|
||||
occurred_at=message.occurred_at,
|
||||
extra=message.extra or {},
|
||||
)
|
||||
for message in messages
|
||||
],
|
||||
)
|
||||
@@ -120,7 +120,13 @@ async def _handle_offer(websocket, payload, peers):
|
||||
video_in_enabled=vision_enabled,
|
||||
)
|
||||
asyncio.create_task(
|
||||
run_pipeline(transport, cfg, vision_enabled=vision_enabled)
|
||||
run_pipeline(
|
||||
transport,
|
||||
cfg,
|
||||
vision_enabled=vision_enabled,
|
||||
assistant_id=offer.assistant_id,
|
||||
channel="webrtc",
|
||||
)
|
||||
)
|
||||
|
||||
answer = pc.get_answer()
|
||||
|
||||
@@ -24,13 +24,16 @@ from starlette.websockets import WebSocketDisconnect
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
async def _resolve_start_config(raw: str) -> AssistantConfig:
|
||||
async def _resolve_start_config(raw: str) -> tuple[AssistantConfig, str | None]:
|
||||
data = json.loads(raw)
|
||||
if data.get("assistant_id"):
|
||||
async with SessionLocal() as session:
|
||||
return await resolve_runtime_config(session, data["assistant_id"])
|
||||
return (
|
||||
await resolve_runtime_config(session, data["assistant_id"]),
|
||||
data["assistant_id"],
|
||||
)
|
||||
if data.get("inline_config"):
|
||||
return AssistantConfig(**data["inline_config"])
|
||||
return AssistantConfig(**data["inline_config"]), None
|
||||
raise ValueError("启动参数缺少 assistant_id 或 inline_config")
|
||||
|
||||
|
||||
@@ -43,10 +46,10 @@ async def voice_stream(websocket: WebSocket):
|
||||
return
|
||||
await websocket.accept()
|
||||
try:
|
||||
cfg = await _resolve_start_config(await websocket.receive_text())
|
||||
cfg, assistant_id = await _resolve_start_config(await websocket.receive_text())
|
||||
transport = build_ws_transport(websocket)
|
||||
# 直接 await:管线持续读这条 WS 的音频帧,直到对端断开
|
||||
await run_pipeline(transport, cfg)
|
||||
await run_pipeline(transport, cfg, assistant_id=assistant_id, channel="websocket")
|
||||
except WebSocketDisconnect:
|
||||
logger.info("WS 音频流断开")
|
||||
except Exception as e:
|
||||
|
||||
@@ -7,6 +7,7 @@ JSON 用 camelCase(modelId/interfaceType/apiUrl/apiKey),Python 内部用 snake_c
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Annotated, Any, Literal, Union
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
@@ -199,3 +200,37 @@ class ModelResourceTestResult(CamelModel):
|
||||
latency_ms: int | None = None
|
||||
message: str
|
||||
detail: str = ""
|
||||
|
||||
|
||||
# ---------- 对话历史 ----------
|
||||
class ConversationMessageOut(CamelModel):
|
||||
id: str
|
||||
sequence: int
|
||||
role: Literal["user", "assistant"]
|
||||
content_type: str
|
||||
content: str
|
||||
occurred_at: datetime
|
||||
extra: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class ConversationOut(CamelModel):
|
||||
id: str
|
||||
assistant_id: str | None
|
||||
assistant_name: str
|
||||
channel: str
|
||||
runtime_mode: str
|
||||
status: str
|
||||
message_count: int
|
||||
started_at: datetime
|
||||
ended_at: datetime | None
|
||||
|
||||
|
||||
class ConversationDetailOut(ConversationOut):
|
||||
messages: list[ConversationMessageOut] = Field(default_factory=list)
|
||||
|
||||
|
||||
class ConversationListOut(CamelModel):
|
||||
items: list[ConversationOut]
|
||||
total: int
|
||||
page: int
|
||||
page_size: int
|
||||
|
||||
137
backend/services/conversation_history.py
Normal file
137
backend/services/conversation_history.py
Normal file
@@ -0,0 +1,137 @@
|
||||
"""对话历史持久化。
|
||||
|
||||
只依赖管线已经发给客户端的最终文本事件,不侵入 Pipecat。媒体历史以后写入
|
||||
conversation_artifacts,并把对象存储地址关联到会话或消息。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from datetime import UTC, datetime
|
||||
from uuid import uuid4
|
||||
|
||||
from db.models import ConversationMessage, ConversationSession
|
||||
from db.session import SessionLocal
|
||||
from loguru import logger
|
||||
|
||||
|
||||
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._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,
|
||||
) -> "ConversationRecorder | None":
|
||||
session_id = 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={},
|
||||
)
|
||||
)
|
||||
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")
|
||||
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}"
|
||||
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(
|
||||
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:
|
||||
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}")
|
||||
@@ -17,6 +17,7 @@ from models import AssistantConfig
|
||||
from openai import AsyncOpenAI
|
||||
from PIL import Image
|
||||
from services.brains import build_brain
|
||||
from services.conversation_history import ConversationRecorder
|
||||
from services.pipecat.service_factory import (
|
||||
create_realtime_service,
|
||||
create_stt,
|
||||
@@ -298,6 +299,20 @@ class RealtimeTextInputProcessor(FrameProcessor):
|
||||
)
|
||||
|
||||
|
||||
class ConversationHistoryProcessor(FrameProcessor):
|
||||
"""从最终客户端事件旁路保存历史,不改变 Pipecat 的上下文与帧语义。"""
|
||||
|
||||
def __init__(self, recorder: ConversationRecorder | None):
|
||||
super().__init__()
|
||||
self._recorder = recorder
|
||||
|
||||
async def process_frame(self, frame, direction: FrameDirection):
|
||||
await super().process_frame(frame, direction)
|
||||
await self.push_frame(frame, direction)
|
||||
if self._recorder and isinstance(frame, OutputTransportMessageUrgentFrame):
|
||||
await self._recorder.record_transport_message(frame.message)
|
||||
|
||||
|
||||
class PassthroughLLMAssistantAggregator(LLMAssistantAggregator):
|
||||
"""聚合 LLM 回复进上下文,同时继续把回复帧交给下游 TTS。"""
|
||||
|
||||
@@ -363,6 +378,8 @@ async def run_pipeline(
|
||||
cfg: AssistantConfig,
|
||||
*,
|
||||
vision_enabled: bool = False,
|
||||
assistant_id: str | None = None,
|
||||
channel: str = "webrtc",
|
||||
) -> None:
|
||||
"""在给定 transport 上构建并运行管线,直到连接结束。
|
||||
|
||||
@@ -388,7 +405,12 @@ async def run_pipeline(
|
||||
if cfg.runtimeMode == "realtime":
|
||||
if vision_enabled:
|
||||
logger.warning("Realtime 模式暂未接入视频帧工具,本次仅启用语音通话")
|
||||
await run_realtime_pipeline(transport, cfg)
|
||||
await run_realtime_pipeline(
|
||||
transport,
|
||||
cfg,
|
||||
assistant_id=assistant_id,
|
||||
channel=channel,
|
||||
)
|
||||
return
|
||||
|
||||
stt = create_stt(cfg)
|
||||
@@ -677,6 +699,12 @@ async def run_pipeline(
|
||||
reason = str(call_end_state["reason"] or "completed")
|
||||
await queue_call_end(reason)
|
||||
|
||||
recorder = await ConversationRecorder.start(
|
||||
assistant_id=assistant_id,
|
||||
assistant_name=cfg.name,
|
||||
channel=channel,
|
||||
runtime_mode=cfg.runtimeMode,
|
||||
)
|
||||
pipeline = Pipeline(
|
||||
[
|
||||
transport.input(),
|
||||
@@ -691,6 +719,7 @@ async def run_pipeline(
|
||||
assistant_aggregator,
|
||||
tts,
|
||||
EndCallAfterSpeech(),
|
||||
ConversationHistoryProcessor(recorder),
|
||||
transport.output(),
|
||||
]
|
||||
)
|
||||
@@ -950,21 +979,42 @@ async def run_pipeline(
|
||||
await worker.queue_frame(EndFrame())
|
||||
|
||||
runner = WorkerRunner(handle_sigint=False)
|
||||
await runner.add_workers(worker)
|
||||
await runner.run()
|
||||
run_status = "completed"
|
||||
try:
|
||||
await runner.add_workers(worker)
|
||||
await runner.run()
|
||||
except Exception:
|
||||
run_status = "failed"
|
||||
raise
|
||||
finally:
|
||||
if recorder:
|
||||
await recorder.finish(status=run_status)
|
||||
logger.info("管线已结束")
|
||||
|
||||
|
||||
async def run_realtime_pipeline(transport, cfg: AssistantConfig) -> None:
|
||||
async def run_realtime_pipeline(
|
||||
transport,
|
||||
cfg: AssistantConfig,
|
||||
*,
|
||||
assistant_id: str | None = None,
|
||||
channel: str = "webrtc",
|
||||
) -> None:
|
||||
"""Run a speech-to-speech model that owns ASR, reasoning, and synthesis."""
|
||||
realtime = create_realtime_service(cfg)
|
||||
text_input = RealtimeTextInputProcessor()
|
||||
|
||||
recorder = await ConversationRecorder.start(
|
||||
assistant_id=assistant_id,
|
||||
assistant_name=cfg.name,
|
||||
channel=channel,
|
||||
runtime_mode=cfg.runtimeMode,
|
||||
)
|
||||
pipeline = Pipeline(
|
||||
[
|
||||
transport.input(),
|
||||
text_input,
|
||||
realtime,
|
||||
ConversationHistoryProcessor(recorder),
|
||||
transport.output(),
|
||||
]
|
||||
)
|
||||
@@ -1017,6 +1067,14 @@ async def run_realtime_pipeline(transport, cfg: AssistantConfig) -> None:
|
||||
await worker.queue_frame(EndFrame())
|
||||
|
||||
runner = WorkerRunner(handle_sigint=False)
|
||||
await runner.add_workers(worker)
|
||||
await runner.run()
|
||||
run_status = "completed"
|
||||
try:
|
||||
await runner.add_workers(worker)
|
||||
await runner.run()
|
||||
except Exception:
|
||||
run_status = "failed"
|
||||
raise
|
||||
finally:
|
||||
if recorder:
|
||||
await recorder.finish(status=run_status)
|
||||
logger.info("Realtime 管线已结束")
|
||||
|
||||
Reference in New Issue
Block a user