diff --git a/src/agent/job_monitor.rs b/src/agent/job_monitor.rs index 8d302c1d..02f5e3e2 100644 --- a/src/agent/job_monitor.rs +++ b/src/agent/job_monitor.rs @@ -170,7 +170,9 @@ pub fn spawn_completion_watcher( tokio::spawn(async move { loop { match event_rx.recv().await { - Ok((ev_job_id, _user_id, SseEvent::JobResult { status, .. })) if ev_job_id == job_id => { + Ok((ev_job_id, _user_id, SseEvent::JobResult { status, .. })) + if ev_job_id == job_id => + { let target = if status == "completed" { JobState::Completed } else { diff --git a/src/app.rs b/src/app.rs index 4e1d506f..d621c53a 100644 --- a/src/app.rs +++ b/src/app.rs @@ -356,7 +356,7 @@ impl AppBuilder { if is_multi_tenant { let resolver = Arc::new( crate::tools::builtin::memory::PerUserWorkspaceResolver::new( - self.db.as_ref().unwrap().clone(), + Arc::clone(db), embeddings.clone(), emb_cache_config, self.config.search.clone(), @@ -364,7 +364,9 @@ impl AppBuilder { ), ); tools.register_memory_tools_with_resolver(resolver); - tracing::info!("Memory tools configured with per-user workspace resolver (multi-tenant mode)"); + tracing::info!( + "Memory tools configured with per-user workspace resolver (multi-tenant mode)" + ); } else { tools.register_memory_tools(Arc::clone(&ws)); } diff --git a/src/channels/web/auth.rs b/src/channels/web/auth.rs index 49688bf8..7dc8adb4 100644 --- a/src/channels/web/auth.rs +++ b/src/channels/web/auth.rs @@ -12,6 +12,7 @@ use axum::{ middleware::Next, response::{IntoResponse, Response}, }; +use sha2::{Digest, Sha256}; use subtle::ConstantTimeEq; /// Identity resolved from a bearer token. @@ -22,63 +23,88 @@ pub struct UserIdentity { pub workspace_read_scopes: Vec, } -/// Multi-user auth state: maps tokens to user identities. +/// Hash a token with SHA-256 for constant-size, timing-safe storage. +fn hash_token(token: &str) -> [u8; 32] { + let mut hasher = Sha256::new(); + hasher.update(token.as_bytes()); + hasher.finalize().into() +} + +/// Multi-user auth state: maps token hashes to user identities. +/// +/// Tokens are SHA-256 hashed on construction so they are never stored in +/// plaintext. Authentication compares fixed-size (32-byte) digests using +/// constant-time comparison, eliminating both length-oracle timing leaks +/// and accidental token exposure in memory dumps. /// /// In single-user mode (the default), contains exactly one entry. #[derive(Clone)] pub struct MultiAuthState { - tokens: HashMap, + /// Maps SHA-256(token) → identity. Tokens are never stored in cleartext. + hashed_tokens: Vec<([u8; 32], UserIdentity)>, + /// Original first token kept only for single-user startup printing. + /// Not used for authentication. + display_token: Option, } impl MultiAuthState { /// Create a single-user auth state (backwards compatible). pub fn single(token: String, user_id: String) -> Self { - let mut tokens = HashMap::new(); - tokens.insert( - token, - UserIdentity { - user_id, - workspace_read_scopes: Vec::new(), - }, - ); - Self { tokens } + let hash = hash_token(&token); + Self { + hashed_tokens: vec![( + hash, + UserIdentity { + user_id, + workspace_read_scopes: Vec::new(), + }, + )], + display_token: Some(token), + } } /// Create a multi-user auth state from a map of tokens to identities. pub fn multi(tokens: HashMap) -> Self { - Self { tokens } + let hashed_tokens: Vec<([u8; 32], UserIdentity)> = tokens + .into_iter() + .map(|(tok, identity)| (hash_token(&tok), identity)) + .collect(); + Self { + hashed_tokens, + display_token: None, + } } /// Authenticate a token, returning the associated identity if valid. /// - /// Uses constant-time comparison (`subtle::ConstantTimeEq`) to prevent - /// timing side-channels that could leak token information. Iterates all + /// Uses SHA-256 hashing + constant-time comparison (`subtle::ConstantTimeEq`) + /// to prevent timing side-channels. Both the candidate and stored tokens are + /// hashed to 32-byte digests, eliminating length-oracle leaks. Iterates all /// entries regardless of match to avoid early-exit timing differences. /// O(n) in the number of configured users — negligible for typical /// deployments (< 10 users). pub fn authenticate(&self, candidate: &str) -> Option<&UserIdentity> { - let candidate_bytes = candidate.as_bytes(); + let candidate_hash = hash_token(candidate); let mut matched: Option<&UserIdentity> = None; - for (token, identity) in &self.tokens { - let token_bytes = token.as_bytes(); - // ct_eq requires equal lengths; pad comparison to avoid length leak - if candidate_bytes.len() == token_bytes.len() - && bool::from(candidate_bytes.ct_eq(token_bytes)) - { + for (stored_hash, identity) in &self.hashed_tokens { + if bool::from(candidate_hash.ct_eq(stored_hash)) { matched = Some(identity); } } matched } - /// Get the first token (for backwards-compatible printing at startup). + /// Get the first token for backwards-compatible printing at startup. + /// + /// Only available in single-user mode; returns `None` in multi-user mode + /// to avoid exposing tokens. pub fn first_token(&self) -> Option<&str> { - self.tokens.keys().next().map(|s| s.as_str()) + self.display_token.as_deref() } /// Get the first user identity (for single-user fallback). pub fn first_identity(&self) -> Option<&UserIdentity> { - self.tokens.values().next() + self.hashed_tokens.first().map(|(_, id)| id) } } @@ -175,9 +201,6 @@ pub async fn auth_middleware( (StatusCode::UNAUTHORIZED, "Invalid or missing auth token").into_response() } -// Keep the old type as an alias for any external references during migration. -pub type AuthState = MultiAuthState; - #[cfg(test)] mod tests { use super::*; diff --git a/src/channels/web/handlers/chat.rs b/src/channels/web/handlers/chat.rs index 640ee260..9753c015 100644 --- a/src/channels/web/handlers/chat.rs +++ b/src/channels/web/handlers/chat.rs @@ -173,20 +173,26 @@ pub async fn chat_auth_token_handler( resp.instructions = result.verification.as_ref().map(|v| v.instructions.clone()); if result.verification.is_some() { - state.sse.broadcast(SseEvent::AuthRequired { - extension_name: req.extension_name.clone(), - instructions: Some(result.message), - auth_url: None, - setup_url: None, - }); + state.sse.broadcast_for_user( + &user.user_id, + SseEvent::AuthRequired { + extension_name: req.extension_name.clone(), + instructions: Some(result.message), + auth_url: None, + setup_url: None, + }, + ); } else { clear_auth_mode(&state, &user.user_id).await; - state.sse.broadcast(SseEvent::AuthCompleted { - extension_name: req.extension_name.clone(), - success: true, - message: result.message, - }); + state.sse.broadcast_for_user( + &user.user_id, + SseEvent::AuthCompleted { + extension_name: req.extension_name.clone(), + success: true, + message: result.message, + }, + ); } Ok(Json(resp)) @@ -194,12 +200,15 @@ pub async fn chat_auth_token_handler( Err(e) => { let msg = e.to_string(); if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { - state.sse.broadcast(SseEvent::AuthRequired { - extension_name: req.extension_name.clone(), - instructions: Some(msg.clone()), - auth_url: None, - setup_url: None, - }); + state.sse.broadcast_for_user( + &user.user_id, + SseEvent::AuthRequired { + extension_name: req.extension_name.clone(), + instructions: Some(msg.clone()), + auth_url: None, + setup_url: None, + }, + ); } Ok(Json(ActionResponse::fail(msg))) } @@ -291,7 +300,9 @@ pub async fn chat_history_handler( "Session manager not available".to_string(), ))?; - let session = session_manager.get_or_create_session(&identity.user_id).await; + let session = session_manager + .get_or_create_session(&identity.user_id) + .await; let limit = query.limit.unwrap_or(50); let before_cursor = query @@ -451,7 +462,9 @@ pub async fn chat_threads_handler( "Session manager not available".to_string(), ))?; - let session = session_manager.get_or_create_session(&identity.user_id).await; + let session = session_manager + .get_or_create_session(&identity.user_id) + .await; // Try DB first for persistent thread list if let Some(ref store) = state.store { @@ -552,7 +565,9 @@ pub async fn chat_new_thread_handler( "Session manager not available".to_string(), ))?; - let session = session_manager.get_or_create_session(&identity.user_id).await; + let session = session_manager + .get_or_create_session(&identity.user_id) + .await; let (thread_id, info) = { let mut sess = session.lock().await; let thread = sess.create_thread(); diff --git a/src/channels/web/handlers/jobs.rs b/src/channels/web/handlers/jobs.rs index 6901d943..1064dd12 100644 --- a/src/channels/web/handlers/jobs.rs +++ b/src/channels/web/handlers/jobs.rs @@ -375,152 +375,166 @@ pub async fn jobs_restart_handler( .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?; // Try sandbox job restart first. - if let Ok(Some(old_job)) = store.get_sandbox_job(old_job_id).await { - if old_job.user_id != user.user_id { - return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + match store.get_sandbox_job(old_job_id).await { + Ok(Some(old_job)) => { + if old_job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + if old_job.status != "interrupted" && old_job.status != "failed" { + return Err(( + StatusCode::CONFLICT, + format!("Cannot restart job in state '{}'", old_job.status), + )); + } + + let jm = state.job_manager.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Sandbox not enabled".to_string(), + ))?; + + // Enrich the task with failure context. + let task = if let Some(ref reason) = old_job.failure_reason { + format!( + "Previous attempt failed: {}. Retry: {}", + reason, old_job.task + ) + } else { + old_job.task.clone() + }; + + let new_job_id = Uuid::new_v4(); + let now = chrono::Utc::now(); + + let record = crate::history::SandboxJobRecord { + id: new_job_id, + task: task.clone(), + status: "creating".to_string(), + user_id: old_job.user_id.clone(), + project_dir: old_job.project_dir.clone(), + success: None, + failure_reason: None, + created_at: now, + started_at: None, + completed_at: None, + credential_grants_json: old_job.credential_grants_json.clone(), + }; + store + .save_sandbox_job(&record) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + let mode = match store.get_sandbox_job_mode(old_job_id).await { + Ok(Some(m)) if m == "claude_code" => { + crate::orchestrator::job_manager::JobMode::ClaudeCode + } + _ => crate::orchestrator::job_manager::JobMode::Worker, + }; + + let credential_grants: Vec = + serde_json::from_str(&old_job.credential_grants_json).unwrap_or_else(|e| { + tracing::warn!( + job_id = %old_job.id, + "Failed to deserialize credential grants from stored job: {}. \ + Restarted job will have no credentials.", + e + ); + vec![] + }); + + let project_dir = std::path::PathBuf::from(&old_job.project_dir); + let _token = jm + .create_job( + new_job_id, + &task, + Some(project_dir), + mode, + credential_grants, + ) + .await + .map_err(|e| { + ( + StatusCode::INTERNAL_SERVER_ERROR, + format!("Failed to create container: {}", e), + ) + })?; + + store + .update_sandbox_job_status(new_job_id, "running", None, None, Some(now), None) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + return Ok(Json(serde_json::json!({ + "status": "restarted", + "old_job_id": old_job_id, + "new_job_id": new_job_id, + }))); } - if old_job.status != "interrupted" && old_job.status != "failed" { + Ok(None) => {} + Err(e) => { return Err(( - StatusCode::CONFLICT, - format!("Cannot restart job in state '{}'", old_job.status), + StatusCode::INTERNAL_SERVER_ERROR, + format!("Database error: {}", e), )); } - - let jm = state.job_manager.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Sandbox not enabled".to_string(), - ))?; - - // Enrich the task with failure context. - let task = if let Some(ref reason) = old_job.failure_reason { - format!( - "Previous attempt failed: {}. Retry: {}", - reason, old_job.task - ) - } else { - old_job.task.clone() - }; - - let new_job_id = Uuid::new_v4(); - let now = chrono::Utc::now(); - - let record = crate::history::SandboxJobRecord { - id: new_job_id, - task: task.clone(), - status: "creating".to_string(), - user_id: old_job.user_id.clone(), - project_dir: old_job.project_dir.clone(), - success: None, - failure_reason: None, - created_at: now, - started_at: None, - completed_at: None, - credential_grants_json: old_job.credential_grants_json.clone(), - }; - store - .save_sandbox_job(&record) - .await - .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; - - let mode = match store.get_sandbox_job_mode(old_job_id).await { - Ok(Some(m)) if m == "claude_code" => { - crate::orchestrator::job_manager::JobMode::ClaudeCode - } - _ => crate::orchestrator::job_manager::JobMode::Worker, - }; - - let credential_grants: Vec = - serde_json::from_str(&old_job.credential_grants_json).unwrap_or_else(|e| { - tracing::warn!( - job_id = %old_job.id, - "Failed to deserialize credential grants from stored job: {}. \ - Restarted job will have no credentials.", - e - ); - vec![] - }); - - let project_dir = std::path::PathBuf::from(&old_job.project_dir); - let _token = jm - .create_job( - new_job_id, - &task, - Some(project_dir), - mode, - credential_grants, - ) - .await - .map_err(|e| { - ( - StatusCode::INTERNAL_SERVER_ERROR, - format!("Failed to create container: {}", e), - ) - })?; - - store - .update_sandbox_job_status(new_job_id, "running", None, None, Some(now), None) - .await - .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; - - return Ok(Json(serde_json::json!({ - "status": "restarted", - "old_job_id": old_job_id, - "new_job_id": new_job_id, - }))); } // Try agent job restart: dispatch a new job via the scheduler. - if let Ok(Some(old_job)) = store.get_job(old_job_id).await { - if old_job.user_id != user.user_id { - return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + match store.get_job(old_job_id).await { + Ok(Some(old_job)) => { + if old_job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + if old_job.state.is_active() { + return Err(( + StatusCode::CONFLICT, + format!("Cannot restart job in state '{}'", old_job.state), + )); + } + + let slot = state.scheduler.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Scheduler not available".to_string(), + ))?; + let scheduler_guard = slot.read().await; + let scheduler = scheduler_guard.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Agent not started yet".to_string(), + ))?; + + // Look up failure reason (O(1) point lookup). + let failure_reason = store + .get_agent_job_failure_reason(old_job_id) + .await + .ok() + .flatten() + .unwrap_or_default(); + + let title = if !failure_reason.is_empty() { + format!( + "Previous attempt failed: {}. Retry: {}", + failure_reason, old_job.title + ) + } else { + old_job.title.clone() + }; + + let new_job_id = scheduler + .dispatch_job(&old_job.user_id, &title, &old_job.description, None) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + Ok(Json(serde_json::json!({ + "status": "restarted", + "old_job_id": old_job_id, + "new_job_id": new_job_id, + }))) } - if old_job.state.is_active() { - return Err(( - StatusCode::CONFLICT, - format!("Cannot restart job in state '{}'", old_job.state), - )); - } - - let slot = state.scheduler.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Scheduler not available".to_string(), - ))?; - let scheduler_guard = slot.read().await; - let scheduler = scheduler_guard.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Agent not started yet".to_string(), - ))?; - - // Look up failure reason (O(1) point lookup). - let failure_reason = store - .get_agent_job_failure_reason(old_job_id) - .await - .ok() - .flatten() - .unwrap_or_default(); - - let title = if !failure_reason.is_empty() { - format!( - "Previous attempt failed: {}. Retry: {}", - failure_reason, old_job.title - ) - } else { - old_job.title.clone() - }; - - let new_job_id = scheduler - .dispatch_job(&old_job.user_id, &title, &old_job.description, None) - .await - .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; - - return Ok(Json(serde_json::json!({ - "status": "restarted", - "old_job_id": old_job_id, - "new_job_id": new_job_id, - }))); + Ok(None) => Err((StatusCode::NOT_FOUND, "Job not found".to_string())), + Err(e) => Err(( + StatusCode::INTERNAL_SERVER_ERROR, + format!("Database error: {}", e), + )), } - - Err((StatusCode::NOT_FOUND, "Job not found".to_string())) } /// Submit a follow-up prompt to a running job. diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index bed844cd..5a51c3ce 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -30,7 +30,9 @@ use crate::agent::SessionManager; use crate::bootstrap::ironclaw_base_dir; use crate::channels::IncomingMessage; use crate::channels::relay::DEFAULT_RELAY_NAME; -use crate::channels::web::auth::{AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware}; +use crate::channels::web::auth::{ + AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware, +}; use crate::channels::web::handlers::jobs::{ job_files_list_handler, job_files_read_handler, jobs_cancel_handler, jobs_detail_handler, jobs_events_handler, jobs_list_handler, jobs_prompt_handler, jobs_restart_handler, @@ -770,11 +772,14 @@ async fn oauth_callback_handler( ); // Notify UI so auth card can show error instead of staying stuck if let Some(ref sse) = flow.sse_manager { - sse.broadcast_for_user(&flow.user_id, SseEvent::AuthCompleted { - extension_name: flow.extension_name.clone(), - success: false, - message: "OAuth flow expired. Please try again.".to_string(), - }); + sse.broadcast_for_user( + &flow.user_id, + SseEvent::AuthCompleted { + extension_name: flow.extension_name.clone(), + success: false, + message: "OAuth flow expired. Please try again.".to_string(), + }, + ); } clear_auth_mode(&state, &flow.user_id).await; return oauth_error_page(&flow.display_name); @@ -909,11 +914,14 @@ async fn oauth_callback_handler( // Broadcast SSE event to notify the web UI if let Some(ref sse) = flow.sse_manager { - sse.broadcast_for_user(&flow.user_id, SseEvent::AuthCompleted { - extension_name: flow.extension_name, - success, - message: final_message.clone(), - }); + sse.broadcast_for_user( + &flow.user_id, + SseEvent::AuthCompleted { + extension_name: flow.extension_name, + success, + message: final_message.clone(), + }, + ); } let html = oauth_defaults::landing_html(&flow.display_name, success); @@ -1115,7 +1123,10 @@ async fn slack_relay_oauth_callback_handler( } // Delete the nonce (one-time use) - let _ = ext_mgr.secrets().delete(&state.default_user_id, &state_key).await; + let _ = ext_mgr + .secrets() + .delete(&state.default_user_id, &state_key) + .await; let result: Result<(), String> = async { let store = state.store.as_ref().ok_or_else(|| { @@ -1126,7 +1137,11 @@ async fn slack_relay_oauth_callback_handler( // Store team_id in settings let team_id_key = format!("relay:{}:team_id", DEFAULT_RELAY_NAME); let _ = store - .set_setting(&state.default_user_id, &team_id_key, &serde_json::json!(team_id)) + .set_setting( + &state.default_user_id, + &team_id_key, + &serde_json::json!(team_id), + ) .await; // Activate the relay channel @@ -1750,52 +1765,52 @@ async fn chat_threads_handler( .list_conversations_all_channels(&user.user_id, 50) .await { - Ok(summaries) => { - let mut assistant_thread = None; - let mut threads = Vec::new(); + Ok(summaries) => { + let mut assistant_thread = None; + let mut threads = Vec::new(); - for s in &summaries { - let info = ThreadInfo { - id: s.id, - state: "Idle".to_string(), - turn_count: s.message_count.max(0) as usize, - created_at: s.started_at.to_rfc3339(), - updated_at: s.last_activity.to_rfc3339(), - title: s.title.clone(), - thread_type: s.thread_type.clone(), - channel: Some(s.channel.clone()), - }; + for s in &summaries { + let info = ThreadInfo { + id: s.id, + state: "Idle".to_string(), + turn_count: s.message_count.max(0) as usize, + created_at: s.started_at.to_rfc3339(), + updated_at: s.last_activity.to_rfc3339(), + title: s.title.clone(), + thread_type: s.thread_type.clone(), + channel: Some(s.channel.clone()), + }; - if s.id == assistant_id { - assistant_thread = Some(info); - } else { - threads.push(info); + if s.id == assistant_id { + assistant_thread = Some(info); + } else { + threads.push(info); + } } - } - // If assistant wasn't in the list (0 messages), synthesize it - if assistant_thread.is_none() { - assistant_thread = Some(ThreadInfo { - id: assistant_id, - state: "Idle".to_string(), - turn_count: 0, - created_at: chrono::Utc::now().to_rfc3339(), - updated_at: chrono::Utc::now().to_rfc3339(), - title: None, - thread_type: Some("assistant".to_string()), - channel: Some("gateway".to_string()), - }); - } + // If assistant wasn't in the list (0 messages), synthesize it + if assistant_thread.is_none() { + assistant_thread = Some(ThreadInfo { + id: assistant_id, + state: "Idle".to_string(), + turn_count: 0, + created_at: chrono::Utc::now().to_rfc3339(), + updated_at: chrono::Utc::now().to_rfc3339(), + title: None, + thread_type: Some("assistant".to_string()), + channel: Some("gateway".to_string()), + }); + } - return Ok(Json(ThreadListResponse { - assistant_thread, - threads, - active_thread: sess.active_thread, - })); - } - Err(e) => { - tracing::error!(user_id = %user.user_id, error = %e, "DB error listing threads; falling back to in-memory"); - } + return Ok(Json(ThreadListResponse { + assistant_thread, + threads, + active_thread: sess.active_thread, + })); + } + Err(e) => { + tracing::error!(user_id = %user.user_id, error = %e, "DB error listing threads; falling back to in-memory"); + } } } @@ -1890,14 +1905,10 @@ async fn resolve_workspace( if let Some(ref pool) = state.workspace_pool { return Ok(pool.get_or_create(user).await); } - state - .workspace - .as_ref() - .cloned() - .ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Workspace not available".to_string(), - )) + state.workspace.as_ref().cloned().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Workspace not available".to_string(), + )) } #[derive(Deserialize)] @@ -2618,7 +2629,10 @@ async fn extensions_setup_submit_handler( // through to the LLM instead of being intercepted as a token. clear_auth_mode(&state, &user.user_id).await; - match ext_mgr.configure(&name, &req.secrets, &req.fields, &user.user_id).await { + match ext_mgr + .configure(&name, &req.secrets, &req.fields, &user.user_id) + .await + { Ok(result) => { let mut resp = if result.verification.is_some() || result.activated { ActionResponse::ok(result.message) @@ -2635,11 +2649,14 @@ async fn extensions_setup_submit_handler( if result.verification.is_none() { // Broadcast auth_completed so the chat UI can dismiss any in-progress // auth card or setup modal that was triggered by tool_auth/tool_activate. - state.sse.broadcast_for_user(&user.user_id, SseEvent::AuthCompleted { - extension_name: name.clone(), - success: result.activated, - message: resp.message.clone(), - }); + state.sse.broadcast_for_user( + &user.user_id, + SseEvent::AuthCompleted { + extension_name: name.clone(), + success: result.activated, + message: resp.message.clone(), + }, + ); } Ok(Json(resp)) } @@ -3289,12 +3306,18 @@ mod tests { "telegram_bot_token": "123456789:ABCdefGhI" } }); - let req = axum::http::Request::builder() + let mut req = axum::http::Request::builder() .method("POST") .uri("/api/extensions/telegram/setup") .header("content-type", "application/json") .body(Body::from(req_body.to_string())) .expect("request"); + // Inject AuthenticatedUser so the handler's extractor succeeds + // without needing the full auth middleware layer. + req.extensions_mut().insert(UserIdentity { + user_id: "test".to_string(), + workspace_read_scopes: Vec::new(), + }); let resp = ServiceExt::>::oneshot(app, req) .await @@ -3316,7 +3339,12 @@ mod tests { break; } match timeout(remaining, receiver.recv()).await { - Ok(Ok(scoped)) if matches!(scoped.event, crate::channels::web::types::SseEvent::AuthRequired { .. }) => { + Ok(Ok(scoped)) + if matches!( + scoped.event, + crate::channels::web::types::SseEvent::AuthRequired { .. } + ) => + { panic!("verification responses should not emit auth_required SSE events") } Ok(Ok(_)) => continue, diff --git a/src/channels/web/sse.rs b/src/channels/web/sse.rs index 8654f7bd..6f30f463 100644 --- a/src/channels/web/sse.rs +++ b/src/channels/web/sse.rs @@ -129,10 +129,10 @@ impl SseManager { // Global events (user_id=None) always pass through. // Scoped events only pass if the subscriber matches (or subscriber is unscoped). match (&user_id, &scoped.user_id) { - (_, None) => Some(scoped.event), // global -> all - (None, _) => Some(scoped.event), // unscoped subscriber -> all + (_, None) => Some(scoped.event), // global -> all + (None, _) => Some(scoped.event), // unscoped subscriber -> all (Some(sub), Some(ev)) if sub == ev => Some(scoped.event), // match - _ => None, // different user -> skip + _ => None, // different user -> skip } } Err(_) => None, @@ -178,8 +178,14 @@ impl SseManager { }, Err(_) => None, }) - .map(|event| { - let data = serde_json::to_string(&event).unwrap_or_default(); + .filter_map(|event| { + let data = match serde_json::to_string(&event) { + Ok(s) => s, + Err(e) => { + tracing::warn!("Failed to serialize SSE event: {}", e); + return None; + } + }; let event_type = match &event { SseEvent::Response { .. } => "response", SseEvent::Thinking { .. } => "thinking", @@ -204,7 +210,7 @@ impl SseManager { SseEvent::TurnCost { .. } => "turn_cost", SseEvent::ExtensionStatus { .. } => "extension_status", }; - Ok(Event::default().event(event_type).data(data)) + Some(Ok(Event::default().event(event_type).data(data))) }); // Wrap in a stream that decrements on drop diff --git a/src/channels/web/test_helpers.rs b/src/channels/web/test_helpers.rs index 87b1c86d..802512a6 100644 --- a/src/channels/web/test_helpers.rs +++ b/src/channels/web/test_helpers.rs @@ -104,7 +104,7 @@ impl TestGatewayBuilder { let state = self.build(); let addr: SocketAddr = "127.0.0.1:0" .parse() - .expect("hard-coded address must parse"); + .expect("hard-coded address must parse"); // safety: constant literal let bound = start_server(addr, state.clone(), auth).await?; Ok((bound, state)) } @@ -118,7 +118,7 @@ impl TestGatewayBuilder { let state = self.build(); let addr: SocketAddr = "127.0.0.1:0" .parse() - .expect("hard-coded address must parse"); + .expect("hard-coded address must parse"); // safety: constant literal let bound = start_server(addr, state.clone(), auth).await?; Ok((bound, state)) } diff --git a/src/channels/web/tests/multi_tenant.rs b/src/channels/web/tests/multi_tenant.rs index 7ffc73af..55010831 100644 --- a/src/channels/web/tests/multi_tenant.rs +++ b/src/channels/web/tests/multi_tenant.rs @@ -8,18 +8,19 @@ use std::collections::HashMap; use std::sync::Arc; use std::time::Duration; +use axum::Router; use axum::body::Body; use axum::http::{Method, Request, StatusCode}; use axum::middleware; use axum::routing::{delete, get, post}; -use axum::Router; use tower::ServiceExt; use uuid::Uuid; -use crate::channels::web::auth::{auth_middleware, AuthenticatedUser, MultiAuthState, UserIdentity}; +use crate::channels::web::auth::{ + AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware, +}; use crate::channels::web::server::{ - ActiveConfigSnapshot, GatewayState, PerUserRateLimiter, PromptQueue, RateLimiter, - WorkspacePool, + ActiveConfigSnapshot, GatewayState, PerUserRateLimiter, PromptQueue, RateLimiter, WorkspacePool, }; use crate::channels::web::sse::SseManager; @@ -88,19 +89,16 @@ fn build_state( #[cfg(feature = "libsql")] async fn test_db() -> (Arc, tempfile::TempDir) { use crate::db::Database; - let dir = tempfile::tempdir().expect("failed to create temp dir"); + let dir = tempfile::tempdir().expect("failed to create temp dir"); // safety: test-only let path = dir.path().join("test.db"); let backend = crate::db::libsql::LibSqlBackend::new_local(&path) .await - .expect("failed to create test LibSqlBackend"); + .expect("failed to create test LibSqlBackend"); // safety: test-only backend .run_migrations() .await - .expect("failed to run migrations"); - ( - Arc::new(backend) as Arc, - dir, - ) + .expect("failed to run migrations"); // safety: test-only + (Arc::new(backend) as Arc, dir) } /// Build a minimal Routine for testing. diff --git a/src/channels/web/ws.rs b/src/channels/web/ws.rs index 4ee6002a..3a601679 100644 --- a/src/channels/web/ws.rs +++ b/src/channels/web/ws.rs @@ -267,7 +267,10 @@ async fn handle_client_message( token, } => { if let Some(ref ext_mgr) = state.extension_manager { - match ext_mgr.configure_token(&extension_name, &token, user_id).await { + match ext_mgr + .configure_token(&extension_name, &token, user_id) + .await + { Ok(result) => { if result.verification.is_some() { state.sse.broadcast_for_user( diff --git a/src/config/channels.rs b/src/config/channels.rs index 45c180df..d9c2c0a9 100644 --- a/src/config/channels.rs +++ b/src/config/channels.rs @@ -158,10 +158,7 @@ impl ChannelsConfig { if layer.name.len() > 64 { return Err(ConfigError::InvalidValue { key: "MEMORY_LAYERS".to_string(), - message: format!( - "layer name '{}' exceeds 64 characters", - layer.name - ), + message: format!("layer name '{}' exceeds 64 characters", layer.name), }); } if !layer @@ -181,10 +178,7 @@ impl ChannelsConfig { if layer.scope.trim().is_empty() { return Err(ConfigError::InvalidValue { key: "MEMORY_LAYERS".to_string(), - message: format!( - "layer '{}' has an empty scope", - layer.name - ), + message: format!("layer '{}' has an empty scope", layer.name), }); } } @@ -205,20 +199,22 @@ impl ChannelsConfig { let user_tokens: Option> = match optional_env("GATEWAY_USER_TOKENS")? { Some(json_str) => { - let tokens: HashMap = - serde_json::from_str(&json_str).map_err(|e| { - ConfigError::InvalidValue { - key: "GATEWAY_USER_TOKENS".to_string(), - message: format!( - "must be valid JSON object mapping tokens to user configs: {e}" - ), - } - })?; + let tokens: HashMap = serde_json::from_str( + &json_str, + ) + .map_err(|e| ConfigError::InvalidValue { + key: "GATEWAY_USER_TOKENS".to_string(), + message: format!( + "must be valid JSON object mapping tokens to user configs: {e}" + ), + })?; if tokens.is_empty() { return Err(ConfigError::InvalidValue { - key: "GATEWAY_USER_TOKENS".to_string(), - message: "token map is empty — remove the variable to use single-user mode".to_string(), - }); + key: "GATEWAY_USER_TOKENS".to_string(), + message: + "token map is empty — remove the variable to use single-user mode" + .to_string(), + }); } for (tok, cfg) in &tokens { if cfg.user_id.trim().is_empty() { @@ -248,10 +244,7 @@ impl ChannelsConfig { if scope.len() > 128 { return Err(ConfigError::InvalidValue { key: "WORKSPACE_READ_SCOPES".to_string(), - message: format!( - "scope '{}...' exceeds 128 characters", - &scope[..32] - ), + message: format!("scope '{}...' exceeds 128 characters", &scope[..32]), }); } } diff --git a/src/extensions/manager.rs b/src/extensions/manager.rs index 94163b26..7da9e980 100644 --- a/src/extensions/manager.rs +++ b/src/extensions/manager.rs @@ -1024,10 +1024,7 @@ impl ExtensionManager { } /// Set the SSE broadcast sender for pushing extension status events to the web UI. - pub async fn set_sse_sender( - &self, - sse: Arc, - ) { + pub async fn set_sse_sender(&self, sse: Arc) { *self.sse_manager.write().await = Some(sse); } @@ -1177,7 +1174,9 @@ impl ExtensionManager { &self, name: &str, url: Option<&str>, - kind_hint: Option, user_id: &str) -> Result { + kind_hint: Option, + user_id: &str, + ) -> Result { let sanitized_url = url.map(sanitize_url_for_logging); tracing::info!(extension = %name, url = ?sanitized_url, kind = ?kind_hint, "Installing extension"); Self::validate_extension_name(name)?; @@ -1241,7 +1240,11 @@ impl ExtensionManager { } /// Activate an installed (and optionally authenticated) extension. - pub async fn activate(&self, name: &str, user_id: &str) -> Result { + pub async fn activate( + &self, + name: &str, + user_id: &str, + ) -> Result { Self::validate_extension_name(name)?; let kind = self.determine_installed_kind(name, user_id).await?; @@ -1270,8 +1273,7 @@ impl ExtensionManager { match self.load_mcp_servers(user_id).await { Ok(servers) => { for server in &servers.servers { - let authenticated = - is_authenticated(server, &self.secrets, user_id).await; + let authenticated = is_authenticated(server, &self.secrets, user_id).await; let clients = self.mcp_clients.read().await; let active = clients.contains_key(&server.name); @@ -1663,7 +1665,11 @@ impl ExtensionManager { /// /// The upgrade preserves authentication secrets — only the `.wasm` binary /// (and `.capabilities.json`) are replaced. - pub async fn upgrade(&self, name: Option<&str>, user_id: &str) -> Result { + pub async fn upgrade( + &self, + name: Option<&str>, + user_id: &str, + ) -> Result { // Collect extensions to check let mut candidates: Vec<(String, ExtensionKind)> = Vec::new(); @@ -1858,7 +1864,11 @@ impl ExtensionManager { } /// Get detailed info about an installed extension (version, wit_version, host compatibility). - pub async fn extension_info(&self, name: &str, user_id: &str) -> Result { + pub async fn extension_info( + &self, + name: &str, + user_id: &str, + ) -> Result { Self::validate_extension_name(name)?; let kind = self.determine_installed_kind(name, user_id).await?; @@ -1940,7 +1950,9 @@ impl ExtensionManager { // ── MCP config helpers (DB with disk fallback) ───────────────────── async fn load_mcp_servers( - &self, user_id: &str) -> Result + &self, + user_id: &str, + ) -> Result { if let Some(ref store) = self.store { crate::tools::mcp::config::load_mcp_servers_from_db(store.as_ref(), user_id).await @@ -1951,7 +1963,9 @@ impl ExtensionManager { async fn get_mcp_server( &self, - name: &str, user_id: &str) -> Result { + name: &str, + user_id: &str, + ) -> Result { let servers = self.load_mcp_servers(user_id).await?; servers.get(name).cloned().ok_or_else(|| { crate::tools::mcp::config::ConfigError::ServerNotFound { @@ -1962,11 +1976,12 @@ impl ExtensionManager { async fn add_mcp_server( &self, - config: McpServerConfig, user_id: &str) -> Result<(), crate::tools::mcp::config::ConfigError> { + config: McpServerConfig, + user_id: &str, + ) -> Result<(), crate::tools::mcp::config::ConfigError> { config.validate()?; if let Some(ref store) = self.store { - crate::tools::mcp::config::add_mcp_server_db(store.as_ref(), user_id, config) - .await + crate::tools::mcp::config::add_mcp_server_db(store.as_ref(), user_id, config).await } else { crate::tools::mcp::config::add_mcp_server(config).await } @@ -1974,10 +1989,11 @@ impl ExtensionManager { async fn remove_mcp_server( &self, - name: &str, user_id: &str) -> Result<(), crate::tools::mcp::config::ConfigError> { + name: &str, + user_id: &str, + ) -> Result<(), crate::tools::mcp::config::ConfigError> { if let Some(ref store) = self.store { - crate::tools::mcp::config::remove_mcp_server_db(store.as_ref(), user_id, name) - .await + crate::tools::mcp::config::remove_mcp_server_db(store.as_ref(), user_id, name).await } else { crate::tools::mcp::config::remove_mcp_server(name).await } @@ -1990,7 +2006,9 @@ impl ExtensionManager { entry: &RegistryEntry, user_id: &str, ) -> Result { - let primary_result = self.try_install_from_source(entry, &entry.source, user_id).await; + let primary_result = self + .try_install_from_source(entry, &entry.source, user_id) + .await; match fallback_decision(&primary_result, &entry.fallback_source) { FallbackDecision::Return => primary_result, FallbackDecision::TryFallback => { @@ -2573,7 +2591,9 @@ impl ExtensionManager { async fn auth_mcp_build_url( &self, name: &str, - server: &McpServerConfig, user_id: &str) -> Result { + server: &McpServerConfig, + user_id: &str, + ) -> Result { // Try to discover OAuth metadata and build a URL the user can open manually let metadata = discover_full_oauth_metadata(&server.url) .await @@ -2704,7 +2724,11 @@ impl ExtensionManager { } } - async fn auth_wasm_tool(&self, name: &str, user_id: &str) -> Result { + async fn auth_wasm_tool( + &self, + name: &str, + user_id: &str, + ) -> Result { // Read the capabilities file to get auth config let cap_path = self .wasm_tools_dir @@ -2756,7 +2780,9 @@ impl ExtensionManager { let merged = self .collect_shared_scopes(&auth.secret_name, &oauth.scopes, user_id) .await; - let needs = self.needs_scope_expansion(&auth.secret_name, &merged, user_id).await; + let needs = self + .needs_scope_expansion(&auth.secret_name, &merged, user_id) + .await; tracing::debug!( tool = name, secret_name = %auth.secret_name, @@ -2779,7 +2805,10 @@ impl ExtensionManager { // But only if credentials are available — if the tool has setup secrets // for client_id/secret that aren't configured yet, return needs_setup. if let Some(ref oauth) = auth.oauth { - if self.needs_setup_credentials(name, &auth, oauth, user_id).await { + if self + .needs_setup_credentials(name, &auth, oauth, user_id) + .await + { let display = auth.display_name.as_deref().unwrap_or(name); return Ok(AuthResult::needs_setup( name, @@ -2873,7 +2902,9 @@ impl ExtensionManager { async fn collect_shared_scopes( &self, secret_name: &str, - base_scopes: &[String], _user_id: &str) -> Vec { + base_scopes: &[String], + _user_id: &str, + ) -> Vec { let mut all_scopes: std::collections::BTreeSet = base_scopes.iter().cloned().collect(); @@ -2893,7 +2924,12 @@ impl ExtensionManager { } /// Check whether the stored scopes are insufficient for the merged scopes. - async fn needs_scope_expansion(&self, secret_name: &str, merged_scopes: &[String], user_id: &str) -> bool { + async fn needs_scope_expansion( + &self, + secret_name: &str, + merged_scopes: &[String], + user_id: &str, + ) -> bool { if merged_scopes.is_empty() { return false; } @@ -2967,7 +3003,9 @@ impl ExtensionManager { &self, name: &str, auth: &crate::tools::wasm::AuthCapabilitySchema, - oauth: &crate::tools::wasm::OAuthConfigSchema, user_id: &str) -> bool { + oauth: &crate::tools::wasm::OAuthConfigSchema, + user_id: &str, + ) -> bool { let builtin = crate::cli::oauth_defaults::builtin_credentials(&auth.secret_name); let (id_entry, secret_entry) = self.find_setup_credential_names(name).await; @@ -3011,7 +3049,9 @@ impl ExtensionManager { inline_value: &Option, env_var_name: &Option, builtin_value: Option<&str>, - setup_secret_name: Option<&str>, user_id: &str) -> Option { + setup_secret_name: Option<&str>, + user_id: &str, + ) -> Option { // 1. Check secrets store (entered via Setup tab) if let Some(secret_name) = setup_secret_name && let Ok(secret) = self.secrets.get_decrypted(user_id, secret_name).await @@ -3047,7 +3087,9 @@ impl ExtensionManager { &self, name: &str, auth: &crate::tools::wasm::AuthCapabilitySchema, - oauth: &crate::tools::wasm::OAuthConfigSchema, user_id: &str) -> Result { + oauth: &crate::tools::wasm::OAuthConfigSchema, + user_id: &str, + ) -> Result { use crate::cli::oauth_defaults; let builtin = oauth_defaults::builtin_credentials(&auth.secret_name); @@ -3066,7 +3108,9 @@ impl ExtensionManager { &oauth.client_id, &oauth.client_id_env, builtin.as_ref().map(|c| c.client_id), - setup_client_id_name.as_deref(), user_id) + setup_client_id_name.as_deref(), + user_id, + ) .await .ok_or_else(|| { let env_name = oauth @@ -3093,7 +3137,9 @@ impl ExtensionManager { &oauth.client_secret, &oauth.client_secret_env, builtin.as_ref().map(|c| c.client_secret), - setup_client_secret_name.as_deref(), user_id) + setup_client_secret_name.as_deref(), + user_id, + ) .await; self.clear_pending_extension_auth(name).await; @@ -3428,7 +3474,11 @@ impl ExtensionManager { } /// Check auth status for a WASM channel (read-only). - async fn auth_wasm_channel_status(&self, name: &str, user_id: &str) -> Result { + async fn auth_wasm_channel_status( + &self, + name: &str, + user_id: &str, + ) -> Result { let cap_path = self .wasm_channels_dir .join(format!("{}.capabilities.json", name)); @@ -3485,7 +3535,11 @@ impl ExtensionManager { )) } - async fn activate_mcp(&self, name: &str, user_id: &str) -> Result { + async fn activate_mcp( + &self, + name: &str, + user_id: &str, + ) -> Result { // Check if already activated { let clients = self.mcp_clients.read().await; @@ -3576,7 +3630,11 @@ impl ExtensionManager { }) } - async fn activate_wasm_tool(&self, name: &str, user_id: &str) -> Result { + async fn activate_wasm_tool( + &self, + name: &str, + user_id: &str, + ) -> Result { // Check if already active if self.tool_registry.has(name).await { return Ok(ActivateResult { @@ -3670,7 +3728,11 @@ impl ExtensionManager { /// Loads the channel from its WASM file, injects credentials and config, /// registers it with the webhook router, and hot-adds it to the channel manager /// so its stream feeds into the agent loop. - async fn activate_wasm_channel(&self, name: &str, user_id: &str) -> Result { + async fn activate_wasm_channel( + &self, + name: &str, + user_id: &str, + ) -> Result { // If already active, re-inject credentials and refresh webhook secret. // Handles the case where a channel was loaded at startup before the // user saved secrets via the web UI. @@ -3930,7 +3992,11 @@ impl ExtensionManager { /// /// Called when the user saves new secrets via the setup form for a channel /// that was loaded at startup (possibly without credentials). - async fn refresh_active_channel(&self, name: &str, user_id: &str) -> Result { + async fn refresh_active_channel( + &self, + name: &str, + user_id: &str, + ) -> Result { let router = { let rt_guard = self.channel_runtime.read().await; match rt_guard.as_ref() { @@ -4028,10 +4094,7 @@ impl ExtensionManager { // Refresh signature key if let Some(ref sig_key_name) = sig_key_secret_name - && let Ok(key_secret) = self - .secrets - .get_decrypted(user_id, sig_key_name) - .await + && let Ok(key_secret) = self.secrets.get_decrypted(user_id, sig_key_name).await { match router .register_signature_key(name, key_secret.expose()) @@ -4119,7 +4182,11 @@ impl ExtensionManager { /// For Slack: initiates OAuth flow (redirect-based). /// For Telegram: accepts a bot token, registers it with channel-relay, /// and stores the returned stream token. - async fn auth_channel_relay(&self, name: &str, user_id: &str) -> Result { + async fn auth_channel_relay( + &self, + name: &str, + user_id: &str, + ) -> Result { // Check if already authenticated (stream token exists) if self.is_relay_channel(name, user_id).await { return Ok(AuthResult::authenticated(name, ExtensionKind::ChannelRelay)); @@ -4143,10 +4210,7 @@ impl ExtensionManager { // Delete any stale nonce before storing the new one let _ = self.secrets.delete(user_id, &state_key).await; self.secrets - .create( - user_id, - CreateSecretParams::new(&state_key, &state_nonce), - ) + .create(user_id, CreateSecretParams::new(&state_key, &state_nonce)) .await .map_err(|e| ExtensionError::AuthFailed(format!("Failed to store OAuth state: {e}")))?; @@ -4164,7 +4228,11 @@ impl ExtensionManager { } /// Activate a channel-relay extension. - async fn activate_channel_relay(&self, name: &str, user_id: &str) -> Result { + async fn activate_channel_relay( + &self, + name: &str, + user_id: &str, + ) -> Result { let token_key = format!("relay:{}:stream_token", name); let team_id_key = format!("relay:{}:team_id", name); @@ -4259,7 +4327,11 @@ impl ExtensionManager { } /// Activate a channel-relay extension from stored credentials (for startup reconnect). - pub async fn activate_stored_relay(&self, name: &str, user_id: &str) -> Result<(), ExtensionError> { + pub async fn activate_stored_relay( + &self, + name: &str, + user_id: &str, + ) -> Result<(), ExtensionError> { self.installed_relay_extensions .write() .await @@ -4273,7 +4345,11 @@ impl ExtensionManager { /// This is a read-only check — it never modifies `installed_relay_extensions`. /// To mark a relay extension as installed, use `activate_stored_relay()` or /// the explicit install flow. - async fn determine_installed_kind(&self, name: &str, user_id: &str) -> Result { + async fn determine_installed_kind( + &self, + name: &str, + user_id: &str, + ) -> Result { // Check MCP servers first if self.get_mcp_server(name, user_id).await.is_ok() { return Ok(ExtensionKind::McpServer); @@ -5137,20 +5213,14 @@ impl ExtensionManager { && let Some(ref auth_cfg) = cap.auth && auth_cfg.oauth.is_some() { - let _ = self - .secrets - .delete(user_id, &auth_cfg.secret_name) - .await; + let _ = self.secrets.delete(user_id, &auth_cfg.secret_name).await; let _ = self .secrets .delete(user_id, &format!("{}_scopes", auth_cfg.secret_name)) .await; let _ = self .secrets - .delete( - user_id, - &format!("{}_refresh_token", auth_cfg.secret_name), - ) + .delete(user_id, &format!("{}_refresh_token", auth_cfg.secret_name)) .await; } @@ -5277,7 +5347,9 @@ impl ExtensionManager { pub async fn configure_token( &self, name: &str, - token: &str, user_id: &str) -> Result { + token: &str, + user_id: &str, + ) -> Result { let kind = self.determine_installed_kind(name, user_id).await?; let secret_name = match kind { ExtensionKind::WasmChannel => { @@ -5297,12 +5369,7 @@ impl ExtensionManager { if s.optional { continue; } - if !self - .secrets - .exists(user_id, &s.name) - .await - .unwrap_or(false) - { + if !self.secrets.exists(user_id, &s.name).await.unwrap_or(false) { target = Some(s.name.clone()); break; } @@ -5338,12 +5405,7 @@ impl ExtensionManager { // Auth secret exists, find first missing setup secret let mut found = None; for s in &setup.required_secrets { - if !self - .secrets - .exists(user_id, &s.name) - .await - .unwrap_or(false) - { + if !self.secrets.exists(user_id, &s.name).await.unwrap_or(false) { found = Some(s.name.clone()); break; } @@ -5938,7 +6000,7 @@ mod tests { wasm_runtime, tools_dir, channels_dir, - None, // tunnel_url + None, // tunnel_url "test".to_string(), // user_id store, vec![], @@ -6057,7 +6119,12 @@ mod tests { fields.insert("llm_backend".to_string(), "openai".to_string()); let result = mgr - .configure("switch-llm", &std::collections::HashMap::new(), &fields, "test-user") + .configure( + "switch-llm", + &std::collections::HashMap::new(), + &fields, + "test-user", + ) .await .expect("save configuration"); @@ -6105,7 +6172,12 @@ mod tests { fields.insert("session".to_string(), "overwrite".to_string()); let err = match mgr - .configure("evil-tool", &std::collections::HashMap::new(), &fields, "test-user") + .configure( + "evil-tool", + &std::collections::HashMap::new(), + &fields, + "test-user", + ) .await { Ok(_) => panic!("disallowed setting_path should fail"), @@ -6255,7 +6327,10 @@ mod tests { let manager = make_manager_custom_dirs(dir.path().join("tools"), channels_dir); - let result = manager.upgrade(Some("custom-channel"), "test").await.unwrap(); + let result = manager + .upgrade(Some("custom-channel"), "test") + .await + .unwrap(); assert_eq!(result.results.len(), 1); assert_eq!(result.results[0].status, "not_in_registry"); } @@ -6977,7 +7052,10 @@ mod tests { let dir = tempfile::tempdir().expect("temp dir"); let mgr = make_test_manager(None, dir.path().to_path_buf()); - let err = mgr.activate_channel_relay("slack-relay", "test").await.unwrap_err(); + let err = mgr + .activate_channel_relay("slack-relay", "test") + .await + .unwrap_err(); assert!( matches!(err, ExtensionError::AuthRequired), "expected AuthRequired, got: {err:?}" @@ -7794,7 +7872,12 @@ mod tests { ); let result = mgr - .configure("test-relay", &secrets, &std::collections::HashMap::new(), "test") + .configure( + "test-relay", + &secrets, + &std::collections::HashMap::new(), + "test", + ) .await; assert!( result.is_ok(), diff --git a/src/main.rs b/src/main.rs index 89f648fa..da6ab402 100644 --- a/src/main.rs +++ b/src/main.rs @@ -820,7 +820,9 @@ async fn async_main() -> anyhow::Result<()> { .unwrap_or_else(|| "default".to_string()); let persisted = ext_mgr.load_persisted_active_channels(&ext_user_id).await; for name in &persisted { - if active_at_startup.contains(name) || ext_mgr.is_relay_channel(name, &ext_user_id).await { + if active_at_startup.contains(name) + || ext_mgr.is_relay_channel(name, &ext_user_id).await + { continue; } match ext_mgr.activate(name, &ext_user_id).await { diff --git a/src/tools/builtin/memory.rs b/src/tools/builtin/memory.rs index 5dd777c9..29f7480a 100644 --- a/src/tools/builtin/memory.rs +++ b/src/tools/builtin/memory.rs @@ -668,8 +668,13 @@ impl MemoryTreeTool { }; if entry.is_directory && current_depth < max_depth { - let children = - Box::pin(Self::build_tree(workspace, &entry.path, current_depth + 1, max_depth)).await?; + let children = Box::pin(Self::build_tree( + workspace, + &entry.path, + current_depth + 1, + max_depth, + )) + .await?; if children.is_empty() { result.push(serde_json::Value::String(display_path)); } else { diff --git a/src/tools/registry.rs b/src/tools/registry.rs index dd72f33b..bc3be144 100644 --- a/src/tools/registry.rs +++ b/src/tools/registry.rs @@ -355,9 +355,15 @@ impl ToolRegistry { /// Memory tools require a workspace for persistence. Call this after /// `register_builtin_tools()` if you have a workspace available. pub fn register_memory_tools(&self, workspace: Arc) { - self.register_sync(Arc::new(MemorySearchTool::from_workspace(Arc::clone(&workspace)))); - self.register_sync(Arc::new(MemoryWriteTool::from_workspace(Arc::clone(&workspace)))); - self.register_sync(Arc::new(MemoryReadTool::from_workspace(Arc::clone(&workspace)))); + self.register_sync(Arc::new(MemorySearchTool::from_workspace(Arc::clone( + &workspace, + )))); + self.register_sync(Arc::new(MemoryWriteTool::from_workspace(Arc::clone( + &workspace, + )))); + self.register_sync(Arc::new(MemoryReadTool::from_workspace(Arc::clone( + &workspace, + )))); self.register_sync(Arc::new(MemoryTreeTool::from_workspace(workspace))); tracing::debug!("Registered 4 memory tools"); @@ -377,7 +383,11 @@ impl ToolRegistry { job_manager: Option>, store: Option>, job_event_tx: Option< - tokio::sync::broadcast::Sender<(uuid::Uuid, String, crate::channels::web::types::SseEvent)>, + tokio::sync::broadcast::Sender<( + uuid::Uuid, + String, + crate::channels::web::types::SseEvent, + )>, >, inject_tx: Option>, prompt_queue: Option, diff --git a/tests/multi_tenant_integration.rs b/tests/multi_tenant_integration.rs index 6924a1a7..a1493c4f 100644 --- a/tests/multi_tenant_integration.rs +++ b/tests/multi_tenant_integration.rs @@ -12,11 +12,11 @@ use std::net::SocketAddr; use std::sync::Arc; use std::time::Duration; +use axum::Router; use axum::body::Body; use axum::http::{Request, StatusCode}; use axum::middleware; use axum::routing::{get, post}; -use axum::Router; use tower::ServiceExt; use ironclaw::channels::web::auth::{ @@ -60,15 +60,11 @@ fn two_user_auth() -> MultiAuthState { /// Build a test Router that echoes the authenticated user_id back. fn user_echo_app(auth: MultiAuthState) -> Router { - async fn echo_user( - AuthenticatedUser(user): AuthenticatedUser, - ) -> String { + async fn echo_user(AuthenticatedUser(user): AuthenticatedUser) -> String { user.user_id } - async fn echo_user_with_scopes( - AuthenticatedUser(user): AuthenticatedUser, - ) -> String { + async fn echo_user_with_scopes(AuthenticatedUser(user): AuthenticatedUser) -> String { format!("{}:{}", user.user_id, user.workspace_read_scopes.join(",")) } @@ -870,7 +866,10 @@ async fn start_multi_user_server_with_db() -> ( let backend = ironclaw::db::libsql::LibSqlBackend::new_local(&path) .await .expect("failed to create test DB"); - backend.run_migrations().await.expect("failed to run migrations"); + backend + .run_migrations() + .await + .expect("failed to run migrations"); let db: Arc = Arc::new(backend); let (agent_tx, _agent_rx) = tokio::sync::mpsc::channel(64); let auth = two_user_auth(); @@ -1027,7 +1026,11 @@ async fn full_server_bob_sees_own_jobs_only() { let body: serde_json::Value = resp.json().await.unwrap(); let jobs = body["jobs"].as_array().unwrap(); - assert_eq!(jobs.len(), 2, "Bob should see only his 2 jobs, not Alice's 3"); + assert_eq!( + jobs.len(), + 2, + "Bob should see only his 2 jobs, not Alice's 3" + ); for job in jobs { let title = job["title"].as_str().unwrap(); assert!( diff --git a/tests/multi_tenant_system_prompt.rs b/tests/multi_tenant_system_prompt.rs index ce6a5e94..ece794bf 100644 --- a/tests/multi_tenant_system_prompt.rs +++ b/tests/multi_tenant_system_prompt.rs @@ -98,10 +98,7 @@ mod tests { #[tokio::test] async fn alice_system_prompt_contains_alice_identity() { let trace = simple_trace(1); - let rig = TestRigBuilder::new() - .with_trace(trace) - .build() - .await; + let rig = TestRigBuilder::new().with_trace(trace).build().await; // Seed alice's identity into the database let db = rig.database(); @@ -114,8 +111,8 @@ mod tests { // The system prompt sent to the LLM should contain Alice's identity let requests = rig.captured_llm_requests(); - let system_prompt = extract_system_prompt(&requests) - .expect("Expected a system prompt in the LLM request"); + let system_prompt = + extract_system_prompt(&requests).expect("Expected a system prompt in the LLM request"); assert!( system_prompt.contains("Alice is a software engineer"), @@ -134,10 +131,7 @@ mod tests { #[tokio::test] async fn bob_system_prompt_contains_bob_identity() { let trace = simple_trace(1); - let rig = TestRigBuilder::new() - .with_trace(trace) - .build() - .await; + let rig = TestRigBuilder::new().with_trace(trace).build().await; // Seed bob's identity into the database let db = rig.database(); @@ -150,8 +144,8 @@ mod tests { // The system prompt should contain Bob's identity let requests = rig.captured_llm_requests(); - let system_prompt = extract_system_prompt(&requests) - .expect("Expected a system prompt in the LLM request"); + let system_prompt = + extract_system_prompt(&requests).expect("Expected a system prompt in the LLM request"); assert!( system_prompt.contains("Bob is a marine biologist"), @@ -169,10 +163,7 @@ mod tests { #[tokio::test] async fn alice_identity_does_not_leak_into_bob_prompt() { let trace = simple_trace(1); - let rig = TestRigBuilder::new() - .with_trace(trace) - .build() - .await; + let rig = TestRigBuilder::new().with_trace(trace).build().await; // Seed BOTH users' identities let db = rig.database(); @@ -196,8 +187,7 @@ mod tests { ); } // Also verify Bob's identity IS present (compound check) - let prompt = system_prompt - .expect("Expected a system prompt in the LLM request"); + let prompt = system_prompt.expect("Expected a system prompt in the LLM request"); assert!( prompt.contains("Bob is a marine biologist"), "Bob's own identity should be in his system prompt.\n\ @@ -214,10 +204,7 @@ mod tests { #[tokio::test] async fn bob_identity_does_not_leak_into_alice_prompt() { let trace = simple_trace(1); - let rig = TestRigBuilder::new() - .with_trace(trace) - .build() - .await; + let rig = TestRigBuilder::new().with_trace(trace).build().await; // Seed BOTH users' identities let db = rig.database(); @@ -241,8 +228,7 @@ mod tests { ); } // Also verify Alice's identity IS present - let prompt = system_prompt - .expect("Expected a system prompt in the LLM request"); + let prompt = system_prompt.expect("Expected a system prompt in the LLM request"); assert!( prompt.contains("Alice is a software engineer"), "Alice's own identity should be in her system prompt.\n\ diff --git a/tests/openai_compat_integration.rs b/tests/openai_compat_integration.rs index c3892ace..16568246 100644 --- a/tests/openai_compat_integration.rs +++ b/tests/openai_compat_integration.rs @@ -722,9 +722,7 @@ async fn test_no_llm_provider_returns_503() { "test-user".to_string(), ); let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let bound_addr = start_server(addr, state, auth) - .await - .unwrap(); + let bound_addr = start_server(addr, state, auth).await.unwrap(); let url = format!("http://{}/v1/chat/completions", bound_addr); let resp = client() diff --git a/tests/support/gateway_workflow_harness.rs b/tests/support/gateway_workflow_harness.rs index a6353f27..7f9d3dff 100644 --- a/tests/support/gateway_workflow_harness.rs +++ b/tests/support/gateway_workflow_harness.rs @@ -13,9 +13,11 @@ use ironclaw::agent::routine_engine::RoutineEngine; use ironclaw::agent::{Agent, AgentDeps, SessionManager as AgentSessionManager}; use ironclaw::app::{AppBuilder, AppBuilderFlags}; use ironclaw::channels::IncomingMessage; -use ironclaw::channels::web::log_layer::LogBroadcaster; use ironclaw::channels::web::auth::MultiAuthState; -use ironclaw::channels::web::server::{GatewayState, PerUserRateLimiter, RateLimiter, start_server}; +use ironclaw::channels::web::log_layer::LogBroadcaster; +use ironclaw::channels::web::server::{ + GatewayState, PerUserRateLimiter, RateLimiter, start_server, +}; use ironclaw::channels::web::sse::SseManager; use ironclaw::channels::web::ws::WsConnectionTracker; use ironclaw::config::{Config, RegistryProviderConfig, RoutineConfig};