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!
-
\ No newline at end of file
+
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