270 lines
9.5 KiB
Python
270 lines
9.5 KiB
Python
"""Event registration for cascade and realtime conversation pipelines."""
|
|
|
|
from collections.abc import Awaitable, Callable
|
|
|
|
from loguru import logger
|
|
|
|
from pipecat.frames.frames import (
|
|
BotStartedSpeakingFrame,
|
|
BotStoppedSpeakingFrame,
|
|
EndFrame,
|
|
OutputTransportMessageUrgentFrame,
|
|
TTSSpeakFrame,
|
|
)
|
|
from pipecat.runner.utils import (
|
|
get_transport_client_id,
|
|
maybe_capture_participant_camera,
|
|
)
|
|
from pipecat.utils.time import time_now_iso8601
|
|
from services.pipecat.processors import UserInput
|
|
|
|
|
|
def bind_cascade_pipeline_events(
|
|
*,
|
|
transport,
|
|
worker,
|
|
brain,
|
|
context,
|
|
text_input,
|
|
user_aggregator,
|
|
assistant_aggregator,
|
|
greeting: str,
|
|
vision_enabled: bool,
|
|
vision_state: dict[str, str | None],
|
|
submit_user_input: Callable[[UserInput], Awaitable[None]] | None = None,
|
|
) -> None:
|
|
"""Connect processors to transport events without owning pipeline assembly."""
|
|
|
|
pending_user_inputs: list[UserInput] = []
|
|
greeting_transcript_sent = False
|
|
greeting_timestamp = ""
|
|
greeting_playback_pending = False
|
|
greeting_playback_started = False
|
|
|
|
# FlowManager already observes downstream frames for its own actions. Add
|
|
# to that filter instead of replacing it, then use the real transport
|
|
# playback boundary to release Workflow startup.
|
|
worker.add_reached_downstream_filter(
|
|
(BotStartedSpeakingFrame, BotStoppedSpeakingFrame)
|
|
)
|
|
|
|
@worker.event_handler("on_frame_reached_downstream")
|
|
async def on_frame_reached_downstream(_worker, frame):
|
|
nonlocal greeting_playback_pending, greeting_playback_started
|
|
if not greeting_playback_pending:
|
|
return
|
|
if isinstance(frame, BotStartedSpeakingFrame):
|
|
greeting_playback_started = True
|
|
return
|
|
if isinstance(frame, BotStoppedSpeakingFrame) and greeting_playback_started:
|
|
greeting_playback_pending = False
|
|
greeting_playback_started = False
|
|
await brain.on_greeting_finished()
|
|
|
|
async def queue_transcript(role: str, content: str, timestamp: str) -> None:
|
|
if not content:
|
|
return
|
|
await worker.queue_frame(
|
|
OutputTransportMessageUrgentFrame(
|
|
message={
|
|
"type": "transcript",
|
|
"role": role,
|
|
"content": content,
|
|
"timestamp": timestamp,
|
|
}
|
|
)
|
|
)
|
|
|
|
async def queue_input_result(
|
|
user_input: UserInput,
|
|
status: str,
|
|
message: str = "",
|
|
) -> None:
|
|
await worker.queue_frame(
|
|
OutputTransportMessageUrgentFrame(
|
|
message={
|
|
"type": "user-input-result",
|
|
"input_id": user_input.input_id,
|
|
"status": status,
|
|
**({"message": message} if message else {}),
|
|
}
|
|
)
|
|
)
|
|
|
|
async def finish_user_input(user_input: UserInput) -> None:
|
|
try:
|
|
if submit_user_input is None:
|
|
raise RuntimeError("用户输入提交器尚未配置")
|
|
await submit_user_input(user_input)
|
|
except Exception as exc: # noqa: BLE001 - input errors must reach the client
|
|
logger.warning(f"用户输入处理失败: {exc}")
|
|
await queue_input_result(user_input, "error", str(exc))
|
|
return
|
|
await queue_input_result(user_input, "accepted")
|
|
|
|
@user_aggregator.event_handler("on_user_turn_stopped")
|
|
async def on_user_turn_stopped(_aggregator, _strategy, message):
|
|
await queue_transcript("user", message.content, message.timestamp)
|
|
|
|
@assistant_aggregator.event_handler("on_assistant_text_start")
|
|
async def on_assistant_text_start(_aggregator, turn_id, timestamp):
|
|
await brain.on_assistant_text_start(turn_id)
|
|
await worker.queue_frame(
|
|
OutputTransportMessageUrgentFrame(
|
|
message={
|
|
"type": "assistant-text-start",
|
|
"turn_id": turn_id,
|
|
"timestamp": timestamp,
|
|
}
|
|
)
|
|
)
|
|
|
|
@assistant_aggregator.event_handler("on_assistant_text_delta")
|
|
async def on_assistant_text_delta(_aggregator, turn_id, delta):
|
|
await worker.queue_frame(
|
|
OutputTransportMessageUrgentFrame(
|
|
message={
|
|
"type": "assistant-text-delta",
|
|
"turn_id": turn_id,
|
|
"delta": delta,
|
|
}
|
|
)
|
|
)
|
|
|
|
@assistant_aggregator.event_handler("on_assistant_text_end")
|
|
async def on_assistant_text_end(_aggregator, turn_id, content, interrupted):
|
|
await worker.queue_frame(
|
|
OutputTransportMessageUrgentFrame(
|
|
message={
|
|
"type": "assistant-text-end",
|
|
"turn_id": turn_id,
|
|
"content": content,
|
|
"interrupted": interrupted,
|
|
}
|
|
)
|
|
)
|
|
await brain.on_assistant_text_end(turn_id, content, interrupted)
|
|
|
|
@text_input.event_handler("on_user_input")
|
|
async def on_user_input(_processor, user_input: UserInput):
|
|
await queue_transcript(
|
|
"user",
|
|
user_input.transcript_text,
|
|
time_now_iso8601(),
|
|
)
|
|
if user_input.run_immediately and user_input.interrupt:
|
|
pending_user_inputs.append(user_input)
|
|
return
|
|
await finish_user_input(user_input)
|
|
|
|
@assistant_aggregator.event_handler("on_interruption_processed")
|
|
async def on_interruption_processed(_aggregator):
|
|
if not pending_user_inputs:
|
|
return
|
|
await finish_user_input(pending_user_inputs.pop(0))
|
|
|
|
@text_input.event_handler("on_client_ready")
|
|
async def on_client_ready(_processor):
|
|
nonlocal greeting_transcript_sent
|
|
if greeting and not greeting_transcript_sent:
|
|
greeting_transcript_sent = True
|
|
await queue_transcript(
|
|
"assistant",
|
|
greeting,
|
|
greeting_timestamp or time_now_iso8601(),
|
|
)
|
|
await brain.on_client_ready()
|
|
|
|
@transport.event_handler("on_client_connected")
|
|
async def on_client_connected(_transport, _client):
|
|
nonlocal greeting_timestamp, greeting_playback_pending
|
|
if vision_enabled:
|
|
try:
|
|
vision_state["client_id"] = get_transport_client_id(
|
|
_transport,
|
|
_client,
|
|
)
|
|
await maybe_capture_participant_camera(_transport, _client)
|
|
logger.info(
|
|
f"视觉理解已接入视频客户端: {vision_state['client_id']}"
|
|
)
|
|
except Exception as exc: # noqa: BLE001 - media availability is optional
|
|
logger.warning(f"视觉理解摄像头捕获初始化失败: {exc}")
|
|
has_greeting = bool(greeting.strip())
|
|
if has_greeting:
|
|
# Preserve the actual playback order. The transcript is delivered
|
|
# later on client-ready, but the preview sorts by this timestamp.
|
|
greeting_timestamp = greeting_timestamp or time_now_iso8601()
|
|
if brain.spec.owns_context:
|
|
brain.prepare_greeting_context(greeting, context)
|
|
greeting_playback_pending = True
|
|
|
|
# Initialize the Workflow before the greeting is queued so a very
|
|
# short TTS response cannot finish before the brain arms its startup
|
|
# gate. Other brain types simply ignore greeting_pending.
|
|
await brain.on_connected(greeting_pending=has_greeting)
|
|
|
|
if has_greeting:
|
|
await worker.queue_frame(
|
|
TTSSpeakFrame(greeting, append_to_context=False)
|
|
)
|
|
|
|
@transport.event_handler("on_client_disconnected")
|
|
async def on_client_disconnected(_transport, _client):
|
|
logger.info("对端断开,结束管线")
|
|
await worker.queue_frame(EndFrame())
|
|
|
|
|
|
def bind_realtime_pipeline_events(
|
|
*,
|
|
transport,
|
|
worker,
|
|
realtime,
|
|
text_input,
|
|
greeting: str,
|
|
) -> None:
|
|
"""Connect text and lifecycle events for a realtime model pipeline."""
|
|
|
|
async def queue_transcript(role: str, content: str) -> None:
|
|
if not content:
|
|
return
|
|
await worker.queue_frame(
|
|
OutputTransportMessageUrgentFrame(
|
|
message={
|
|
"type": "transcript",
|
|
"role": role,
|
|
"content": content,
|
|
"timestamp": time_now_iso8601(),
|
|
}
|
|
)
|
|
)
|
|
|
|
@text_input.event_handler("on_user_input")
|
|
async def on_user_input(_processor, user_input: UserInput):
|
|
await queue_transcript("user", user_input.text)
|
|
if user_input.run_immediately and user_input.interrupt:
|
|
await realtime.interrupt()
|
|
await realtime.send_text(
|
|
user_input.text,
|
|
run_immediately=user_input.run_immediately,
|
|
)
|
|
await worker.queue_frame(
|
|
OutputTransportMessageUrgentFrame(
|
|
message={
|
|
"type": "user-input-result",
|
|
"input_id": user_input.input_id,
|
|
"status": "accepted",
|
|
}
|
|
)
|
|
)
|
|
|
|
@transport.event_handler("on_client_connected")
|
|
async def on_client_connected(_transport, _client):
|
|
if greeting:
|
|
await realtime.speak(greeting)
|
|
|
|
@transport.event_handler("on_client_disconnected")
|
|
async def on_client_disconnected(_transport, _client):
|
|
logger.info("Realtime 对端断开,结束管线")
|
|
await worker.queue_frame(EndFrame())
|