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

509 lines
19 KiB
Python

"""StepFun StepAudio realtime speech-to-speech Pipecat service."""
from __future__ import annotations
import asyncio
import base64
import json
from collections.abc import Awaitable, Callable
from typing import Any
from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit
from uuid import uuid4
from loguru import logger
from pipecat.frames.frames import (
CancelFrame,
EndFrame,
Frame,
InputAudioRawFrame,
InterruptionFrame,
LLMMessagesAppendFrame,
OutputTransportMessageUrgentFrame,
StartFrame,
TTSAudioRawFrame,
)
from pipecat.processors.frame_processor import FrameDirection
from pipecat.services.ai_service import AIService
from pipecat.services.settings import ServiceSettings
from pipecat.utils.time import time_now_iso8601
from websockets.asyncio.client import connect as websocket_connect
from websockets.protocol import State
from services.pipecat.realtime_tools import (
RealtimeTool,
RealtimeToolDispatcher,
RealtimeToolSession,
)
DEFAULT_STEPFUN_REALTIME_URL = "wss://api.stepfun.com/v1/realtime"
SpeechStartedHandler = Callable[[], Awaitable[None]]
class StepFunRealtimeService(AIService):
"""Bridge Pipecat audio frames to StepFun's Realtime WebSocket events."""
def __init__(
self,
*,
api_key: str,
model: str,
base_url: str = DEFAULT_STEPFUN_REALTIME_URL,
instructions: str = "",
voice: str = "linjiajiejie",
input_sample_rate: int = 24000,
output_sample_rate: int = 24000,
prefix_padding_ms: int = 500,
silence_duration_ms: int = 300,
energy_awakeness_threshold: int = 2500,
**kwargs,
) -> None:
super().__init__(settings=ServiceSettings(model=model), **kwargs)
self._api_key = api_key
self._model = model
self._base_url = base_url or DEFAULT_STEPFUN_REALTIME_URL
self._instructions = instructions
self._voice = voice
self._input_sample_rate = input_sample_rate
self._output_sample_rate = output_sample_rate
self._prefix_padding_ms = prefix_padding_ms
self._silence_duration_ms = silence_duration_ms
self._energy_awakeness_threshold = energy_awakeness_threshold
self._warned_input_sample_rate = False
self._websocket = None
self._receive_task: asyncio.Task | None = None
self._session_ready = asyncio.Event()
self._pending_events: list[dict[str, Any]] = []
self._assistant_turn_id: str | None = None
self._assistant_text = ""
self._assistant_timestamp = ""
self._tools: list[RealtimeTool] = []
self._tool_session = RealtimeToolSession(self._send_tool_event)
self._fixed_speech_completion: asyncio.Future[None] | None = None
self._suppress_response_transcript = False
self._speech_started_handler: SpeechStartedHandler | None = None
self._function_names: dict[str, str] = {}
async def start(self, frame: StartFrame) -> None:
await super().start(frame)
if not self._api_key or not self._model:
await self.push_error(
"StepFun Realtime requires api_key and model", fatal=True
)
return
await self._connect()
async def stop(self, frame: EndFrame) -> None:
await self._disconnect()
await super().stop(frame)
async def cancel(self, frame: CancelFrame) -> None:
await self._disconnect()
await super().cancel(frame)
async def cleanup(self) -> None:
await self._disconnect()
await super().cleanup()
async def process_frame(self, frame: Frame, direction: FrameDirection) -> None:
await super().process_frame(frame, direction)
if isinstance(frame, InputAudioRawFrame):
if (
frame.sample_rate != self._input_sample_rate
and not self._warned_input_sample_rate
):
self._warned_input_sample_rate = True
logger.warning(
"StepFun Realtime expected {} Hz input, received {} Hz",
self._input_sample_rate,
frame.sample_rate,
)
await self._send_event(
{
"type": "input_audio_buffer.append",
"audio": base64.b64encode(frame.audio).decode("ascii"),
}
)
return
if isinstance(frame, LLMMessagesAppendFrame):
for message in frame.messages:
text = self._message_text(message)
if text:
await self.send_text(text, run_immediately=frame.run_llm is not False)
return
if isinstance(frame, InterruptionFrame):
await self._send_event({"type": "response.cancel"}, wait_until_ready=False)
await self._finish_assistant_text(interrupted=True)
self._resolve_fixed_speech()
await self.push_frame(frame, direction)
async def send_text(self, text: str, *, run_immediately: bool = True) -> None:
await self._send_event(
{
"type": "conversation.item.create",
"item": {
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": text}],
},
}
)
if run_immediately:
await self._send_event({"type": "response.create"})
async def interrupt(self) -> None:
await self._send_event({"type": "response.cancel"}, wait_until_ready=False)
await self._finish_assistant_text(interrupted=True)
self._resolve_fixed_speech()
await self.broadcast_interruption()
async def request_response(self) -> None:
await self._send_event({"type": "response.create"})
def set_speech_started_handler(
self,
handler: SpeechStartedHandler | None,
) -> None:
self._speech_started_handler = handler
async def speak(self, text: str) -> None:
"""Ask the realtime model to voice a fixed greeting."""
await self.speak_fixed(text, suppress_transcript=False)
async def speak_fixed(
self,
text: str,
*,
suppress_transcript: bool = True,
) -> Awaitable[None] | None:
"""Speak configured text and expose the provider response boundary."""
if not text:
return None
completion = asyncio.get_running_loop().create_future()
self._resolve_fixed_speech()
self._fixed_speech_completion = completion
self._suppress_response_transcript = suppress_transcript
await self._send_event(
{
"type": "response.create",
"session": {
"instructions": f"请原样无修改地输出下面的话:\n{text}",
},
}
)
return completion
async def _connect(self) -> None:
if self._websocket and self._websocket.state is State.OPEN:
return
try:
self._websocket = await websocket_connect(
self._connection_url(),
additional_headers={"Authorization": f"Bearer {self._api_key}"},
max_size=None,
open_timeout=10,
)
self._receive_task = self.create_task(
self._receive_messages(), name="stepfun_realtime_receive"
)
except Exception as exc:
self._websocket = None
await self.push_error(
f"StepFun Realtime connection failed: {exc}",
exception=exc,
fatal=True,
)
async def _disconnect(self) -> None:
current_task = asyncio.current_task()
task = self._receive_task
self._receive_task = None
if task and task is not current_task:
await self.cancel_task(task)
websocket = self._websocket
self._websocket = None
self._session_ready.clear()
self._tool_session.clear()
self._function_names.clear()
self._resolve_fixed_speech()
if websocket and websocket.state is State.OPEN:
try:
await websocket.close()
except Exception:
pass
def _connection_url(self) -> str:
parts = urlsplit(self._base_url)
query = dict(parse_qsl(parts.query))
query["model"] = self._model
return urlunsplit(
(parts.scheme, parts.netloc, parts.path, urlencode(query), parts.fragment)
)
async def _receive_messages(self) -> None:
websocket = self._websocket
if not websocket:
return
try:
async for raw_message in websocket:
payload = json.loads(raw_message)
await self._handle_server_event(payload)
except Exception as exc:
if self._websocket is websocket:
await self.push_error(
f"StepFun Realtime receive failed: {exc}", exception=exc
)
finally:
if self._websocket is websocket:
self._websocket = None
self._session_ready.clear()
if self._receive_task is asyncio.current_task():
self._receive_task = None
async def _handle_server_event(self, event: dict[str, Any]) -> None:
event_type = event.get("type")
if event_type == "session.created":
await self._send_session_update()
elif event_type == "session.updated":
self._session_ready.set()
pending, self._pending_events = self._pending_events, []
for payload in pending:
await self._send_event(payload, wait_until_ready=False)
elif event_type == "response.audio.delta":
audio = event.get("delta")
if audio:
await self.push_frame(
TTSAudioRawFrame(
base64.b64decode(audio),
self._output_sample_rate,
1,
)
)
elif event_type in {"response.audio_transcript.delta", "response.text.delta"}:
if not self._suppress_response_transcript:
await self._append_assistant_text(str(event.get("delta") or ""))
elif event_type in {"response.audio_transcript.done", "response.text.done"}:
transcript = str(event.get("transcript") or event.get("text") or "")
if transcript and not self._suppress_response_transcript:
if not self._assistant_turn_id:
await self._append_assistant_text(transcript)
else:
self._assistant_text = transcript
await self._finish_assistant_text(interrupted=False)
elif event_type == "conversation.item.input_audio_transcription.completed":
await self._send_transcript("user", str(event.get("transcript") or ""))
elif event_type == "input_audio_buffer.speech_started":
await self._send_event({"type": "response.cancel"}, wait_until_ready=False)
await self.broadcast_interruption()
self._resolve_fixed_speech()
if self._speech_started_handler is not None:
await self._speech_started_handler()
elif event_type == "response.done":
response = event.get("response")
interrupted = isinstance(response, dict) and response.get("status") in {
"cancelled",
"incomplete",
"interrupted",
}
await self._finish_assistant_text(interrupted=interrupted)
self._resolve_fixed_speech()
elif event_type == "response.output_item.added":
self._remember_function_call(event)
elif event_type in {
"response.function_call_arguments.done",
"response.output_item.done",
}:
await self._handle_function_call_event(event)
elif event_type == "error":
error = event.get("error")
message = error.get("message") if isinstance(error, dict) else str(error)
if "cancel" not in str(message).lower():
await self.push_error(f"StepFun Realtime error: {message}")
self._resolve_fixed_speech()
async def _send_session_update(self) -> None:
await self._send_event(
{
"type": "session.update",
"session": {
"modalities": ["text", "audio"],
"instructions": self._instructions,
"voice": self._voice,
"input_audio_format": "pcm16",
"output_audio_format": "pcm16",
"turn_detection": {
"type": "server_vad",
"prefix_padding_ms": self._prefix_padding_ms,
"silence_duration_ms": self._silence_duration_ms,
"energy_awakeness_threshold": self._energy_awakeness_threshold,
},
"tools": [tool.provider_schema() for tool in self._tools],
"tool_choice": "auto",
},
},
wait_until_ready=False,
)
async def update_instructions(self, instructions: str) -> None:
"""Refresh model instructions without rebuilding the realtime session."""
self._instructions = instructions
if self._session_ready.is_set():
await self._send_event(
{
"type": "session.update",
"session": {"instructions": instructions},
},
wait_until_ready=False,
)
async def update_session(
self,
instructions: str,
tools: list[RealtimeTool],
) -> None:
"""Atomically replace the active Workflow prompt and tool catalog."""
self._instructions = instructions
self._tools = list(tools)
if self._session_ready.is_set():
await self._send_event(
{
"type": "session.update",
"session": {
"instructions": instructions,
"tools": [tool.provider_schema() for tool in tools],
"tool_choice": "auto",
},
},
wait_until_ready=False,
)
def set_tool_dispatcher(
self,
dispatcher: RealtimeToolDispatcher | None,
) -> None:
self._tool_session.set_dispatcher(dispatcher)
async def _send_tool_event(self, payload: dict[str, Any]) -> None:
await self._send_event(payload, wait_until_ready=False)
async def _handle_function_call_event(self, event: dict[str, Any]) -> None:
item = event.get("item")
source = item if isinstance(item, dict) else event
if isinstance(item, dict) and item.get("type") != "function_call":
return
call_id = str(
source.get("call_id")
or event.get("call_id")
or source.get("id")
or ""
)
name = str(
source.get("name")
or event.get("name")
or self._function_names.get(call_id)
or ""
)
if not name:
return
await self._tool_session.handle_call(
name=name,
call_id=call_id,
arguments=source.get("arguments", event.get("arguments")),
)
self._function_names.pop(call_id, None)
def _remember_function_call(self, event: dict[str, Any]) -> None:
item = event.get("item")
if not isinstance(item, dict) or item.get("type") != "function_call":
return
call_id = str(item.get("call_id") or item.get("id") or "")
name = str(item.get("name") or "")
if call_id and name:
self._function_names[call_id] = name
def _resolve_fixed_speech(self) -> None:
completion = self._fixed_speech_completion
self._fixed_speech_completion = None
self._suppress_response_transcript = False
if completion is not None and not completion.done():
completion.set_result(None)
async def _send_event(
self, payload: dict[str, Any], *, wait_until_ready: bool = True
) -> None:
if wait_until_ready and not self._session_ready.is_set():
self._pending_events.append(payload)
return
if not self._websocket or self._websocket.state is not State.OPEN:
return
payload = {"event_id": uuid4().hex, **payload}
await self._websocket.send(json.dumps(payload, ensure_ascii=False))
async def _append_assistant_text(self, delta: str) -> None:
if not delta:
return
if not self._assistant_turn_id:
self._assistant_turn_id = uuid4().hex
self._assistant_timestamp = time_now_iso8601()
await self._send_transport_message(
{
"type": "assistant-text-start",
"turn_id": self._assistant_turn_id,
"timestamp": self._assistant_timestamp,
}
)
self._assistant_text += delta
await self._send_transport_message(
{
"type": "assistant-text-delta",
"turn_id": self._assistant_turn_id,
"delta": delta,
}
)
async def _finish_assistant_text(self, *, interrupted: bool) -> None:
if not self._assistant_turn_id:
return
await self._send_transport_message(
{
"type": "assistant-text-end",
"turn_id": self._assistant_turn_id,
"content": self._assistant_text,
"interrupted": interrupted,
}
)
self._assistant_turn_id = None
self._assistant_text = ""
self._assistant_timestamp = ""
async def _send_transcript(self, role: str, content: str) -> None:
if content:
await self._send_transport_message(
{
"type": "transcript",
"role": role,
"content": content,
"timestamp": time_now_iso8601(),
}
)
async def _send_transport_message(self, message: dict[str, Any]) -> None:
await self.push_frame(OutputTransportMessageUrgentFrame(message=message))
@staticmethod
def _message_text(message: Any) -> str:
if not isinstance(message, dict):
return ""
content = message.get("content")
if isinstance(content, str):
return content.strip()
if isinstance(content, list):
return "\n".join(
str(part.get("text") or "")
for part in content
if isinstance(part, dict) and part.get("type") == "text"
).strip()
return ""