"""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 ""