diff --git a/src-tauri/src/cartesia/client.rs b/src-tauri/src/cartesia/client.rs index 29ce152..f70b016 100644 --- a/src-tauri/src/cartesia/client.rs +++ b/src-tauri/src/cartesia/client.rs @@ -49,4 +49,30 @@ impl CartesiaClient { stream } + + pub async fn open_tts_connection(&self) -> WebSocketStream> { + let mut request = Url::parse("wss://api.cartesia.ai/tts/websocket") + .expect("failed to parse TTS connection URL") + .as_str() + .into_client_request() + .expect("failed to instantiate TTS WebSocket request"); + + let headers = request.headers_mut(); + let api_key = self + .secrets_manager + .get_secret(SecretName::CartesiaApiKey) + .expect("failed to retrieve API key"); + + headers.insert( + "X-API-Key", + HeaderValue::from_str(api_key.as_str()).expect("could not convert key to header value"), + ); + headers.insert("Cartesia-Version", "2025-04-16".parse().unwrap()); + + let (stream, _) = connect_async(request) + .await + .expect("failed to open TTS websocket connection"); + + stream + } } diff --git a/src-tauri/src/cartesia/mod.rs b/src-tauri/src/cartesia/mod.rs index b71c998..0a47b60 100644 --- a/src-tauri/src/cartesia/mod.rs +++ b/src-tauri/src/cartesia/mod.rs @@ -1,3 +1,4 @@ pub mod client; pub mod commands; pub mod stt; +pub mod tts; diff --git a/src-tauri/src/cartesia/tts.rs b/src-tauri/src/cartesia/tts.rs new file mode 100644 index 0000000..f7a223d --- /dev/null +++ b/src-tauri/src/cartesia/tts.rs @@ -0,0 +1,232 @@ +use std::{collections::HashMap, sync::Arc}; + +use futures_util::{ + stream::{SplitSink, SplitStream}, + SinkExt, StreamExt, +}; +use serde::{Deserialize, Serialize}; +use serde_json::{from_str, json}; +use tokio::{ + net::TcpStream, + sync::{ + mpsc::{unbounded_channel, UnboundedReceiver, UnboundedSender}, + Mutex, + }, +}; +use tokio_tungstenite::{MaybeTlsStream, WebSocketStream}; +use ts_rs::TS; +use tungstenite::Message; + +use crate::cartesia::client::CartesiaClient; + +#[derive(Serialize, Deserialize, TS)] +pub struct TtsTimestamp { + words: Vec, + start: Vec, + end: Vec, +} + +#[derive(Serialize, Deserialize, TS)] +pub struct TtsPhonemeTimestamp { + phonemes: Vec, + start: Vec, + end: Vec, +} + +#[derive(Serialize, Deserialize, TS)] +#[serde(tag = "type", rename_all = "snake_case")] +#[ts(export)] +pub enum TtsMessage { + Chunk { + data: String, + done: bool, + status_code: u16, + step_time: f32, + context_id: Option, + }, + FlushDone { + done: bool, + flush_done: bool, + flush_id: u16, + status_code: u16, + context_id: Option, + }, + Done { + done: bool, + status_code: u16, + context_id: Option, + }, + Timestamps { + done: bool, + status_code: u16, + context_id: Option, + word_timestamps: Option>, + }, + Error { + done: bool, + error: String, + status_code: u16, + context_id: Option, + }, + PhonemeTimestamps { + done: bool, + status_code: u16, + context_id: Option, + phoneme_timestamps: Option>, + }, +} + +pub struct TtsContext { + pub id: String, + reader: Mutex>, + writer: Arc>>, +} +impl TtsContext { + async fn send(&self, content: String, is_final: bool) { + let id = self.id.clone(); + let tx = self.writer.lock().await; + + tx.send(TtsInputMessage { + context_id: id.clone(), + content, + done: is_final, + }) + .expect(format!("failed to send content to context {id}").as_str()) + } + + async fn recv(&self) -> Option { + self.reader.lock().await.recv().await + } +} + +#[derive(Serialize, Deserialize, TS)] +#[ts(export)] +pub struct TtsInputMessage { + context_id: String, + content: String, + done: bool, +} + +pub struct TtsManager { + client: Arc, + contexts: Arc>>>>>, + input: Option>>>, +} +impl TtsManager { + pub fn new(client: Arc) -> Self { + Self { + client, + contexts: Arc::new(Mutex::new(HashMap::new().into())), + input: None, + } + } + + async fn connect(&mut self) { + match &self.input { + Some(_) => (), + None => { + let stream = self.client.open_tts_connection().await; + let (tx_in, rx_in) = unbounded_channel::(); + let (tx_ws, rx_ws) = stream.split(); + + tokio::spawn(handle_incoming(self.contexts.clone(), rx_ws)); + tokio::spawn(handle_outgoing(tx_ws, rx_in)); + + self.input = Some(Arc::new(Mutex::new(tx_in))) + } + } + } + + async fn new_context(&self, id: String) -> Result { + match &self.input { + Some(input) => { + let (tx_out, rx_out) = unbounded_channel::(); + + { + let mut map = self.contexts.lock().await; + map.insert(id.clone(), Arc::new(Mutex::new(tx_out))); + } + + Ok(TtsContext { + id, + reader: Mutex::new(rx_out), + writer: input.clone(), + }) + } + None => { + todo!() + } + } + } +} + +async fn handle_incoming( + contexts: Arc>>>>>, + mut reader: SplitStream>>, +) { + while let Some(event) = reader.next().await { + match event { + Ok(msg) => { + let text = msg.into_text().expect("failed to decode websocket message"); + + println!("Got message: {}", text); + + match from_str::(&text).expect("failed to decode TTS message") { + TtsMessage::Chunk { + data, + done, + status_code, + step_time, + context_id: Some(id), + } => match contexts.lock().await.get(&id) { + Some(ctx) => { + ctx.lock() + .await + .send(TtsMessage::Chunk { + data, + done, + status_code, + step_time, + context_id: Some(id), + }) + .expect("failed to forward TTS chunk to contextual writer"); + } + None => { + eprintln!("no matching TTS sender context") + } + }, + _ => println!("got not chunk"), + } + } + Err(err) => { + eprintln!("got error from websocket: {}", err) + } + } + } +} + +async fn handle_outgoing( + mut writer: SplitSink>, Message>, + mut reader: UnboundedReceiver, +) { + while let Some(msg) = reader.recv().await { + let payload = json!({ + "model_id": "sonic-2", + "transcript": msg.content, + "voice": "TODO", + "output_format": { + "container": "raw", + "encoding": "pcm_f32le", + "sample_rate": 44100 + }, + "continue": msg.done, + "max_buffer_delay": 250, + "context_id": msg.context_id, + }); + + writer + .send(Message::Text(payload.to_string().into())) + .await + .expect("failed to send TTS generation request"); + } +} diff --git a/src-tauri/src/state.rs b/src-tauri/src/state.rs index e842296..15e9981 100644 --- a/src-tauri/src/state.rs +++ b/src-tauri/src/state.rs @@ -5,7 +5,7 @@ use tauri_plugin_http::reqwest::Client; use tauri_plugin_store::Store; use crate::{ - cartesia::{client::CartesiaClient, stt::SttManager}, + cartesia::{client::CartesiaClient, stt::SttManager, tts::TtsManager}, devices::{input::InputDeviceManager, output::OutputDeviceManager, types::AudioDeviceError}, letta::LettaManager, secrets::SecretsManager, @@ -14,6 +14,7 @@ use crate::{ pub struct AppState { pub cartesia_client: Arc, pub stt_manager: Arc, + pub tts_manager: Arc, pub letta_manager: Arc, pub secrets_manager: Arc, pub input_device_manager: Arc, @@ -38,6 +39,7 @@ impl AppState { cartesia_client.clone(), input_device_manager.clone(), )); + let tts_manager = Arc::new(TtsManager::new(cartesia_client.clone())); Ok(AppState { input_device_manager, @@ -46,6 +48,7 @@ impl AppState { letta_manager, cartesia_client, stt_manager, + tts_manager, }) } }