diff --git a/src-tauri/src/letta/commands.rs b/src-tauri/src/letta/commands.rs index 500b6fb..7b5f804 100644 --- a/src-tauri/src/letta/commands.rs +++ b/src-tauri/src/letta/commands.rs @@ -1,11 +1,19 @@ +use tauri::ipc::Channel; + use crate::{ - letta::{LettaAgentInfo, LettaConfigKey}, + letta::{types::LettaCompletionMessage, LettaAgentInfo, LettaConfigKey}, state::AppState, }; #[tauri::command] pub async fn list_agents(state: tauri::State<'_, AppState>) -> Result, ()> { - Ok(state.letta_manager.list_agents().await) + match state.letta_manager.list_agents().await { + Ok(res) => Ok(res), + Err(err) => { + eprintln!("failed to list agents: {}", err); + Ok(Vec::new()) + } + } } #[tauri::command] @@ -58,8 +66,22 @@ pub async fn set_letta_agent_id( pub async fn start_llm_completion( state: tauri::State<'_, AppState>, message: String, + on_event: Channel, ) -> Result<(), ()> { - state.letta_manager.start_completion(message).await; + match state.letta_manager.start_completion(message).await { + Ok(mut rx) => { + while let Some(ev) = rx.recv().await { + on_event + .send(ev) + .expect("failed to forward event to channel"); + } - Ok(()) + Ok(()) + } + Err(err) => { + eprintln!("failed to start completion: {}", err); + + Ok(()) + } + } } diff --git a/src-tauri/src/letta/mod.rs b/src-tauri/src/letta/mod.rs index 5c1c6cd..5d0132d 100644 --- a/src-tauri/src/letta/mod.rs +++ b/src-tauri/src/letta/mod.rs @@ -4,12 +4,12 @@ use futures_util::StreamExt; use reqwest::Client; use reqwest_eventsource::{Event, EventSource}; use serde_json::{from_str, json}; -use tauri::async_runtime::Mutex; +use tauri::async_runtime::{channel, spawn, Mutex, Receiver, Sender}; use tauri::Wry; use tauri_plugin_store::Store; use crate::secrets::SecretsManager; -use types::{LettaAgentInfo, LettaCompletionMessage, LettaConfigKey}; +use types::{LettaAgentInfo, LettaCompletionMessage, LettaConfigKey, LettaError}; pub mod commands; pub mod types; @@ -57,92 +57,93 @@ impl LettaManager { } } - pub async fn list_agents(&self) -> Vec { - match self + pub async fn list_agents(&self) -> Result, LettaError> { + let api_key = self .secrets_manager - .get_secret(crate::secrets::SecretName::LettaApiKey) - { - Ok(api_key) => { - let base_url = self.base_url.lock().await.to_owned(); - let req = self - .http_client - .get(format!("{base_url}/v1/agents/")) - .header("Authorization", format!("Bearer {api_key}")); - - let res = req.send().await.expect("failed to list Letta agents"); - let out = res - .json::>() - .await - .expect("failed to deserialize agent info list"); - - out - } - Err(_) => Vec::new(), - } + .get_secret(crate::secrets::SecretName::LettaApiKey)?; + let base_url = self.base_url.lock().await.to_owned(); + let req = self + .http_client + .get(format!("{base_url}/v1/agents/")) + .header("Authorization", format!("Bearer {api_key}")); + + let res = req.send().await.expect("failed to list Letta agents"); + let out = res.json::>().await?; + + Ok(out) } - pub async fn start_completion(&self, msg: String) { - match self + pub async fn start_completion( + &self, + msg: String, + ) -> Result, LettaError> { + let api_key = self .secrets_manager - .get_secret(crate::secrets::SecretName::LettaApiKey) - { - Ok(api_key) => { - let base_url = self.base_url.lock().await.to_owned(); - let agent_id = self.agent_id.lock().await.to_owned(); - let body = &json!({ - "messages": [ - { - "role": "user", - "content": [{ - "type": "text", - "text": msg, - }] - } - ], - "stream_tokens": true - }); - - println!("body: {:?}", body); - - let req = self - .http_client - .post(format!("{base_url}/v1/agents/{agent_id}/messages/stream")) - .header("Authorization", format!("Bearer {api_key}")) - .header("Content-Type", "application/json") - .json(body); - - let mut source = - EventSource::new(req).expect("could not convert request to event source"); - - while let Some(event) = source.next().await { - match event { - Ok(Event::Open) => println!("stream opened"), - Ok(Event::Message(msg)) => { - if msg.data == "[DONE]" { - continue; - }; - - match from_str::(&msg.data) { - Ok(content) => { - println!("parsed content: {:?}", content) - } - Err(err) => { - eprintln!( - "failed to parse message: {:?}, err: {}", - msg.data.clone(), - err - ) - } - } - } - Err(err) => { - eprintln!("got stream error: {}", err); - source.close(); - } + .get_secret(crate::secrets::SecretName::LettaApiKey)?; + let base_url = self.base_url.lock().await.to_owned(); + let agent_id = self.agent_id.lock().await.to_owned(); + let body = &json!({ + "messages": [ + { + "role": "user", + "content": [{ + "type": "text", + "text": msg, + }] + } + ], + "stream_tokens": true + }); + + println!("body: {:?}", body); + + let req = self + .http_client + .post(format!("{base_url}/v1/agents/{agent_id}/messages/stream")) + .header("Authorization", format!("Bearer {api_key}")) + .header("Content-Type", "application/json") + .json(body); + + let source = EventSource::new(req).expect("failed to clone request"); + + let (tx, rx) = channel::(100); + let handler = handle_completion_messages(source, tx); + + spawn(handler); + + return Ok(rx); + } +} + +async fn handle_completion_messages( + mut source: EventSource, + sender: Sender, +) { + while let Some(event) = source.next().await { + match event { + Ok(Event::Open) => println!("stream opened"), + Ok(Event::Message(msg)) => { + if msg.data == "[DONE]" { + continue; + }; + + match from_str::(&msg.data) { + Ok(content) => { + sender.send(content).await.expect("failed to forward event"); + } + Err(err) => { + eprintln!( + "failed to parse message: {:?}, err: {}", + msg.data.clone(), + err + ) } } } - Err(err) => eprintln!("could not fetch Letta API key: {}", err), + Err(err) => { + eprintln!("got stream error: {}", err); + source.close(); + } } } } diff --git a/src-tauri/src/letta/types.rs b/src-tauri/src/letta/types.rs index 44977f0..3dbbb36 100644 --- a/src-tauri/src/letta/types.rs +++ b/src-tauri/src/letta/types.rs @@ -1,8 +1,47 @@ +use core::fmt; + use serde::{Deserialize, Serialize}; use serde_string_enum::{DeserializeLabeledStringEnum, SerializeLabeledStringEnum}; use strum::{Display, EnumString}; use ts_rs::TS; +#[derive(Debug)] +pub enum LettaError { + ConfigurationError(String), + ApiKeyError(keyring::Error), + RequestError(reqwest::Error), + EventSourceError(reqwest_eventsource::Error), +} + +impl fmt::Display for LettaError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + LettaError::ConfigurationError(msg) => write!(f, "Configuration error: {}", msg), + LettaError::ApiKeyError(err) => write!(f, "API key error: {}", err), + LettaError::RequestError(err) => write!(f, "Letta request error: {}", err), + LettaError::EventSourceError(err) => write!(f, "Letta event source error: {}", err), + } + } +} + +impl From for LettaError { + fn from(err: keyring::Error) -> Self { + LettaError::ApiKeyError(err) + } +} + +impl From for LettaError { + fn from(err: reqwest::Error) -> Self { + LettaError::RequestError(err) + } +} + +impl From for LettaError { + fn from(err: reqwest_eventsource::Error) -> Self { + LettaError::EventSourceError(err) + } +} + #[derive(Serialize, Deserialize, Clone, Display, EnumString)] #[strum(prefix = "letta")] pub enum LettaConfigKey { diff --git a/src/lib/rust/LettaCompletionMessage.ts b/src/lib/rust/LettaCompletionMessage.ts index 6cea7ca..e4dfba7 100644 --- a/src/lib/rust/LettaCompletionMessage.ts +++ b/src/lib/rust/LettaCompletionMessage.ts @@ -5,4 +5,137 @@ import type { LettaReasoningSource } from "./LettaReasoningSource"; import type { LettaToolCall } from "./LettaToolCall"; import type { LettaToolReturnStatus } from "./LettaToolReturnStatus"; -export type LettaCompletionMessage = { "message_type": "system_message", id: string, date: string, content: string, name: string | null, otid: string | null, sender_id: string | null, step_id: string | null, is_err: boolean | null, seq_id: bigint | null, run_id: string | null, } | { "message_type": "user_message", id: string, date: string, content: Array, name: string | null, otid: string | null, sender_id: string | null, step_id: string | null, is_err: boolean | null, seq_id: bigint | null, run_id: string | null, } | { "message_type": "reasoning_message", id: string, date: string, reasoning: string, name: string | null, otid: string | null, sender_id: string | null, step_id: string | null, is_err: boolean | null, seq_id: bigint | null, run_id: string | null, source: LettaReasoningSource | null, signature: string | null, } | { "message_type": "hidden_reasoning_message", id: string, date: string, state: LettaHiddenReasoningState, name: string | null, otid: string | null, sender_id: string | null, step_id: string | null, is_err: boolean | null, seq_id: bigint | null, run_id: string | null, hidden_reasoning: string | null, } | { "message_type": "tool_call_message", id: string, date: string, tool_call: LettaToolCall, name: string | null, otid: string | null, sender_id: string | null, step_id: string | null, is_err: boolean | null, seq_id: bigint | null, run_id: string | null, } | { "message_type": "tool_return_message", id: string, date: string, tool_return: string, status: LettaToolReturnStatus, tool_call_id: string, name: string | null, otid: string | null, sender_id: string | null, step_id: string | null, is_err: boolean | null, seq_id: bigint | null, run_id: string | null, stdout: Array | null, stderr: Array | null, } | { "message_type": "assistant_message", id: string, date: string, content: Array, name: string | null, otid: string | null, sender_id: string | null, step_id: string | null, is_err: boolean | null, seq_id: bigint | null, run_id: string | null, } | { "message_type": "approval_request_message", id: string, date: string, tool_call: LettaToolCall, name: string | null, otid: string | null, sender_id: string | null, step_id: string | null, is_err: boolean | null, seq_id: bigint | null, run_id: string | null, } | { "message_type": "approval_response_message", id: string, date: string, approve: boolean, approval_request_id: string, name: string | null, otid: string | null, sender_id: string | null, step_id: string | null, is_err: boolean | null, seq_id: bigint | null, run_id: string | null, reason: string | null, } | { "message_type": "stop_reason", stop_reason: string, } | { "message_type": "usage_statistics", completion_tokens: bigint, prompt_tokens: bigint, step_count: bigint, }; +export type LettaCompletionMessage = + | { + message_type: "system_message"; + id: string; + date: string; + content: string; + name: string | null; + otid: string | null; + sender_id: string | null; + step_id: string | null; + is_err: boolean | null; + seq_id: bigint | null; + run_id: string | null; + } + | { + message_type: "user_message"; + id: string; + date: string; + content: Array; + name: string | null; + otid: string | null; + sender_id: string | null; + step_id: string | null; + is_err: boolean | null; + seq_id: bigint | null; + run_id: string | null; + } + | { + message_type: "reasoning_message"; + id: string; + date: string; + reasoning: string; + name: string | null; + otid: string | null; + sender_id: string | null; + step_id: string | null; + is_err: boolean | null; + seq_id: bigint | null; + run_id: string | null; + source: LettaReasoningSource | null; + signature: string | null; + } + | { + message_type: "hidden_reasoning_message"; + id: string; + date: string; + state: LettaHiddenReasoningState; + name: string | null; + otid: string | null; + sender_id: string | null; + step_id: string | null; + is_err: boolean | null; + seq_id: bigint | null; + run_id: string | null; + hidden_reasoning: string | null; + } + | { + message_type: "tool_call_message"; + id: string; + date: string; + tool_call: LettaToolCall; + name: string | null; + otid: string | null; + sender_id: string | null; + step_id: string | null; + is_err: boolean | null; + seq_id: bigint | null; + run_id: string | null; + } + | { + message_type: "tool_return_message"; + id: string; + date: string; + tool_return: string; + status: LettaToolReturnStatus; + tool_call_id: string; + name: string | null; + otid: string | null; + sender_id: string | null; + step_id: string | null; + is_err: boolean | null; + seq_id: bigint | null; + run_id: string | null; + stdout: Array | null; + stderr: Array | null; + } + | { + message_type: "assistant_message"; + id: string; + date: string; + content: Array; + name: string | null; + otid: string | null; + sender_id: string | null; + step_id: string | null; + is_err: boolean | null; + seq_id: bigint | null; + run_id: string | null; + } + | { + message_type: "approval_request_message"; + id: string; + date: string; + tool_call: LettaToolCall; + name: string | null; + otid: string | null; + sender_id: string | null; + step_id: string | null; + is_err: boolean | null; + seq_id: bigint | null; + run_id: string | null; + } + | { + message_type: "approval_response_message"; + id: string; + date: string; + approve: boolean; + approval_request_id: string; + name: string | null; + otid: string | null; + sender_id: string | null; + step_id: string | null; + is_err: boolean | null; + seq_id: bigint | null; + run_id: string | null; + reason: string | null; + } + | { message_type: "stop_reason"; stop_reason: string } + | { + message_type: "usage_statistics"; + completion_tokens: bigint; + prompt_tokens: bigint; + step_count: bigint; + }; diff --git a/src/routes/+page.svelte b/src/routes/+page.svelte index ae17e80..3fade17 100644 --- a/src/routes/+page.svelte +++ b/src/routes/+page.svelte @@ -1,5 +1,6 @@