feat(workflow): add edge-tool routing and realtime runtime
This commit is contained in:
@@ -5,6 +5,7 @@ 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
|
||||
@@ -28,7 +29,14 @@ 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):
|
||||
@@ -68,6 +76,12 @@ class StepFunRealtimeService(AIService):
|
||||
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)
|
||||
@@ -120,6 +134,7 @@ class StepFunRealtimeService(AIService):
|
||||
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)
|
||||
|
||||
@@ -140,12 +155,35 @@ class StepFunRealtimeService(AIService):
|
||||
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
|
||||
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",
|
||||
@@ -154,6 +192,7 @@ class StepFunRealtimeService(AIService):
|
||||
},
|
||||
}
|
||||
)
|
||||
return completion
|
||||
|
||||
async def _connect(self) -> None:
|
||||
if self._websocket and self._websocket.state is State.OPEN:
|
||||
@@ -186,6 +225,9 @@ class StepFunRealtimeService(AIService):
|
||||
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()
|
||||
@@ -240,10 +282,11 @@ class StepFunRealtimeService(AIService):
|
||||
)
|
||||
)
|
||||
elif event_type in {"response.audio_transcript.delta", "response.text.delta"}:
|
||||
await self._append_assistant_text(str(event.get("delta") or ""))
|
||||
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:
|
||||
if transcript and not self._suppress_response_transcript:
|
||||
if not self._assistant_turn_id:
|
||||
await self._append_assistant_text(transcript)
|
||||
else:
|
||||
@@ -254,6 +297,9 @@ class StepFunRealtimeService(AIService):
|
||||
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 {
|
||||
@@ -262,11 +308,20 @@ class StepFunRealtimeService(AIService):
|
||||
"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(
|
||||
@@ -284,6 +339,8 @@ class StepFunRealtimeService(AIService):
|
||||
"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,
|
||||
@@ -293,7 +350,85 @@ class StepFunRealtimeService(AIService):
|
||||
"""Refresh model instructions without rebuilding the realtime session."""
|
||||
self._instructions = instructions
|
||||
if self._session_ready.is_set():
|
||||
await self._send_session_update()
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user