diff --git a/riva/client/argparse_utils.py b/riva/client/argparse_utils.py index d1cdf98..8d1f6e8 100644 --- a/riva/client/argparse_utils.py +++ b/riva/client/argparse_utils.py @@ -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 diff --git a/riva/client/realtime.py b/riva/client/realtime.py index bd41313..58f29af 100644 --- a/riva/client/realtime.py +++ b/riva/client/realtime.py @@ -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 ." + 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." +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.``. 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.""" @@ -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() @@ -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: @@ -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" @@ -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: @@ -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 @@ -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 @@ -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" @@ -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: diff --git a/scripts/tts/realtime_tts_client.py b/scripts/tts/realtime_tts_client.py index 3144450..25fbfda 100644 --- a/scripts/tts/realtime_tts_client.py +++ b/scripts/tts/realtime_tts_client.py @@ -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: @@ -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")