Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions riva/client/argparse_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -189,4 +189,9 @@ def add_connection_argparse_parameters(parser: argparse.ArgumentParser) -> argpa
def add_realtime_config_argparse_parameters(parser: argparse.ArgumentParser) -> argparse.ArgumentParser:
parser.add_argument("--endpoint", default="/v1/realtime", help="Endpoint to WebSocket server endpoint.")
parser.add_argument("--query-params", default="intent=transcription", help="Query parameters to WebSocket server endpoint.")
parser.add_argument(
"--mint-api-key",
default=None,
help="API key for POST /v1/realtime/*_sessions (Authorization: Bearer).",
)
return parser
176 changes: 124 additions & 52 deletions riva/client/realtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,116 @@
logger = logging.getLogger(__name__)


def _mint_auth_headers(args: argparse.Namespace) -> Dict[str, str]:
"""Build HTTP headers for session mint endpoints, including optional tier-1 auth."""
headers = {"Content-Type": "application/json"}
api_key = (getattr(args, "mint_api_key", None) or "").strip()
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
return headers


def _server_error_detail(response) -> str:
"""Extract the FastAPI ``detail`` message from an error response, falling back to raw text."""
try:
payload = response.json()
except ValueError:
return response.text.strip()
if isinstance(payload, dict) and payload.get("detail"):
return str(payload["detail"])
return response.text.strip()


def _raise_for_session_status(response, provided_key: bool) -> None:
"""Raise a clear, actionable error when the session mint request fails."""
if response.status_code == 200:
return
if response.status_code in (401, 403):
if provided_key:
hint = "The --mint-api-key value was rejected by the server. Verify it matches the server's REALTIME_AUTH_MINT_API_KEY."
else:
hint = "This server requires an API key on the session endpoint. Pass it with --mint-api-key <key>."
raise Exception(
f"Session authentication failed (HTTP {response.status_code}: {_server_error_detail(response)}). {hint}"
)
raise Exception(
f"Failed to initialize session. Status: {response.status_code}, Error: {_server_error_detail(response)}"
)


def _extract_client_secret_token(session_data: Dict[str, Any]) -> Optional[str]:
"""Return the ephemeral client_secret token when present, else ``None``.

Older servers that do not mint ``client_secret`` keep the legacy
unauthenticated WebSocket flow. New servers that return a secret require
it on ``Sec-WebSocket-Protocol``.
"""
client_secret = session_data.get("client_secret")
if not isinstance(client_secret, dict):
return None
token = client_secret.get("value")
if not token:
return None
return str(token)

TOKEN_SUBPROTO_PREFIX = "realtime-token."
Comment thread
virajkarandikar marked this conversation as resolved.
REALTIME_SUBPROTOCOL = "realtime"

def _build_websocket_url(
server: str,
endpoint: str,
query_params: str,
use_ssl: bool,
) -> str:
"""Build the WebSocket URL (intent only; auth token goes in Sec-WebSocket-Protocol)."""
query = query_params.strip()
scheme = "wss" if use_ssl else "ws"
if query:
return f"{scheme}://{server}{endpoint}?{query}"
return f"{scheme}://{server}{endpoint}"


def _websocket_subprotocols(token: str) -> List[str]:
"""Return Sec-WebSocket-Protocol entries for the realtime handshake.

Offers ``realtime`` plus ``realtime-token.<client_secret>``. The server
validates the token entry and echoes back ``realtime`` so the secret
never appears in the response headers.
"""
return [REALTIME_SUBPROTOCOL, f"{TOKEN_SUBPROTO_PREFIX}{token}"]


def _build_ssl_context(args: argparse.Namespace):
"""Build an optional SSL context for WebSocket connections."""
if not getattr(args, "use_ssl", False):
return None
ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
if getattr(args, "ssl_root_cert", None):
ssl_context.load_verify_locations(args.ssl_root_cert)
if getattr(args, "ssl_client_cert", None) and getattr(args, "ssl_client_key", None):
ssl_context.load_cert_chain(args.ssl_client_cert, args.ssl_client_key)
ssl_context.check_hostname = False
return ssl_context


async def _connect_websocket(
args: argparse.Namespace,
client_secret_token: Optional[str],
):
"""Open the realtime WebSocket, attaching subprotocol auth when available."""
ws_url = _build_websocket_url(
args.server,
args.endpoint,
args.query_params,
args.use_ssl,
)
ssl_context = _build_ssl_context(args)
connect_kwargs = {"ssl": ssl_context}
if client_secret_token:
connect_kwargs["subprotocols"] = _websocket_subprotocols(client_secret_token)
return await websockets.connect(ws_url, **connect_kwargs)


class RealtimeClientASR:
"""Client for real-time transcription via WebSocket connection."""

Expand All @@ -33,6 +143,7 @@ def __init__(self, args: argparse.Namespace):
self.args = args
self.websocket = None
self.session_config = None
self._client_secret_token: Optional[str] = None

# Input audio playback
self.input_audio_queue = queue.Queue()
Expand All @@ -50,9 +161,12 @@ async def connect(self):
# Initialize session via HTTP POST
session_data = await self._initialize_http_session()
self.session_config = session_data
self._client_secret_token = _extract_client_secret_token(session_data)

# Connect to WebSocket
await self._connect_websocket()
self.websocket = await _connect_websocket(
self.args, self._client_secret_token,
)
await self._initialize_session()

except requests.exceptions.RequestException as e:
Expand All @@ -67,7 +181,7 @@ async def connect(self):

async def _initialize_http_session(self) -> Dict[str, Any]:
"""Initialize session via HTTP POST request."""
headers = {"Content-Type": "application/json"}
headers = _mint_auth_headers(self.args)
uri = f"http://{self.args.server}/v1/realtime/transcription_sessions"
if self.args.use_ssl:
uri = f"https://{self.args.server}/v1/realtime/transcription_sessions"
Expand All @@ -80,37 +194,12 @@ async def _initialize_http_session(self) -> Dict[str, Any]:
verify=self.args.ssl_root_cert if self.args.ssl_root_cert else True
)

if response.status_code != 200:
raise Exception(
f"Failed to initialize session. Status: {response.status_code}, "
f"Error: {response.text}"
)
_raise_for_session_status(response, provided_key=bool(getattr(self.args, "mint_api_key", None)))

session_data = response.json()
logger.debug("Session initialized: %s", session_data)
return session_data

async def _connect_websocket(self):
"""Connect to WebSocket endpoint."""
ssl_context = None
ws_url = f"ws://{self.args.server}{self.args.endpoint}?{self.args.query_params}"
if self.args.use_ssl:
ws_url = f"wss://{self.args.server}{self.args.endpoint}?{self.args.query_params}"

ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
# Load a custom CA certificate bundle
if self.args.ssl_root_cert:
ssl_context.load_verify_locations(self.args.ssl_root_cert)
# Load a client certificate and key
if self.args.ssl_client_cert and self.args.ssl_client_key:
ssl_context.load_cert_chain(self.args.ssl_client_cert, self.args.ssl_client_key)
# Disable hostname verification
ssl_context.check_hostname = False
# ssl_context.verify_mode = ssl.CERT_REQUIRED

logger.debug("Connecting to WebSocket: %s", ws_url)
self.websocket = await websockets.connect(ws_url, ssl=ssl_context)

async def _initialize_session(self):
"""Initialize the WebSocket session."""
try:
Expand Down Expand Up @@ -536,6 +625,7 @@ def __init__(self, args: argparse.Namespace):
self.args = args
self.websocket = None
self.session_config = None
self._client_secret_token: Optional[str] = None
self.audio_data = []
self.is_synthesis_complete = False
self.wav_file = None # WAV file handle for streaming write
Expand Down Expand Up @@ -564,11 +654,14 @@ async def connect(self):
logger.info("Initializing HTTP session...")
session_data = await self._initialize_http_session()
self.session_config = session_data
self._client_secret_token = _extract_client_secret_token(session_data)
logger.info("HTTP session initialized successfully")

# Connect to WebSocket
logger.info("Connecting to WebSocket...")
await self._connect_websocket()
self.websocket = await _connect_websocket(
self.args, self._client_secret_token,
)
logger.info("WebSocket connected successfully")

# Initialize WebSocket session
Expand All @@ -593,7 +686,7 @@ async def connect(self):

async def _initialize_http_session(self) -> Dict[str, Any]:
"""Initialize session via HTTP POST request."""
headers = {"Content-Type": "application/json"}
headers = _mint_auth_headers(self.args)
uri = f"http://{self.args.server}/v1/realtime/synthesis_sessions"
if self.args.use_ssl:
uri = f"https://{self.args.server}/v1/realtime/synthesis_sessions"
Expand Down Expand Up @@ -631,33 +724,12 @@ async def _initialize_http_session(self) -> Dict[str, Any]:
logger.error("HTTP request failed: %s", e)
raise

if response.status_code != 200:
raise Exception(
f"Failed to initialize session. Status: {response.status_code}, "
f"Error: {response.text}"
)
_raise_for_session_status(response, provided_key=bool(getattr(self.args, "mint_api_key", None)))

session_data = response.json()
logger.info("Session initialized: %s", session_data)
return session_data

async def _connect_websocket(self):
"""Connect to WebSocket endpoint."""
ssl_context = None
ws_url = f"ws://{self.args.server}{self.args.endpoint}?{self.args.query_params}"
if self.args.use_ssl:
ws_url = f"wss://{self.args.server}{self.args.endpoint}?{self.args.query_params}"

ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
if self.args.ssl_root_cert:
ssl_context.load_verify_locations(self.args.ssl_root_cert)
if self.args.ssl_client_cert and self.args.ssl_client_key:
ssl_context.load_cert_chain(self.args.ssl_client_cert, self.args.ssl_client_key)
ssl_context.check_hostname = False

logger.info("Connecting to WebSocket: %s", ws_url)
self.websocket = await websockets.connect(ws_url, ssl=ssl_context)

async def _initialize_session(self):
"""Initialize the WebSocket session."""
try:
Expand Down
6 changes: 5 additions & 1 deletion scripts/tts/realtime_tts_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,10 @@
import websockets
from websockets.exceptions import WebSocketException

from riva.client.argparse_utils import add_connection_argparse_parameters
from riva.client.argparse_utils import (
add_connection_argparse_parameters,
add_realtime_config_argparse_parameters,
)
try:
from riva.client.argparse_utils import cli_main
except ImportError:
Expand Down Expand Up @@ -154,6 +157,7 @@ def parse_args() -> argparse.Namespace:

# Add connection parameters
parser = add_connection_argparse_parameters(parser)
parser = add_realtime_config_argparse_parameters(parser)

# Override default server for realtime TTS (WebSocket endpoint, not gRPC)
parser.set_defaults(server="localhost:9000")
Expand Down