From 99d46164ab80d3894b4bc1fb2fbb4efa0387e900 Mon Sep 17 00:00:00 2001 From: Graham Barber Date: Tue, 23 Sep 2025 11:59:21 -0700 Subject: [PATCH] wip: stub conversation manager --- src-tauri/Cargo.lock | 1 + src-tauri/Cargo.toml | 1 + src-tauri/src/cartesia/tts.rs | 21 ++-- src-tauri/src/conversation/mod.rs | 178 ++++++++++++++++++++++++++++ src-tauri/src/conversation/types.rs | 52 ++++++++ src-tauri/src/letta/types.rs | 12 +- src-tauri/src/lib.rs | 1 + 7 files changed, 248 insertions(+), 18 deletions(-) create mode 100644 src-tauri/src/conversation/mod.rs create mode 100644 src-tauri/src/conversation/types.rs diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 3c65ac5..5fb04e8 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -2421,6 +2421,7 @@ dependencies = [ name = "miwiwi" version = "0.1.0" dependencies = [ + "base64 0.22.1", "cpal", "dasp", "futures-util", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 709ea73..2026a73 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -41,6 +41,7 @@ reqwest-eventsource = "0.6.0" ts-rs = "11.0.1" serde_string_enum = "0.2.1" unicase = "2.8.1" +base64 = "0.22.1" [target.'cfg(not(any(target_os = "android", target_os = "ios")))'.dependencies] tauri-plugin-positioner = "2" diff --git a/src-tauri/src/cartesia/tts.rs b/src-tauri/src/cartesia/tts.rs index f7a223d..213dcc3 100644 --- a/src-tauri/src/cartesia/tts.rs +++ b/src-tauri/src/cartesia/tts.rs @@ -76,13 +76,14 @@ pub enum TtsMessage { }, } +#[derive(Clone)] pub struct TtsContext { pub id: String, - reader: Mutex>, + pub reader: Arc>>, writer: Arc>>, } impl TtsContext { - async fn send(&self, content: String, is_final: bool) { + pub async fn send(&self, content: String, is_final: bool) { let id = self.id.clone(); let tx = self.writer.lock().await; @@ -93,18 +94,14 @@ impl TtsContext { }) .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 context_id: String, + pub content: String, + pub done: bool, } pub struct TtsManager { @@ -121,7 +118,7 @@ impl TtsManager { } } - async fn connect(&mut self) { + pub async fn connect(&mut self) { match &self.input { Some(_) => (), None => { @@ -137,7 +134,7 @@ impl TtsManager { } } - async fn new_context(&self, id: String) -> Result { + pub async fn new_context(&self, id: String) -> Result { match &self.input { Some(input) => { let (tx_out, rx_out) = unbounded_channel::(); @@ -149,7 +146,7 @@ impl TtsManager { Ok(TtsContext { id, - reader: Mutex::new(rx_out), + reader: Arc::new(Mutex::new(rx_out)), writer: input.clone(), }) } diff --git a/src-tauri/src/conversation/mod.rs b/src-tauri/src/conversation/mod.rs new file mode 100644 index 0000000..784a315 --- /dev/null +++ b/src-tauri/src/conversation/mod.rs @@ -0,0 +1,178 @@ +use std::{io::Cursor, sync::Arc}; + +use base64::prelude::{Engine as _, BASE64_STANDARD}; +use tauri::async_runtime::spawn; +use tokio::sync::{ + mpsc::{unbounded_channel, UnboundedReceiver}, + Mutex, RwLock, +}; + +use crate::{ + cartesia::tts::{TtsContext, TtsInputMessage, TtsManager, TtsMessage}, + conversation::types::{Turn, TurnMessage}, + letta::{ + types::{LettaCompletionMessage, LettaMessageContent}, + LettaManager, + }, +}; + +mod types; + +pub struct ConversationManager { + turn: Option>>, + current_msg_index: RwLock, + letta_manager: Arc, + tts_manager: Arc, +} +impl ConversationManager { + pub fn new(letta_manager: Arc, tts_manager: Arc) -> Self { + Self { + turn: None, + current_msg_index: RwLock::new(0), + letta_manager, + tts_manager, + } + } + + pub fn is_idle(&self) -> bool { + self.turn.is_none() + } + + /// Start a new conversation turn + pub async fn start_turn(&mut self, prompt: String) { + if !self.is_idle() { + return; + }; + + let turn = Arc::new(RwLock::new(Turn::new())); + + { + let mut idx = self.current_msg_index.write().await; + + *idx = 0; + self.turn = Some(turn.clone()); + } + + spawn(handle_letta_messages( + prompt, + self.letta_manager.clone(), + self.tts_manager.clone(), + turn.clone(), + )); + } +} + +async fn handle_letta_messages( + prompt: String, + letta: Arc, + tts: Arc, + turn: Arc>, +) { + match letta.start_completion(prompt).await { + Ok(mut iter) => { + while let Some(msg) = iter.recv().await { + let turn = turn.read().await; + + match turn.latest() { + Some(latest) => match latest { + TurnMessage::TextMessage { + id, + reader: _, + writer, + } if id == msg.id => { + // Add current message to chunks + } + TurnMessage::AudioMessage { + id, + reader: _, + writer, + context, + cursor: _, + timestamps: _, + } => { + // Add current message to chunks + } + }, + None => { + // Create a new message + // Append message to chunks + } + } + } + } + Err(err) => eprintln!("failed to start completion: {}", err), + } +} + +async fn create_turn_message_from_letta( + src: LettaCompletionMessage, + tts: Arc, +) -> Option { + let (writer, reader) = unbounded_channel::(); + + writer.send(src.clone()); + + match src { + LettaCompletionMessage::ApprovalRequestMessage { id, .. } + | LettaCompletionMessage::ApprovalResponseMessage { id, .. } + | LettaCompletionMessage::HiddenReasoningMessage { id, .. } + | LettaCompletionMessage::SystemMessage { id, .. } + | LettaCompletionMessage::ToolCallMessage { id, .. } + | LettaCompletionMessage::ToolReturnMessage { id, .. } => { + Some(TurnMessage::TextMessage { id, reader, writer }) + } + LettaCompletionMessage::AssistantMessage { + id, + content: blocks, + .. + } => { + let cursor = Cursor::new(Vec::new()); + let mut content = "".to_owned(); + + for b in blocks { + match b { + LettaMessageContent::Text { text } => content.push_str(&text), + _ => (), + } + } + + let context = tts + .new_context(id.clone()) + .await + .expect("failed to create new TTS context"); + + context.send(content, false).await; + + // Spawn task for handling audio generation + spawn((async |reader: Arc< + Mutex>, + >| { + // Listen to context and append to cursor + while let Some(msg) = reader.lock().await.recv().await { + match msg { + TtsMessage::Chunk { data, .. } => { + // Decode chunk and write to cursor + let out = BASE64_STANDARD.decode_vec(data, cursor.get_mut()); + } + TtsMessage::Timestamps { + word_timestamps, .. + } => { + // Append timestamps + } + _ => (), + } + } + })(context.reader.clone())); + + Some(TurnMessage::AudioMessage { + id: id.clone(), + reader, + writer, + context, + cursor, + timestamps: Vec::new(), + }) + } + _ => None, + } +} diff --git a/src-tauri/src/conversation/types.rs b/src-tauri/src/conversation/types.rs new file mode 100644 index 0000000..96875e1 --- /dev/null +++ b/src-tauri/src/conversation/types.rs @@ -0,0 +1,52 @@ +use std::io::Cursor; + +use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender}; + +use crate::{ + cartesia::tts::{TtsContext, TtsTimestamp}, + letta::types::LettaCompletionMessage, +}; + +pub enum TurnMessage { + /// Represents a message to be displayed without audio + TextMessage { + id: String, + reader: UnboundedReceiver, + writer: UnboundedSender, + }, + + /// Represents a message with associated audio data + AudioMessage { + id: String, + reader: UnboundedReceiver, + writer: UnboundedSender, + context: TtsContext, + cursor: Cursor>, + timestamps: Vec, + }, +} + +pub struct Turn { + messages: Vec, + done: bool, +} +impl Turn { + pub fn new() -> Self { + Self { + messages: Vec::new(), + done: false, + } + } + + pub fn complete(&mut self) { + self.done = true; + } + + pub fn latest(&self) -> Option<&TurnMessage> { + self.messages.last() + } + + pub fn add_message(&mut self, msg: TurnMessage) { + self.messages.push(msg); + } +} diff --git a/src-tauri/src/letta/types.rs b/src-tauri/src/letta/types.rs index 3dbbb36..f7af779 100644 --- a/src-tauri/src/letta/types.rs +++ b/src-tauri/src/letta/types.rs @@ -56,14 +56,14 @@ pub struct LettaAgentInfo { name: String, } -#[derive(Serialize, Deserialize, TS, Debug)] +#[derive(Serialize, Deserialize, TS, Debug, Clone)] #[serde(tag = "type", rename_all = "lowercase")] pub enum LettaMessageContent { Text { text: String }, Image { source: String }, } -#[derive(TS, Debug, SerializeLabeledStringEnum, DeserializeLabeledStringEnum)] +#[derive(TS, Debug, SerializeLabeledStringEnum, DeserializeLabeledStringEnum, Clone, Copy)] pub enum LettaReasoningSource { #[string = "reasoner_model"] ReasonerModel, @@ -72,7 +72,7 @@ pub enum LettaReasoningSource { NonReasonerModel, } -#[derive(SerializeLabeledStringEnum, DeserializeLabeledStringEnum, TS, Debug)] +#[derive(SerializeLabeledStringEnum, DeserializeLabeledStringEnum, TS, Debug, Clone, Copy)] pub enum LettaHiddenReasoningState { #[string = "redacted"] Redacted, @@ -81,7 +81,7 @@ pub enum LettaHiddenReasoningState { Omitted, } -#[derive(Serialize, Deserialize, TS, Debug)] +#[derive(Serialize, Deserialize, TS, Debug, Clone)] #[serde(untagged)] pub enum LettaToolCall { Call { @@ -96,7 +96,7 @@ pub enum LettaToolCall { }, } -#[derive(SerializeLabeledStringEnum, DeserializeLabeledStringEnum, TS, Debug)] +#[derive(SerializeLabeledStringEnum, DeserializeLabeledStringEnum, TS, Debug, Clone, Copy)] pub enum LettaToolReturnStatus { #[string = "success"] Success, @@ -105,7 +105,7 @@ pub enum LettaToolReturnStatus { Error, } -#[derive(Serialize, Deserialize, TS, Debug)] +#[derive(Serialize, Deserialize, TS, Debug, Clone)] #[serde(tag = "message_type", rename_all = "snake_case")] #[ts(export)] pub enum LettaCompletionMessage { diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 8c20601..09bd591 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -5,6 +5,7 @@ use tauri_plugin_store::StoreExt; use tauri_plugin_window_state::{StateFlags, WindowExt as StateWindowExt}; mod cartesia; +mod conversation; mod devices; mod letta; mod secrets; -- 2.51.2