diff --git a/README.md b/README.md index 0882b27..a2cf0f7 100644 --- a/README.md +++ b/README.md @@ -44,6 +44,9 @@ The API supports multiple AI characters: ### Authentication +> Note: WebSocket calls may reject anonymous Firebase users. Prefer tokens from a logged-in web session. +> The client now also auto-runs `PUT /api/external/user` after auth (same as web app bootstrap). + ```python from sesame_ai import SesameAI, TokenManager @@ -60,7 +63,7 @@ print(f"User ID: {lookup_response.local_id}") # For easier token management, use TokenManager token_manager = TokenManager(client, token_file="token.json") -id_token = token_manager.get_valid_token() +id_token = token_manager.get_valid_token(allow_anonymous=False) ``` ### Voice Chat Example @@ -75,7 +78,7 @@ import numpy as np # Get authentication token using TokenManager api_client = SesameAI() token_manager = TokenManager(api_client, token_file="token.json") -id_token = token_manager.get_valid_token() +id_token = token_manager.get_valid_token(allow_anonymous=False) # Connect to WebSocket (choose character: "Miles" or "Maya") ws = SesameWebSocket(id_token=id_token, character="Maya") @@ -178,6 +181,10 @@ Command-line options: - `--output-device`: Output device index - `--list-devices`: List audio devices and exit - `--token-file`: Path to token storage file +- `--id-token`: Firebase ID token (or `SESAME_ID_TOKEN`) +- `--refresh-token`: Firebase refresh token (or `SESAME_REFRESH_TOKEN`) +- `--firebase-auth-json`: Firebase web auth JSON path/raw JSON (or `SESAME_FIREBASE_AUTH_JSON`) +- `--no-anonymous`: Disable anonymous fallback and require explicit token/auth JSON - `--debug`: Enable debug logging ## API Reference @@ -196,7 +203,9 @@ The main API client for authentication. Manages authentication tokens with automatic refresh and persistence. - `TokenManager(api_client=None, token_file=None)` - Create a token manager -- `get_valid_token(force_new=False)` - Get a valid token, refreshing if needed +- `get_valid_token(force_new=False, allow_anonymous=True)` - Get a valid token, refreshing if needed +- `set_tokens(id_token, refresh_token=None, ...)` - Inject tokens directly +- `load_firebase_auth_user(firebase_auth_user)` - Load tokens from Firebase web auth JSON - `clear_tokens()` - Clear stored tokens ### SesameWebSocket @@ -263,7 +272,8 @@ If you have trouble connecting: 1. Check your internet connection 2. Verify your authentication token is valid -3. Ensure the SesameAI service is available +3. If you see `Anonymous users not allowed`, pass `--id-token` (and optionally `--refresh-token`) or `--firebase-auth-json` +4. Ensure the SesameAI service is available ## Legal Disclaimer @@ -279,4 +289,4 @@ This project is licensed under the MIT License - see the [LICENSE](LICENSE) file If you find this project helpful, consider buying me a coffee! -Buy Me A Coffee \ No newline at end of file +Buy Me A Coffee diff --git a/examples/voice_chat.py b/examples/voice_chat.py index 64cde36..e1a0023 100644 --- a/examples/voice_chat.py +++ b/examples/voice_chat.py @@ -20,6 +20,7 @@ import sys import os +import json import time import threading import argparse @@ -27,21 +28,31 @@ import logging import numpy as np import pyaudio -from sesame_ai import SesameAI, SesameWebSocket, TokenManager, InvalidTokenError, NetworkError, APIError +from sesame_ai import ( + SesameAI, + SesameWebSocket, + TokenManager, + AuthenticationError, + InvalidTokenError, + NetworkError, + APIError, +) logger = logging.getLogger('sesame.examples.voice_chat') class VoiceChat: """Voice chat application using SesameAI""" - + # Available characters AVAILABLE_CHARACTERS = ["Miles", "Maya"] - - def __init__(self, character="Miles", input_device=None, output_device=None, - token_file=None): + + def __init__(self, character="Miles", input_device=None, output_device=None, + token_file=None, id_token=None, refresh_token=None, + firebase_auth_json=None, allow_anonymous=True, + client_name="Consumer-Web-App"): """ Initialize the voice chat application - + Args: character (str): Character to chat with ("Miles" or "Maya") input_device (int, optional): Input device index @@ -54,88 +65,138 @@ def __init__(self, character="Miles", input_device=None, output_device=None, if character not in self.AVAILABLE_CHARACTERS: print(f"Warning: '{character}' is not in the list of known characters. Using anyway.") print(f"Known characters: {', '.join(self.AVAILABLE_CHARACTERS)}") - + self.character = character self.input_device_index = input_device self.output_device_index = output_device self.token_file = token_file - + self.provided_id_token = id_token or os.environ.get("SESAME_ID_TOKEN") + self.provided_refresh_token = refresh_token or os.environ.get("SESAME_REFRESH_TOKEN") + self.firebase_auth_json = firebase_auth_json or os.environ.get("SESAME_FIREBASE_AUTH_JSON") + self.allow_anonymous = allow_anonymous + self.client_name = client_name or os.environ.get("SESAME_CLIENT_NAME", "Consumer-Web-App") + # Audio settings self.chunk_size = 1024 self.sample_format = pyaudio.paInt16 self.channels = 1 self.input_rate = 16000 self.output_rate = 24000 # Will be updated from server - + # Voice activity detection self.amplitude_threshold = 500 self.silence_counter = 0 self.silence_limit = 50 # Number of consecutive silent chunks before sending silence - + # PyAudio instance self.p = pyaudio.PyAudio() - + # Streams self.input_stream = None self.output_stream = None - + # SesameAI client self.api_client = SesameAI() - + # Initialize token manager with token_file (which may be None) self.token_manager = TokenManager(self.api_client, token_file=self.token_file) - + self.id_token = None self.ws = None - + # Thread control self.running = False self.threads = [] # Logging logger.debug(f"VoiceChat initialized with character: {character}") - + + def _load_firebase_auth_json(self): + """Load Firebase auth user JSON from a file path or raw JSON string.""" + source = self.firebase_auth_json + if not source: + return + + raw_json = source + if os.path.exists(source): + with open(source, "r") as f: + raw_json = f.read() + + try: + payload = json.loads(raw_json) + except json.JSONDecodeError as e: + raise AuthenticationError( + f"Invalid JSON in --firebase-auth-json (or SESAME_FIREBASE_AUTH_JSON): {e}" + ) + + self.token_manager.load_firebase_auth_user(payload, save=bool(self.token_file)) + def authenticate(self): """Authenticate with SesameAI and get a token""" logger.info("Authenticating with SesameAI...") try: - # If no token file is specified, force a new token - force_new = self.token_file is None - + if self.provided_id_token: + logger.info("Using ID token provided via CLI or environment variable") + self.token_manager.set_tokens( + id_token=self.provided_id_token, + refresh_token=self.provided_refresh_token, + save=bool(self.token_file) + ) + elif self.firebase_auth_json: + logger.info("Loading Firebase auth JSON from CLI or environment variable") + self._load_firebase_auth_json() + # Get a valid token using the token manager - self.id_token = self.token_manager.get_valid_token(force_new=force_new) + self.id_token = self.token_manager.get_valid_token( + force_new=False, + allow_anonymous=self.allow_anonymous + ) + + # Web app does this before starting a call; required for fresh anonymous users. + self.api_client.ensure_external_user( + id_token=self.id_token, + client_name=self.client_name, + ) + logger.info("Authentication successful!") return True + except AuthenticationError as e: + logger.error(f"Authentication failed: {e}") + logger.error( + "This endpoint now rejects anonymous users. Provide --id-token " + "(optionally with --refresh-token) or --firebase-auth-json." + ) + return False except InvalidTokenError: logger.error("Authentication failed: Token expired and couldn't be refreshed") return False except (NetworkError, APIError) as e: logger.error(f"Authentication failed: {e}") return False - + def list_audio_devices(self): """List available audio devices""" logger.info("Listing available audio devices") print("\nAvailable audio devices:") print("-" * 60) - + for i in range(self.p.get_device_count()): dev_info = self.p.get_device_info_by_index(i) name = dev_info.get('name', 'Unknown') inputs = dev_info.get('maxInputChannels', 0) outputs = dev_info.get('maxOutputChannels', 0) - + if inputs > 0: print(f"ID {i}: {name} (Input)") if outputs > 0: print(f"ID {i}: {name} (Output)") - + print("-" * 60) - + def select_devices(self): """Select input and output devices""" self.list_audio_devices() - + # If devices weren't specified in constructor, ask user if self.input_device_index is None: try: @@ -144,7 +205,7 @@ def select_devices(self): except ValueError: logger.warning("Invalid input. Using default device.") self.input_device_index = None - + if self.output_device_index is None: try: self.output_device_index = int(input("Select output device ID: ")) @@ -152,54 +213,58 @@ def select_devices(self): except ValueError: logger.warning("Invalid input. Using default device.") self.output_device_index = None - + def on_connect(self): """Callback when WebSocket connection is established""" logger.info(f"Connected to {self.character}!") # Update output rate from server self.output_rate = self.ws.server_sample_rate logger.debug(f"Server sample rate: {self.output_rate}") - + # Initialize audio streams after connection self.setup_audio_streams() - + # Start audio threads self.start_audio_threads() - + def on_disconnect(self): """Callback when WebSocket connection is disconnected""" logger.info(f"Disconnected from {self.character}") - + # Stop the application if it's still running if self.running: self.stop() - + def connect(self): """Connect to SesameAI WebSocket""" logger.info(f"Connecting to SesameAI as character '{self.character}'...") - + # Create WebSocket client self.ws = SesameWebSocket( id_token=self.id_token, - character=self.character + character=self.character, + client_name=self.client_name, ) - + # Set up callbacks self.ws.set_connect_callback(self.on_connect) self.ws.set_disconnect_callback(self.on_disconnect) - + # Connect to server if self.ws.connect(): logger.debug("WebSocket connection established") return True else: - logger.error("Failed to connect to SesameAI") + if getattr(self.ws, "last_error", None): + logger.error(f"Failed to connect to SesameAI: {self.ws.last_error}") + else: + logger.error("Failed to connect to SesameAI") return False - + def setup_audio_streams(self): """Set up audio input and output streams""" logger.debug("Setting up audio streams") - + # Input stream (microphone) self.input_stream = self.p.open( format=self.sample_format, @@ -209,7 +274,7 @@ def setup_audio_streams(self): frames_per_buffer=self.chunk_size, input_device_index=self.input_device_index ) - + # Output stream (speaker) self.output_stream = self.p.open( format=self.sample_format, @@ -218,26 +283,26 @@ def setup_audio_streams(self): output=True, output_device_index=self.output_device_index ) - + logger.debug("Audio streams initialized") - + def capture_microphone(self): """Capture audio from microphone and send to SesameAI""" logger.debug("Microphone capture started") - + while self.running: if not self.ws.is_connected(): time.sleep(0.1) continue - + try: # Read audio data from microphone data = self.input_stream.read(self.chunk_size, exception_on_overflow=False) - + # Check audio level for voice activity detection audio_samples = np.frombuffer(data, dtype=np.int16) rms_val = np.sqrt(np.mean(audio_samples.astype(np.float32) ** 2)) - + if rms_val > self.amplitude_threshold: # Voice detected self.silence_counter = 0 @@ -257,11 +322,11 @@ def capture_microphone(self): if self.running: logger.error(f"Error capturing microphone: {e}", exc_info=True) time.sleep(0.1) - + def play_audio(self): """Play audio received from SesameAI""" logger.debug("Audio playback started") - + while self.running: try: # Get audio chunk from WebSocket buffer with a short timeout @@ -272,7 +337,7 @@ def play_audio(self): except Exception as e: if self.running: logger.error(f"Error playing audio: {e}", exc_info=True) - + def start_audio_threads(self): """Start audio capture and playback threads""" # Microphone capture thread @@ -280,61 +345,61 @@ def start_audio_threads(self): mic_thread.daemon = True mic_thread.start() self.threads.append(mic_thread) - + # Audio playback thread playback_thread = threading.Thread(target=self.play_audio) playback_thread.daemon = True playback_thread.start() self.threads.append(playback_thread) - + logger.debug("Audio threads started") - + def start(self): """Start the voice chat""" # Authenticate if not self.authenticate(): return False - + # Select audio devices self.select_devices() - + # Set running flag self.running = True - + # Connect to WebSocket (will trigger on_connect callback) if not self.connect(): self.running = False return False - + logger.info(f"Voice chat with {self.character} started! Press Ctrl+C to exit.") return True - + def stop(self): """Stop the voice chat""" if not self.running: return - + self.running = False logger.info("Stopping voice chat...") - + # Disconnect from WebSocket if self.ws and self.ws.is_connected(): self.ws.disconnect() - + # Close audio streams if self.input_stream: self.input_stream.stop_stream() self.input_stream.close() - + if self.output_stream: self.output_stream.stop_stream() self.output_stream.close() - + # Terminate PyAudio self.p.terminate() - + logger.info("Voice chat stopped") - + def run(self): """Run the voice chat application""" try: @@ -360,37 +425,61 @@ def main(): # Set websocket-client logger to DEBUG level logging.getLogger('websocket').setLevel(logging.WARNING) - + parser = argparse.ArgumentParser(description="SesameAI Voice Chat Example") parser.add_argument("--character", default="Miles", choices=VoiceChat.AVAILABLE_CHARACTERS, - help=f"Character to chat with (default: Miles, options: {', '.join(VoiceChat.AVAILABLE_CHARACTERS)})") + help=f"Character to chat with (default: Miles, options: {', '.join(VoiceChat.AVAILABLE_CHARACTERS)})") parser.add_argument("--input-device", type=int, help="Input device index") parser.add_argument("--output-device", type=int, help="Output device index") parser.add_argument("--list-devices", action="store_true", help="List audio devices and exit") parser.add_argument("--token-file", help="Path to token storage file") + parser.add_argument("--id-token", help="Firebase ID token (or set SESAME_ID_TOKEN)") + parser.add_argument("--refresh-token", help="Firebase refresh token (or set SESAME_REFRESH_TOKEN)") + parser.add_argument( + "--firebase-auth-json", + help=( + "Path to Firebase auth user JSON or raw JSON string " + "(or set SESAME_FIREBASE_AUTH_JSON)" + ), + ) + parser.add_argument( + "--no-anonymous", + action="store_true", + help="Disable anonymous token fallback (require explicit token/auth JSON)", + ) + parser.add_argument( + "--client-name", + default="Consumer-Web-App", + help="WebSocket client_name query param (default: Consumer-Web-App)", + ) parser.add_argument("--debug", action="store_true", help="Enable debug logging") - + args = parser.parse_args() - + # Set debug level if requested if args.debug: logging.getLogger('sesame').setLevel(logging.DEBUG) - + # Create voice chat instance voice_chat = VoiceChat( character=args.character, input_device=args.input_device, output_device=args.output_device, - token_file=args.token_file + token_file=args.token_file, + id_token=args.id_token, + refresh_token=args.refresh_token, + firebase_auth_json=args.firebase_auth_json, + allow_anonymous=not args.no_anonymous, + client_name=args.client_name, ) - + # List devices and exit if requested if args.list_devices: voice_chat.list_audio_devices() return - + # Run the voice chat voice_chat.run() if __name__ == "__main__": - main() \ No newline at end of file + main() diff --git a/sesame_ai/api.py b/sesame_ai/api.py index dbd1bba..f1bd88f 100644 --- a/sesame_ai/api.py +++ b/sesame_ai/api.py @@ -8,32 +8,33 @@ class SesameAI: """ SesameAI API Client - Unofficial Python client for the SesameAI API - + Provides authentication and account management functionality for SesameAI services. """ - + def __init__(self, api_key=None): """ Initialize the SesameAI API client - + Args: - api_key (str, optional): Firebase API key. If not provided, + api_key (str, optional): Firebase API key. If not provided, will use the default key from config. """ self.api_key = api_key - + self.base_url = "https://sesameai.app" + def _make_auth_request(self, request_type, payload, is_form_data=False): """ Make a request to the Firebase Authentication API - + Args: request_type (str): Type of request ('signup', 'lookup', etc.) payload (dict): Request payload is_form_data (bool): Whether payload should be sent as form data - + Returns: dict: API response as JSON - + Raises: NetworkError: If a network error occurs APIError: If the API returns an error response @@ -42,7 +43,7 @@ def _make_auth_request(self, request_type, payload, is_form_data=False): headers = get_headers(request_type) params = get_params(request_type, self.api_key) url = get_endpoint_url(request_type) - + try: if is_form_data: response = requests.post( @@ -58,29 +59,29 @@ def _make_auth_request(self, request_type, payload, is_form_data=False): headers=headers, json=payload, ) - + # Check for HTTP errors response.raise_for_status() - + # Parse the response response_json = response.json() - + # Check for API errors if 'error' in response_json: self._handle_api_error(response_json['error']) - + return response_json - + except requests.exceptions.RequestException as e: raise NetworkError(f"Network error: {str(e)}") - + def _handle_api_error(self, error): """ Handle API error responses - + Args: error (dict): Error information from API - + Raises: InvalidTokenError: If a token is invalid APIError: For other API errors @@ -88,21 +89,21 @@ def _handle_api_error(self, error): error_code = error.get('code', 400) error_message = error.get('message', 'Unknown error') error_details = error.get('errors', []) - + # Handle specific error types if error_message in ('INVALID_ID_TOKEN', 'INVALID_REFRESH_TOKEN'): raise InvalidTokenError() - + # Generic API error raise APIError(error_code, error_message, error_details) - + def create_anonymous_account(self): """ Create an anonymous account - + Returns: SignupResponse: Object containing authentication tokens - + Raises: NetworkError: If a network error occurs APIError: If the API returns an error response @@ -112,17 +113,17 @@ def create_anonymous_account(self): } response_json = self._make_auth_request('signup', payload) return SignupResponse(response_json) - + def refresh_authentication_token(self, refresh_token): """ Refresh an ID token using a refresh token - + Args: refresh_token (str): Firebase refresh token - + Returns: RefreshTokenResponse: Object containing new tokens - + Raises: NetworkError: If a network error occurs APIError: If the API returns an error response @@ -132,20 +133,20 @@ def refresh_authentication_token(self, refresh_token): 'grant_type': 'refresh_token', 'refresh_token': refresh_token } - + response_json = self._make_auth_request('refresh', payload, is_form_data=True) return RefreshTokenResponse(response_json) - + def get_account_info(self, id_token): """ Get account information using an ID token - + Args: id_token (str): Firebase ID token - + Returns: LookupResponse: Object containing account information - + Raises: NetworkError: If a network error occurs APIError: If the API returns an error response @@ -154,7 +155,32 @@ def get_account_info(self, id_token): payload = { 'idToken': id_token } - + response_json = self._make_auth_request('lookup', payload) return LookupResponse(response_json) - \ No newline at end of file + + def ensure_external_user(self, id_token, client_name="Consumer-Web-App"): + """ + Ensure a Sesame backend user exists for a Firebase ID token. + + This mirrors the web app bootstrap call: + PUT /api/external/user + """ + url = f"{self.base_url}/api/external/user" + headers = { + "Authorization": f"Bearer {id_token}", + "Client-Name": client_name, + "Content-Type": "application/json", + } + payload = {} + + try: + response = requests.put(url, headers=headers, json=payload) + response.raise_for_status() + return response.json() + except requests.exceptions.HTTPError as e: + status_code = e.response.status_code if e.response is not None else 500 + error_text = e.response.text if e.response is not None else str(e) + raise APIError(status_code, f"Failed to ensure external user: {error_text}") + except requests.exceptions.RequestException as e: + raise NetworkError(f"Network error while ensuring external user: {str(e)}") diff --git a/sesame_ai/token_manager.py b/sesame_ai/token_manager.py index e067996..7c88224 100644 --- a/sesame_ai/token_manager.py +++ b/sesame_ai/token_manager.py @@ -5,24 +5,24 @@ import time import logging from .api import SesameAI -from .exceptions import InvalidTokenError, NetworkError, APIError +from .exceptions import AuthenticationError, InvalidTokenError, NetworkError, APIError logger = logging.getLogger('sesame.token_manager') class TokenManager: """ Manages authentication tokens for SesameAI API - + Handles: - Token storage and retrieval - Token validation - Automatic token refresh """ - + def __init__(self, api_client=None, token_file=None): """ Initialize the token manager - + Args: api_client (SesameAI, optional): API client instance. If None, creates a new one. token_file (str, optional): Path to token storage file. @@ -30,11 +30,11 @@ def __init__(self, api_client=None, token_file=None): self.api_client = api_client if api_client else SesameAI() self.token_file = token_file if token_file else None self.tokens = self._load_tokens() - + def _load_tokens(self): """ Load tokens from storage file - + Returns: dict: Token data or empty dict if file doesn't exist """ @@ -47,19 +47,19 @@ def _load_tokens(self): logger.warning(f"Failed to load tokens: {e}") return {} return {} - + def _save_tokens(self): """Save tokens to storage file""" try: # If no token file is specified, return early if self.token_file is None: return - + # Make sure the directory exists directory = os.path.dirname(self.token_file) if directory: # Only try to create directory if there is one os.makedirs(directory, exist_ok=True) - + # Write the tokens to the file with open(self.token_file, 'w') as f: logger.debug(f"Saving tokens to {self.token_file}") @@ -67,14 +67,88 @@ def _save_tokens(self): logger.debug(f"Tokens successfully saved to {self.token_file}") except Exception as e: logger.warning(f"Could not save tokens: {e}", exc_info=True) - + + def set_tokens(self, id_token, refresh_token=None, user_id=None, expires_in=None, save=True): + """ + Set tokens directly (for example from browser Firebase auth state). + + Args: + id_token (str): Firebase ID token + refresh_token (str, optional): Firebase refresh token + user_id (str, optional): Firebase user ID + expires_in (str|int, optional): Token TTL + save (bool): Persist to token_file when available + """ + if not id_token: + raise AuthenticationError("id_token is required") + + self.tokens = { + "id_token": id_token, + "refresh_token": refresh_token, + "user_id": user_id, + "expires_in": expires_in, + "timestamp": int(time.time()), + } + + if save: + self._save_tokens() + + def load_firebase_auth_user(self, firebase_auth_user, save=True): + """ + Load token data from Firebase Web Auth user JSON (localStorage payload). + + Expected shape includes: + - stsTokenManager.accessToken + - stsTokenManager.refreshToken + """ + if not isinstance(firebase_auth_user, dict): + raise AuthenticationError("firebase_auth_user must be a JSON object") + + sts = firebase_auth_user.get("stsTokenManager", {}) + id_token = ( + sts.get("accessToken") + or firebase_auth_user.get("idToken") + or firebase_auth_user.get("id_token") + ) + refresh_token = ( + sts.get("refreshToken") + or firebase_auth_user.get("refreshToken") + or firebase_auth_user.get("refresh_token") + ) + user_id = ( + firebase_auth_user.get("uid") + or firebase_auth_user.get("localId") + or firebase_auth_user.get("user_id") + ) + + expiration_time_ms = sts.get("expirationTime") + expires_in = firebase_auth_user.get("expiresIn") + if expiration_time_ms and not expires_in: + try: + expires_in = max(0, int((int(expiration_time_ms) / 1000) - time.time())) + except (TypeError, ValueError): + expires_in = None + + if not id_token: + raise AuthenticationError( + "No id token found in Firebase auth JSON. Expected stsTokenManager.accessToken." + ) + + self.set_tokens( + id_token=id_token, + refresh_token=refresh_token, + user_id=user_id, + expires_in=expires_in, + save=save, + ) + def _is_token_expired(self, id_token): """ Check if an ID token is expired - + Args: id_token (str): Firebase ID token - + Returns: bool: True if token is expired or invalid """ @@ -87,17 +161,18 @@ def _is_token_expired(self, id_token): except (NetworkError, APIError) as e: # If lookup fails, raise the error raise e - - def get_valid_token(self, force_new=False): + + def get_valid_token(self, force_new=False, allow_anonymous=True): """ Get a valid ID token, refreshing if necessary - + Args: force_new (bool): If True, creates a new account regardless of existing tokens - + allow_anonymous (bool): If False, do not create a new anonymous account + Returns: str: Valid ID token - + Raises: InvalidTokenError: If token refresh fails NetworkError: If a network error occurs @@ -105,13 +180,17 @@ def get_valid_token(self, force_new=False): """ # If force_new is True, create a new account if force_new: + if not allow_anonymous: + raise AuthenticationError( + "Cannot force a new token when anonymous auth is disabled." + ) logger.debug("Forcing creation of new account") return self._create_new_account() - + # Check if we have an existing ID token id_token = self.tokens.get('id_token') refresh_token = self.tokens.get('refresh_token') - + if id_token: # Check if the token is still valid try: @@ -123,13 +202,13 @@ def get_valid_token(self, force_new=False): logger.warning(f"Error checking token validity: {e}") # If we can't check, assume it's still valid return id_token - + # Token is expired, try to refresh if refresh_token: try: logger.info("Refreshing expired token") refresh_response = self.api_client.refresh_authentication_token(refresh_token) - + # Update tokens self.tokens = { 'id_token': refresh_response.id_token, @@ -139,34 +218,38 @@ def get_valid_token(self, force_new=False): 'timestamp': int(time.time()) } self._save_tokens() - + logger.info("Token refreshed successfully") return refresh_response.id_token except (InvalidTokenError, NetworkError, APIError) as e: logger.error(f"Token refresh failed: {e}") - raise InvalidTokenError("Token refresh failed") + raise InvalidTokenError() else: logger.warning("Token expired and no refresh token available") - raise InvalidTokenError("Token expired and no refresh token available") + raise InvalidTokenError() else: # No existing token, create a new account + if not allow_anonymous: + raise AuthenticationError( + "No saved token found and anonymous auth is disabled." + ) logger.debug("No existing token, creating new account") return self._create_new_account() def _create_new_account(self): """ Create a new anonymous account - + Returns: str: New ID token - + Raises: NetworkError: If a network error occurs APIError: If the API returns an error """ logger.debug("Creating new anonymous account") signup_response = self.api_client.create_anonymous_account() - + # Save the new tokens self.tokens = { 'id_token': signup_response.id_token, @@ -176,7 +259,7 @@ def _create_new_account(self): 'timestamp': int(time.time()) } self._save_tokens() - + logger.debug("New account created successfully") return signup_response.id_token @@ -184,4 +267,4 @@ def clear_tokens(self): """Clear stored tokens""" logger.info("Clearing stored tokens") self.tokens = {} - self._save_tokens() \ No newline at end of file + self._save_tokens() diff --git a/sesame_ai/websocket.py b/sesame_ai/websocket.py index a2c8dd6..2ad98f0 100644 --- a/sesame_ai/websocket.py +++ b/sesame_ai/websocket.py @@ -9,6 +9,7 @@ import queue import time import logging +import os import websocket as websocket_module logger = logging.getLogger('sesame.websocket') @@ -17,95 +18,112 @@ class SesameWebSocket: """ WebSocket client for real-time communication with SesameAI """ - - def __init__(self, id_token, character="Miles", client_name="RP-Web"): + + def __init__( + self, + id_token, + character="Miles", + client_name="Consumer-Web-App", + origin="https://app.sesame.com", + preset=None, + ): """ Initialize the WebSocket client - + Args: id_token (str): Firebase ID token for authentication character (str, optional): Character to interact with. Defaults to "Miles". - client_name (str, optional): Client identifier. Defaults to "RP-Web". + client_name (str, optional): Client identifier. Defaults to "Consumer-Web-App". + origin (str, optional): Origin header sent during WebSocket handshake. + preset (str, optional): Optional preset name for call settings. """ self.id_token = id_token self.character = character self.client_name = client_name - + self.origin = origin + self.preset = preset + # WebSocket connection self.ws = None self.session_id = None self.call_id = None - + # Audio settings self.client_sample_rate = 16000 self.server_sample_rate = 24000 # Default, will be updated from server self.audio_codec = "none" - + # Connection state self.reconnect = False self.is_private = False self.user_agent = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/133.0.0.0 Safari/537.36" - + self.timezone = os.environ.get("TZ", "America/Chicago") + # Audio buffer for received audio self.audio_buffer = queue.Queue(maxsize=1000) - + # Message tracking self.last_sent_message_type = None self.received_since_last_sent = False self.first_audio_received = False - + # Event for tracking connection state self.connected_event = threading.Event() - + self.connection_complete_event = threading.Event() + self.last_error = None + # Callbacks self.on_connect_callback = None self.on_disconnect_callback = None - + def connect(self, blocking=True): """ Connect to the SesameAI WebSocket server - + Args: blocking (bool, optional): If True, blocks until connected. Defaults to True. - + Returns: bool: True if connection was successful """ # Reset connection state self.connected_event.clear() - + self.connection_complete_event.clear() + self.last_error = None + # Start connection in a separate thread connection_thread = threading.Thread(target=self._connect_websocket) connection_thread.daemon = True connection_thread.start() - + if blocking: - # Wait for connection to be established - return self.connected_event.wait(timeout=10) - + # Wait until connected or failed + self.connection_complete_event.wait(timeout=10) + return self.connected_event.is_set() + return True - + def _connect_websocket(self): """Internal method to establish WebSocket connection""" headers = { - 'Origin': 'https://www.sesame.com', + 'Origin': self.origin, 'User-Agent': self.user_agent, } params = { 'id_token': self.id_token, 'client_name': self.client_name, - 'usercontext': json.dumps({"timezone": "America/Chicago"}), + 'usercontext': json.dumps({"timezone": self.timezone}), 'character': self.character, } # Construct the WebSocket URL with query parameters base_url = 'wss://sesameai.app/agent-service-0/v1/connect' - + # Convert params to URL query string query_string = '&'.join([f"{key}={urllib.parse.quote(value)}" for key, value in params.items()]) ws_url = f"{base_url}?{query_string}" - + # Create WebSocket connection self.ws = websocket_module.WebSocketApp( ws_url, @@ -118,24 +136,24 @@ def _connect_websocket(self): # Run the WebSocket self.ws.run_forever( - sslopt={"cert_reqs": ssl.CERT_NONE}, + sslopt={"cert_reqs": ssl.CERT_NONE}, skip_utf8_validation=True, suppress_origin=False ) - + def _on_open(self, ws): """Callback when WebSocket connection is opened""" logger.debug("WebSocket connection opened") - + def _on_message(self, ws, message): """Callback when a message is received from the WebSocket""" try: # Parse the message as JSON data = json.loads(message) - + # Handle different message types message_type = data.get('type') - + if message_type == 'initialize': self._handle_initialize(data) elif message_type == 'call_connect_response': @@ -146,28 +164,37 @@ def _on_message(self, ws, message): self._handle_audio(data) elif message_type == 'call_disconnect_response': self._handle_call_disconnect_response(data) + elif message_type == 'error': + self._handle_server_error(data) else: logger.debug(f"Received message type: {message_type}") - + except json.JSONDecodeError: logger.warning(f"Received non-JSON message: {message}") except Exception as e: logger.error(f"Error handling message: {e}", exc_info=True) - + def _on_error(self, ws, error): - """Callback when a WebSocket error occurs""" - logger.error(f"WebSocket error: {error}") - self.connected_event.clear() - + """Callback when a WebSocket error occurs""" + logger.error(f"WebSocket error: {error}") + self.last_error = str(error) + if "Anonymous users not allowed" in str(error): + logger.error( + "Sesame rejected anonymous auth. Use a non-anonymous Firebase token." + ) + self.connected_event.clear() + self.connection_complete_event.set() + def _on_close(self, ws, close_status_code, close_msg): """Callback when the WebSocket connection is closed""" logger.debug(f"WebSocket closed: {close_status_code} - {close_msg}") self.connected_event.clear() - + self.connection_complete_event.set() + # Call the disconnect callback if set if self.on_disconnect_callback: self.on_disconnect_callback() - + # Message handlers def _handle_initialize(self, data): """Handle initialize message from server""" @@ -177,7 +204,7 @@ def _handle_initialize(self, data): # Send location and call_connect self._send_client_location_state() self._send_call_connect() - + def _handle_call_connect_response(self, data): """Handle call_connect_response message from server""" self.session_id = data.get('session_id') @@ -187,19 +214,30 @@ def _handle_call_connect_response(self, data): self.audio_codec = content.get('audio_codec', 'none') logger.debug(f"Connected: Session ID: {self.session_id}, Call ID: {self.call_id}") - + # Signal that we're connected self.connected_event.set() - + self.connection_complete_event.set() + # Call the connect callback if set if self.on_connect_callback: self.on_connect_callback() - - + + def _handle_server_error(self, data): + """Handle server error message""" + content = data.get('content', {}) if isinstance(data, dict) else {} + code = content.get('code') + message = content.get('message') + detail = content.get('detail') + self.last_error = f"Server error {code}: {message} ({detail})" + logger.error(f"Server error: code={code}, message={message}, detail={detail}") + self.connection_complete_event.set() + + def _handle_ping_response(self, data): """Handle ping_response message from server""" pass - + def _handle_audio(self, data): """Handle audio message from server""" audio_data = data.get('content', {}).get('audio_data', '') @@ -217,7 +255,7 @@ def _handle_audio(self, data): self.audio_buffer.put_nowait(audio_bytes) except queue.Empty: pass - + if not self.first_audio_received: self.first_audio_received = True logger.debug("First audio received, sending initialization chunks") @@ -227,16 +265,16 @@ def _handle_audio(self, data): self._send_audio(chunk_of_As) except Exception as e: logger.error(f"Error processing audio: {e}", exc_info=True) - + def _handle_call_disconnect_response(self, data): """Handle call_disconnect_response message from server""" logger.debug("Call disconnected") self.call_id = None - + # Call the disconnect callback if set if self.on_disconnect_callback: self.on_disconnect_callback() - + # Methods to send messages def _send_ping(self): """Send ping message to server""" @@ -252,7 +290,7 @@ def _send_ping(self): } self._send_data(message) - + def _send_client_location_state(self): """Send client_location_state message to server""" if not self.session_id: @@ -266,15 +304,15 @@ def _send_client_location_state(self): "latitude": 0, "longitude": 0, "address": "", - "timezone": "America/Chicago" + "timezone": self.timezone } } self._send_data(message) - + def _send_audio(self, data): """ Send audio data to server - + Args: data (str): Base64-encoded audio data """ @@ -291,66 +329,72 @@ def _send_audio(self, data): } self._send_data(message) - + def _send_call_connect(self): """Send call_connect message to server""" if not self.session_id: return - + + content = { + "sample_rate": self.client_sample_rate, + "audio_codec": "none", + "reconnect": self.reconnect, + "is_private": self.is_private, + "client_name": self.client_name, + "client_metadata": { + "language": "en-US", + "user_agent": self.user_agent, + "mobile_browser": False, + "media_devices": self._get_media_devices() + } + } + + # Match current web app behavior: + # settings = {character: } unless a specific preset is provided. + if self.preset: + content["settings"] = {"preset": self.preset} + else: + content["settings"] = {"character": self.character} + message = { "type": "call_connect", "session_id": self.session_id, "call_id": None, "request_id": self._generate_request_id(), - "content": { - "sample_rate": self.client_sample_rate, - "audio_codec": "none", - "reconnect": self.reconnect, - "is_private": self.is_private, - "client_name": self.client_name, - "settings": { - "preset": f"{self.character}" - }, - "client_metadata": { - "language": "en-US", - "user_agent": self.user_agent, - "mobile_browser": False, - "media_devices": self._get_media_devices() - } - } + "content": content } - + self._send_data(message) - + def send_audio_data(self, raw_audio_bytes): """ Send raw audio data to the AI - + Args: raw_audio_bytes (bytes): Raw audio data (16-bit PCM) - + Returns: bool: True if audio was sent successfully """ if not self.session_id or not self.call_id: return False - + # Encode the raw audio data in base64 encoded_data = base64.b64encode(raw_audio_bytes).decode('utf-8') self._send_audio(encoded_data) return True - + def disconnect(self): """ Disconnect from the server - + Returns: bool: True if disconnect message was sent successfully """ if not self.session_id or not self.call_id: logger.warning("Cannot disconnect: Not connected") return False - + message = { "type": "call_disconnect", "session_id": self.session_id, @@ -360,11 +404,11 @@ def disconnect(self): "reason": "user_request" } } - + logger.debug("Sending disconnect request") self._send_data(message) return True - + def _send_message(self, message): """Send a raw message to the WebSocket""" if self.ws and self.ws.sock and self.ws.sock.connected: @@ -374,7 +418,7 @@ def _send_message(self, message): else: logger.warning("WebSocket is not connected") return False - + def _send_data(self, message): """Send data with proper ping handling""" try: @@ -382,24 +426,24 @@ def _send_data(self, message): # Send pings for non-control messages after connection is established if self.call_id is not None and data_type not in ["ping", "call_connect", "call_disconnect"]: - if (self.last_sent_message_type is None - or self.received_since_last_sent - or (data_type != self.last_sent_message_type)): + if (self.last_sent_message_type is None + or self.received_since_last_sent + or (data_type != self.last_sent_message_type)): self._send_ping() - + self.last_sent_message_type = data_type self.received_since_last_sent = False return self._send_message(message) - + except Exception as e: logger.error(f"Error sending data: {e}", exc_info=True) return False - + def _generate_request_id(self): """Generate a unique request ID""" return str(uuid.uuid4()) - + def _get_media_devices(self): """Get a list of media devices for the client metadata""" # Simplified version - in a real implementation, this would detect actual devices @@ -417,14 +461,14 @@ def _get_media_devices(self): "groupId": "default" } ] - + def get_next_audio_chunk(self, timeout=None): """ Get the next audio chunk from the buffer - + Args: timeout (float, optional): Timeout in seconds. None means block indefinitely. - + Returns: bytes: Audio data, or None if timeout occurred """ @@ -432,30 +476,30 @@ def get_next_audio_chunk(self, timeout=None): return self.audio_buffer.get(timeout=timeout) except queue.Empty: return None - + def set_connect_callback(self, callback): """ Set callback for connection established events - + Args: callback (callable): Function with no arguments """ self.on_connect_callback = callback - + def set_disconnect_callback(self, callback): """ Set callback for disconnection events - + Args: callback (callable): Function with no arguments """ self.on_disconnect_callback = callback - + def is_connected(self): """ Check if the WebSocket is connected - + Returns: bool: True if connected """ - return self.session_id is not None and self.call_id is not None \ No newline at end of file + return self.session_id is not None and self.call_id is not None