Files
ai-video-fullstack/backend/services/pipecat/processors.py

639 lines
23 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.
"""Reusable frame processors shared by cascade and realtime pipelines."""
import asyncio
from collections.abc import Callable
from dataclasses import dataclass
from uuid import uuid4
from loguru import logger
from models import AssistantConfig
from services.brains import Brain
from services.conversation_history import ConversationRecorder
from services.knowledge import search as search_knowledge
from db.session import SessionLocal
from pipecat.frames.frames import (
BotStartedSpeakingFrame,
BotStoppedSpeakingFrame,
FunctionCallCancelFrame,
FunctionCallResultFrame,
FunctionCallsStartedFrame,
InputTransportMessageFrame,
InterruptionFrame,
LLMContextFrame,
LLMFullResponseEndFrame,
LLMFullResponseStartFrame,
LLMTextFrame,
OutputTransportMessageUrgentFrame,
TextFrame,
UserImageRawFrame,
UserImageRequestFrame,
)
from pipecat.processors.aggregators.llm_response_universal import (
LLMAssistantAggregator,
)
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
from pipecat.turns.user_mute.base_user_mute_strategy import BaseUserMuteStrategy
from pipecat.utils.time import time_now_iso8601
KNOWLEDGE_CONTEXT_MARKER = "<!-- knowledge-context -->"
@dataclass(frozen=True)
class UserInput:
"""Validated application input submitted as one user turn."""
input_id: str
text: str
has_camera_frame: bool
run_immediately: bool
interrupt: bool
@property
def prompt_text(self) -> str:
if self.text:
return self.text
return "请根据用户刚提交的图片进行回复。"
@property
def transcript_text(self) -> str:
if not self.has_camera_frame:
return self.text
return f"{self.text}\n已发送一张图片".strip()
class UserInputError(ValueError):
def __init__(self, message: str, *, input_id: str = "") -> None:
super().__init__(message)
self.input_id = input_id
def parse_user_input(message) -> UserInput | None:
"""Parse the sole public user-input wire format."""
if not isinstance(message, dict):
return None
if message.get("type") != "user-input":
return None
input_id = str(message.get("input_id") or "").strip()
if not input_id or len(input_id) > 128:
raise UserInputError("user-input 缺少有效的 input_id", input_id=input_id)
if message.get("schema_version") != 1:
raise UserInputError("不支持的 user-input schema_version", input_id=input_id)
parts = message.get("parts")
if not isinstance(parts, list) or not parts:
raise UserInputError("user-input parts 不能为空", input_id=input_id)
text = ""
has_camera_frame = False
for part in parts:
if not isinstance(part, dict):
raise UserInputError("user-input part 格式不正确", input_id=input_id)
part_type = part.get("type")
if part_type == "input_text":
if text:
raise UserInputError("P0 只支持一个 input_text", input_id=input_id)
text = str(part.get("text") or "").strip()
if not text:
raise UserInputError("input_text 不能为空", input_id=input_id)
elif part_type == "input_image":
if has_camera_frame:
raise UserInputError("P0 只支持一张图片", input_id=input_id)
source = part.get("source")
if (
not isinstance(source, dict)
or source.get("type") != "camera_frame"
or source.get("frame") != "current"
):
raise UserInputError(
"P0 只支持当前摄像头画面",
input_id=input_id,
)
has_camera_frame = True
else:
raise UserInputError(f"不支持的输入类型: {part_type}", input_id=input_id)
if not text and not has_camera_frame:
raise UserInputError("user-input 没有有效内容", input_id=input_id)
options = message.get("options")
options = options if isinstance(options, dict) else {}
run_immediately = bool(options.get("run_immediately", True))
interrupt = bool(options.get("interrupt", True))
return UserInput(
input_id=input_id,
text=text,
has_camera_frame=has_camera_frame,
run_immediately=run_immediately,
interrupt=interrupt,
)
def message_text(message: dict) -> str:
"""Extract textual content from plain or multimodal LLM messages."""
content = message.get("content")
if isinstance(content, str):
return content.strip()
if not isinstance(content, list):
return ""
return "\n".join(
str(part.get("text") or "").strip()
for part in content
if isinstance(part, dict)
and part.get("type") in {"text", "input_text"}
and str(part.get("text") or "").strip()
).strip()
class UserInputProcessor(FrameProcessor):
"""Validate app inputs and coordinate their interruption boundary."""
def __init__(self, should_ignore_input: Callable[[], bool] | None = None):
super().__init__()
self._should_ignore_input = should_ignore_input or (lambda: False)
self._register_event_handler("on_user_input")
self._register_event_handler("on_client_ready")
async def process_frame(self, frame, direction: FrameDirection):
await super().process_frame(frame, direction)
if not isinstance(frame, InputTransportMessageFrame):
await self.push_frame(frame, direction)
return
if isinstance(frame.message, dict) and frame.message.get("type") == "client-ready":
await self._call_event_handler("on_client_ready")
return
try:
user_input = parse_user_input(frame.message)
except UserInputError as exc:
await self._emit_result(exc.input_id, "error", str(exc))
return
if user_input is None:
await self.push_frame(frame, direction)
return
if self._should_ignore_input():
logger.debug("通话正在结束,忽略后续文字输入")
await self._emit_result(user_input.input_id, "error", "当前不能接收新的用户输入")
return
await self._call_event_handler("on_user_input", user_input)
if user_input.run_immediately and user_input.interrupt:
await self.broadcast_interruption()
async def _emit_result(self, input_id: str, status: str, message: str = "") -> None:
await self.push_frame(
OutputTransportMessageUrgentFrame(
message={
"type": "user-input-result",
"input_id": input_id,
"status": status,
**({"message": message} if message else {}),
}
)
)
class CallEndingUserMuteStrategy(BaseUserMuteStrategy):
"""Keep user media muted after an end-call tool starts terminating a call."""
def __init__(self, is_call_ending: Callable[[], bool]):
super().__init__()
self._is_call_ending = is_call_ending
async def process_frame(self, frame) -> bool:
await super().process_frame(frame)
return self._is_call_ending()
@dataclass
class _BlockedToolCall:
tool_finished: bool = False
response_started: bool = False
response_finished: bool = False
class ToolInterruptionUserMuteStrategy(BaseUserMuteStrategy):
"""Mute user media for selected tools and their associated bot response.
Pipecat's ``cancel_on_interruption`` also controls whether a tool is
synchronous or asynchronous. This strategy keeps the independent product
setting ("allow user interruptions") at the user-input boundary.
"""
def __init__(
self,
blocked_tools: dict[str, str] | set[str],
):
super().__init__()
self._blocked_tools = (
dict(blocked_tools)
if isinstance(blocked_tools, dict)
else {name: "immediate" for name in blocked_tools}
)
self._calls: dict[str, _BlockedToolCall] = {}
self._bot_speaking = False
@property
def is_muted(self) -> bool:
return bool(self._calls)
async def process_frame(self, frame) -> bool:
await super().process_frame(frame)
if isinstance(frame, BotStartedSpeakingFrame):
self._bot_speaking = True
for call in self._calls.values():
call.response_started = True
elif isinstance(frame, BotStoppedSpeakingFrame):
self._bot_speaking = False
for call in self._calls.values():
if call.response_started:
call.response_finished = True
self._release_completed()
elif isinstance(frame, FunctionCallsStartedFrame):
for call in frame.function_calls:
execution_mode = self._blocked_tools.get(call.function_name)
if execution_mode:
self._calls[call.tool_call_id] = _BlockedToolCall(
# An async tool may continue within the current model
# turn. An immediate tool's current speech is a preamble;
# its protected follow-up starts only after the result.
response_started=(
self._bot_speaking and execution_mode == "async"
),
)
elif isinstance(frame, FunctionCallResultFrame):
call = self._calls.get(frame.tool_call_id)
if call is not None:
call.tool_finished = True
self._release_completed()
elif isinstance(frame, FunctionCallCancelFrame):
# A canceled call cannot reliably produce a follow-up response.
self._calls.pop(frame.tool_call_id, None)
return self.is_muted
def _release_completed(self) -> None:
completed = [
call_id
for call_id, state in self._calls.items()
if state.tool_finished and state.response_finished
]
for call_id in completed:
self._calls.pop(call_id, None)
class VisionCaptureProcessor(FrameProcessor):
"""Capture one requested video frame for auxiliary vision-model analysis."""
def __init__(self, timeout_s: float = 3.0):
super().__init__()
self._timeout_s = timeout_s
self._pending: dict[str, asyncio.Future[UserImageRawFrame]] = {}
async def request_image(
self,
requester: FrameProcessor,
request: UserImageRequestFrame,
) -> UserImageRawFrame:
key = request.tool_call_id or str(uuid4())
request.tool_call_id = key
request.append_to_context = False
request.result_callback = None
loop = asyncio.get_running_loop()
future: asyncio.Future[UserImageRawFrame] = loop.create_future()
self._pending[key] = future
await requester.push_frame(request, FrameDirection.UPSTREAM)
try:
return await asyncio.wait_for(future, timeout=self._timeout_s)
finally:
self._pending.pop(key, None)
async def process_frame(self, frame, direction: FrameDirection):
await super().process_frame(frame, direction)
if (
isinstance(frame, UserImageRawFrame)
and frame.request
and frame.request.tool_call_id
and frame.request.tool_call_id in self._pending
):
future = self._pending[frame.request.tool_call_id]
if not future.done():
future.set_result(frame)
return
await self.push_frame(frame, direction)
class RealtimeDynamicVariableProcessor(FrameProcessor):
"""Keep realtime system turn/history variables current between responses."""
def __init__(self, brain: Brain, cfg: AssistantConfig, realtime):
super().__init__()
self._brain = brain
self._cfg = cfg
self._realtime = realtime
async def _refresh_instructions(self) -> None:
update = getattr(self._realtime, "update_instructions", None)
if callable(update):
await update(self._brain.system_prompt(self._cfg))
async def process_frame(self, frame, direction: FrameDirection):
await super().process_frame(frame, direction)
if isinstance(frame, OutputTransportMessageUrgentFrame):
message = frame.message
if isinstance(message, dict):
event_type = message.get("type")
if event_type == "transcript" and message.get("role") == "user":
content = str(message.get("content") or "").strip()
if content:
self._brain.record_user_message(content)
await self._refresh_instructions()
elif event_type == "assistant-text-end":
await self._brain.on_assistant_text_end(
str(message.get("turn_id") or ""),
str(message.get("content") or ""),
bool(message.get("interrupted", False)),
)
await self._refresh_instructions()
await self.push_frame(frame, direction)
class RealtimeUserInputProcessor(FrameProcessor):
"""Route text-only user-input messages to a realtime service."""
def __init__(self):
super().__init__()
self._register_event_handler("on_user_input")
async def process_frame(self, frame, direction: FrameDirection):
await super().process_frame(frame, direction)
if not isinstance(frame, InputTransportMessageFrame):
await self.push_frame(frame, direction)
return
try:
user_input = parse_user_input(frame.message)
except UserInputError as exc:
await self._emit_error(exc.input_id, str(exc))
return
if user_input is None:
await self.push_frame(frame, direction)
return
if user_input.has_camera_frame:
await self._emit_error(
user_input.input_id,
"Realtime 模式暂不支持图片输入",
)
return
await self._call_event_handler("on_user_input", user_input)
async def _emit_error(self, input_id: str, message: str) -> None:
await self.push_frame(
OutputTransportMessageUrgentFrame(
message={
"type": "user-input-result",
"input_id": input_id,
"status": "error",
"message": message,
}
)
)
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 KnowledgeRetrievalProcessor(FrameProcessor):
"""Retrieve before local LLM inference without changing Pipecat internals."""
def __init__(
self,
knowledge_base_id: str | None,
top_n: int = 5,
score_threshold: float = 0.0,
):
super().__init__()
self._knowledge_base_id = knowledge_base_id
self._top_n = top_n
self._score_threshold = score_threshold
self._mode = "automatic" if knowledge_base_id else "disabled"
self._last_signature = ""
def set_scope(self, scope: dict) -> None:
self._knowledge_base_id = scope.get("knowledge_base_id") or None
self._mode = str(scope.get("mode") or "disabled")
self._top_n = int(scope.get("top_n") or 5)
self._score_threshold = float(scope.get("score_threshold") or 0.0)
self._last_signature = ""
def _clear_context(self, messages: list[dict]) -> None:
# Remove the legacy Workflow knowledge message so an in-flight context
# created before this compatibility fix cannot keep sending that role.
messages[:] = [
message
for message in messages
if not (
message.get("role") == "developer"
and KNOWLEDGE_CONTEXT_MARKER in str(message.get("content") or "")
)
]
system_message = next(
(message for message in messages if message.get("role") == "system"),
None,
)
if system_message is not None:
content = str(system_message.get("content") or "")
system_message["content"] = content.split(KNOWLEDGE_CONTEXT_MARKER, 1)[0].rstrip()
def _set_context(self, messages: list[dict], block: str) -> None:
"""Store retrieved knowledge in a provider-compatible system message."""
self._clear_context(messages)
system_message = next(
(message for message in messages if message.get("role") == "system"),
None,
)
if system_message is None:
messages.insert(0, {"role": "system", "content": block})
return
content = str(system_message.get("content") or "").rstrip()
system_message["content"] = f"{content}\n\n{block}" if content else block
async def process_frame(self, frame, direction: FrameDirection):
await super().process_frame(frame, direction)
if not isinstance(frame, LLMContextFrame):
await self.push_frame(frame, direction)
return
messages = frame.context.get_messages()
if self._mode != "automatic" or not self._knowledge_base_id:
self._clear_context(messages)
await self.push_frame(frame, direction)
return
user_messages = [message for message in messages if message.get("role") == "user"]
if not user_messages:
await self.push_frame(frame, direction)
return
query = message_text(user_messages[-1])
signature = f"{len(user_messages)}:{query}"
if not query or signature == self._last_signature:
await self.push_frame(frame, direction)
return
self._last_signature = signature
try:
async with SessionLocal() as session:
results = await search_knowledge(
session,
self._knowledge_base_id,
query,
top_k=self._top_n,
score_threshold=self._score_threshold,
)
except Exception as exc:
logger.warning(f"自动知识库检索失败: {exc}")
results = []
sources = "\n\n".join(
f"[{index + 1}] 来源:{item['document']}(相关度 {item['score']}\n{item['content']}"
for index, item in enumerate(results)
) or "未检索到相关资料。"
block = f"{KNOWLEDGE_CONTEXT_MARKER}\n当前问题的知识库检索结果:\n{sources}"
self._set_context(messages, block)
await self.push_frame(frame, direction)
class UserTurnRoutingProcessor(FrameProcessor):
"""Give a brain first right of refusal before a new user turn reaches the LLM."""
def __init__(self, brain: Brain):
super().__init__()
self._brain = brain
self._last_user_message: dict | None = None
async def process_frame(self, frame, direction: FrameDirection):
await super().process_frame(frame, direction)
if direction != FrameDirection.DOWNSTREAM or not isinstance(
frame, LLMContextFrame
):
await self.push_frame(frame, direction)
return
user_message = next(
(
message
for message in reversed(frame.context.get_messages())
if message.get("role") == "user"
and message_text(message)
),
None,
)
if user_message is None:
await self.push_frame(frame, direction)
return
if user_message is self._last_user_message:
# Programmatic LLMRunFrame after a node transition reuses the same
# user message. It is a response run, not another routing event.
await self.push_frame(frame, direction)
return
self._last_user_message = user_message
content = message_text(user_message)
handled = await self._brain.on_user_turn_end(content)
if not handled:
await self.push_frame(frame, direction)
class PassthroughLLMAssistantAggregator(LLMAssistantAggregator):
"""聚合 LLM 回复进上下文,同时继续把回复帧交给下游 TTS。"""
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self._register_event_handler("on_interruption_processed")
self._register_event_handler("on_assistant_text_start")
self._register_event_handler("on_assistant_text_delta")
self._register_event_handler("on_assistant_text_end")
self._stream_turn_id: str | None = None
self._stream_timestamp = ""
self._stream_text = ""
async def process_frame(self, frame, direction: FrameDirection):
await super().process_frame(frame, direction)
if isinstance(frame, LLMFullResponseStartFrame):
self._stream_turn_id = uuid4().hex
self._stream_timestamp = time_now_iso8601()
self._stream_text = ""
await self._call_event_handler(
"on_assistant_text_start",
self._stream_turn_id,
self._stream_timestamp,
)
elif isinstance(frame, LLMTextFrame) and self._stream_turn_id:
self._stream_text += frame.text
await self._call_event_handler(
"on_assistant_text_delta",
self._stream_turn_id,
frame.text,
)
elif isinstance(frame, LLMFullResponseEndFrame):
await self._finish_text_stream(interrupted=False)
# LLMAssistantAggregator 默认会消费这些帧。放在 TTS 前用于中断时保存
# 已生成前缀时,必须显式透传,否则 TTS 收不到任何 LLM 回复。
if isinstance(
frame,
(LLMFullResponseStartFrame, LLMFullResponseEndFrame, TextFrame),
):
await self.push_frame(frame, direction)
elif isinstance(frame, InterruptionFrame):
await self._finish_text_stream(interrupted=True)
await self._call_event_handler("on_interruption_processed")
async def _finish_text_stream(self, *, interrupted: bool):
if not self._stream_turn_id:
return
await self._call_event_handler(
"on_assistant_text_end",
self._stream_turn_id,
self._stream_text,
interrupted,
)
self._stream_turn_id = None
self._stream_timestamp = ""
self._stream_text = ""
class WorkflowAggregatorPair:
"""Small public-shape adapter required by Pipecat FlowManager."""
def __init__(self, user_aggregator, assistant_aggregator):
self._user = user_aggregator
self._assistant = assistant_aggregator
def user(self):
return self._user
def assistant(self):
return self._assistant