diff --git a/Cargo.lock b/Cargo.lock index 35065644..ba886089 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2686,6 +2686,7 @@ dependencies = [ "tokio", "tracing", "uuid", + "windows-dpapi", "zeroize", ] @@ -6312,6 +6313,17 @@ dependencies = [ "windows-strings 0.5.1", ] +[[package]] +name = "windows-dpapi" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "162a325089267c13a318d5b0356c785e0f548ca5a5584e1b6b7b49ecd163121a" +dependencies = [ + "anyhow", + "log", + "winapi", +] + [[package]] name = "windows-future" version = "0.2.1" diff --git a/Cargo.toml b/Cargo.toml index 49636fca..7131d2bf 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -83,3 +83,4 @@ lto = true codegen-units = 1 strip = true + diff --git a/apps/desktop/src-tauri/src/commands/gateway.rs b/apps/desktop/src-tauri/src/commands/gateway.rs index c75aba4f..fd20f311 100644 --- a/apps/desktop/src-tauri/src/commands/gateway.rs +++ b/apps/desktop/src-tauri/src/commands/gateway.rs @@ -9,7 +9,6 @@ use mcpmux_gateway::{ ConnectionContext, ConnectionResult, FeatureService, InstalledServerInfo, PoolService, ResolvedTransport, ServerKey, }; -use mcpmux_storage::{JwtSecretProvider, KeychainJwtSecretProvider}; use serde::Serialize; use std::sync::Arc; use tauri::{AppHandle, Emitter, State}; @@ -466,11 +465,11 @@ fn create_gateway_dependencies( app_state: &AppState, _app_handle: tauri::AppHandle, ) -> Result { - // Load JWT signing secret from keychain (or create if first run) - let jwt_secret = match KeychainJwtSecretProvider::new() { + // Load JWT signing secret (DPAPI on Windows, keychain elsewhere) + let jwt_secret = match mcpmux_storage::create_jwt_secret_provider(app_state.data_dir()) { Ok(provider) => match provider.get_or_create_secret() { Ok(secret) => { - info!("[Gateway] JWT signing secret loaded from keychain"); + info!("[Gateway] JWT signing secret loaded"); Some(secret) } Err(e) => { @@ -479,7 +478,7 @@ fn create_gateway_dependencies( } }, Err(e) => { - warn!("[Gateway] Failed to create keychain provider: {}", e); + warn!("[Gateway] Failed to create JWT secret provider: {}", e); None } }; @@ -1021,11 +1020,15 @@ pub async fn connect_all_enabled_servers( } }; - // Check if has OAuth credentials + // Check if has OAuth credentials (access token) let has_credentials = matches!( app_state .credential_repository - .get(&space.id, &installed.server_id) + .get( + &space.id, + &installed.server_id, + &mcpmux_core::CredentialType::AccessToken + ) .await, Ok(Some(_)) ); diff --git a/apps/desktop/src-tauri/src/commands/server_manager.rs b/apps/desktop/src-tauri/src/commands/server_manager.rs index 19146c29..58801b5d 100644 --- a/apps/desktop/src-tauri/src/commands/server_manager.rs +++ b/apps/desktop/src-tauri/src/commands/server_manager.rs @@ -478,10 +478,10 @@ pub async fn logout_server( .oauth_manager() .cancel_flow_for_space(space_uuid, &server_id); - // 3. Clear OAuth tokens from credential repository + // 3. Clear all credentials for this server if let Err(e) = app_state .credential_repository - .delete(&space_uuid, &server_id) + .delete_all(&space_uuid, &server_id) .await { warn!( @@ -489,7 +489,7 @@ pub async fn logout_server( server_id, e ); } else { - info!("[ServerManager] Cleared OAuth tokens for {}", server_id); + info!("[ServerManager] Cleared credentials for {}", server_id); } // 4. Clear oauth_connected flag diff --git a/apps/desktop/src-tauri/src/lib.rs b/apps/desktop/src-tauri/src/lib.rs index 3e498f1e..21abbeb4 100644 --- a/apps/desktop/src-tauri/src/lib.rs +++ b/apps/desktop/src-tauri/src/lib.rs @@ -300,12 +300,11 @@ pub fn run() { let url = format!("http://localhost:{}", final_port); info!("Auto-starting gateway on {}", url); - // Load JWT signing secret from keychain (or create if first run) - use mcpmux_storage::{JwtSecretProvider, KeychainJwtSecretProvider}; - let jwt_secret = match KeychainJwtSecretProvider::new() { + // Load JWT signing secret (DPAPI on Windows, keychain elsewhere) + let jwt_secret = match mcpmux_storage::create_jwt_secret_provider(&app_data_dir) { Ok(provider) => match provider.get_or_create_secret() { Ok(secret) => { - info!("[Gateway] JWT signing secret loaded from keychain"); + info!("[Gateway] JWT signing secret loaded"); Some(secret) } Err(e) => { @@ -314,7 +313,7 @@ pub fn run() { } }, Err(e) => { - warn!("[Gateway] Failed to create keychain provider: {}. Token signing disabled.", e); + warn!("[Gateway] Failed to create JWT secret provider: {}. Token signing disabled.", e); None } }; diff --git a/apps/desktop/src-tauri/src/state/mod.rs b/apps/desktop/src-tauri/src/state/mod.rs index ae5599a9..f4dae3aa 100644 --- a/apps/desktop/src-tauri/src/state/mod.rs +++ b/apps/desktop/src-tauri/src/state/mod.rs @@ -11,10 +11,9 @@ use mcpmux_core::{ SpaceService, }; use mcpmux_storage::{ - Database, FieldEncryptor, KeychainKeyProvider, MasterKeyProvider, SqliteAppSettingsRepository, - SqliteCredentialRepository, SqliteFeatureSetRepository, SqliteInboundMcpClientRepository, - SqliteInstalledServerRepository, SqliteOutboundOAuthRepository, SqliteServerFeatureRepository, - SqliteSpaceRepository, + Database, FieldEncryptor, SqliteAppSettingsRepository, SqliteCredentialRepository, + SqliteFeatureSetRepository, SqliteInboundMcpClientRepository, SqliteInstalledServerRepository, + SqliteOutboundOAuthRepository, SqliteServerFeatureRepository, SqliteSpaceRepository, }; use std::path::PathBuf; use std::sync::Arc; @@ -67,9 +66,9 @@ impl AppState { // Ensure data directory exists std::fs::create_dir_all(&data_dir)?; - // Get or create master key from OS keychain - info!("Retrieving master key from keychain..."); - let key_provider = KeychainKeyProvider::new()?; + // Get or create master key (DPAPI on Windows, OS Keychain elsewhere) + info!("Retrieving master key..."); + let key_provider = mcpmux_storage::create_key_provider(&data_dir)?; let master_key = key_provider.get_or_create_key()?; info!("Master key retrieved successfully"); @@ -87,8 +86,9 @@ impl AppState { let space_repository: Arc = Arc::new(SqliteSpaceRepository::new(db.clone())); - let installed_server_repository: Arc = - Arc::new(SqliteInstalledServerRepository::new(db.clone())); + let installed_server_repository: Arc = Arc::new( + SqliteInstalledServerRepository::new(db.clone(), encryptor.clone()), + ); let credential_repository: Arc = Arc::new( SqliteCredentialRepository::new(db.clone(), encryptor.clone()), diff --git a/crates/mcpmux-core/src/application/server.rs b/crates/mcpmux-core/src/application/server.rs index 62c9891e..a9a5d580 100644 --- a/crates/mcpmux-core/src/application/server.rs +++ b/crates/mcpmux-core/src/application/server.rs @@ -175,9 +175,9 @@ impl ServerAppService { } } - // Delete credentials + // Delete all credentials for this server if let Some(ref cred_repo) = self.credential_repo { - if let Err(e) = cred_repo.delete(&space_id, server_id).await { + if let Err(e) = cred_repo.delete_all(&space_id, server_id).await { warn!( server_id = server_id, error = %e, diff --git a/crates/mcpmux-core/src/domain/credential.rs b/crates/mcpmux-core/src/domain/credential.rs index 4a98b344..e19bda6c 100644 --- a/crates/mcpmux-core/src/domain/credential.rs +++ b/crates/mcpmux-core/src/domain/credential.rs @@ -1,45 +1,100 @@ //! Credential entity - secure credential storage //! +//! Each credential is a typed entry: one row per (space, server, type). +//! This allows separate lifecycle management for access tokens vs refresh tokens, +//! and keeps metadata (expiry, scope) as plaintext while only encrypting the secret value. +//! //! Note: OAuth client registration (client_id, endpoints) is stored separately //! in the `oauth_clients` table via OAuthClient entity. use chrono::{DateTime, Utc}; use serde::{Deserialize, Serialize}; +use std::fmt; use uuid::Uuid; -/// Credential type - stores tokens/keys, NOT client registration -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] -#[serde(tag = "type", rename_all = "snake_case")] -pub enum CredentialValue { - /// Simple API key - ApiKey { key: String }, - - /// OAuth tokens (client registration is in oauth_clients table) - OAuth { - access_token: String, - refresh_token: Option, - expires_at: Option>, - token_type: String, - scope: Option, - }, +/// Type of credential entry. +/// +/// Extensible: add new variants for session tokens, client certificates, etc. +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum CredentialType { + /// OAuth access token (~1h lifetime) + AccessToken, + /// OAuth refresh token (~90d lifetime) + RefreshToken, + /// Simple API key (no expiry) + ApiKey, + /// Basic auth username + BasicAuthUser, + /// Basic auth password + BasicAuthPass, +} + +impl CredentialType { + /// Convert to database string representation. + pub fn as_str(&self) -> &'static str { + match self { + Self::AccessToken => "access_token", + Self::RefreshToken => "refresh_token", + Self::ApiKey => "api_key", + Self::BasicAuthUser => "basic_auth_user", + Self::BasicAuthPass => "basic_auth_pass", + } + } + + /// Parse from database string representation. + pub fn parse(s: &str) -> Option { + match s { + "access_token" => Some(Self::AccessToken), + "refresh_token" => Some(Self::RefreshToken), + "api_key" => Some(Self::ApiKey), + "basic_auth_user" => Some(Self::BasicAuthUser), + "basic_auth_pass" => Some(Self::BasicAuthPass), + _ => None, + } + } + + /// Whether this type represents an OAuth token (access or refresh). + pub fn is_oauth(&self) -> bool { + matches!(self, Self::AccessToken | Self::RefreshToken) + } +} - /// Basic authentication - BasicAuth { username: String, password: String }, +impl fmt::Display for CredentialType { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.as_str()) + } } -/// Credential for a specific (Space, Server) combination. +/// Individual credential entry — one per (space, server, type). /// -/// Credentials are stored locally only, never synced to cloud. -#[derive(Debug, Clone, Serialize, Deserialize)] +/// The `value` field contains the secret (token, key, password) in plaintext +/// at the domain level. Encryption is handled by the storage layer. +/// +/// Metadata fields (expires_at, token_type, scope) are non-sensitive and +/// stored as plaintext in the database for queryability. +#[derive(Debug, Clone)] pub struct Credential { - /// Space ID + /// Space this credential belongs to pub space_id: Uuid, - /// Server ID + /// Server this credential is for pub server_id: String, - /// The credential value - pub value: CredentialValue, + /// Type of credential (access_token, refresh_token, api_key, etc.) + pub credential_type: CredentialType, + + /// The secret value (plaintext at domain level, encrypted at storage level) + pub value: String, + + /// When this credential expires (plaintext in DB for queryability) + pub expires_at: Option>, + + /// Token type, e.g. "Bearer" (only for access_token) + pub token_type: Option, + + /// OAuth scope (only for access_token) + pub scope: Option, /// When the credential was created pub created_at: DateTime, @@ -52,88 +107,78 @@ pub struct Credential { } impl Credential { - /// Create a new API key credential + /// Create a new API key credential. pub fn api_key(space_id: Uuid, server_id: impl Into, key: impl Into) -> Self { let now = Utc::now(); Self { space_id, server_id: server_id.into(), - value: CredentialValue::ApiKey { key: key.into() }, + credential_type: CredentialType::ApiKey, + value: key.into(), + expires_at: None, + token_type: None, + scope: None, created_at: now, updated_at: now, last_used: None, } } - /// Create an OAuth credential (tokens only, client registration is separate) - pub fn oauth( + /// Create an OAuth access token credential. + pub fn access_token( space_id: Uuid, server_id: impl Into, - access_token: impl Into, - refresh_token: Option, + token: impl Into, expires_at: Option>, ) -> Self { let now = Utc::now(); Self { space_id, server_id: server_id.into(), - value: CredentialValue::OAuth { - access_token: access_token.into(), - refresh_token, - expires_at, - token_type: "Bearer".to_string(), - scope: None, - }, + credential_type: CredentialType::AccessToken, + value: token.into(), + expires_at, + token_type: Some("Bearer".to_string()), + scope: None, created_at: now, updated_at: now, last_used: None, } } - /// Get the credential key for storage lookup - pub fn key(&self) -> String { - format!("{}:{}", self.space_id, self.server_id) - } - - /// Check if this credential is expired (for OAuth) - pub fn is_expired(&self) -> bool { - match &self.value { - CredentialValue::OAuth { - expires_at: Some(exp), - .. - } => *exp < Utc::now(), - _ => false, + /// Create an OAuth refresh token credential. + pub fn refresh_token( + space_id: Uuid, + server_id: impl Into, + token: impl Into, + expires_at: Option>, + ) -> Self { + let now = Utc::now(); + Self { + space_id, + server_id: server_id.into(), + credential_type: CredentialType::RefreshToken, + value: token.into(), + expires_at, + token_type: None, + scope: None, + created_at: now, + updated_at: now, + last_used: None, } } - /// Check if this credential can be refreshed - pub fn can_refresh(&self) -> bool { - matches!( - &self.value, - CredentialValue::OAuth { - refresh_token: Some(_), - .. - } - ) - } - - /// Update the access token (after refresh) - pub fn update_token(&mut self, access_token: String, expires_at: Option>) { - if let CredentialValue::OAuth { - access_token: ref mut at, - expires_at: ref mut exp, - .. - } = self.value - { - *at = access_token; - *exp = expires_at; - self.updated_at = Utc::now(); + /// Check if this credential is expired. + pub fn is_expired(&self) -> bool { + match self.expires_at { + Some(exp) => exp < Utc::now(), + None => false, } } - /// Check if this is an OAuth credential + /// Check if this is an OAuth credential (access or refresh token). pub fn is_oauth(&self) -> bool { - matches!(self.value, CredentialValue::OAuth { .. }) + self.credential_type.is_oauth() } } @@ -147,38 +192,63 @@ mod tests { let cred = Credential::api_key(space_id, "github", "ghp_xxx"); assert_eq!(cred.server_id, "github"); - assert!(matches!(cred.value, CredentialValue::ApiKey { .. })); + assert_eq!(cred.credential_type, CredentialType::ApiKey); + assert_eq!(cred.value, "ghp_xxx"); assert!(!cred.is_expired()); - assert!(!cred.can_refresh()); + assert!(!cred.is_oauth()); } #[test] - fn test_oauth_credential() { + fn test_access_token_credential() { let space_id = Uuid::new_v4(); - let cred = Credential::oauth( + let cred = Credential::access_token( space_id, "atlassian", - "access_token", - Some("refresh_token".to_string()), + "access_token_xyz", Some(Utc::now() + chrono::Duration::hours(1)), ); + assert_eq!(cred.credential_type, CredentialType::AccessToken); assert!(!cred.is_expired()); - assert!(cred.can_refresh()); + assert!(cred.is_oauth()); + assert_eq!(cred.token_type, Some("Bearer".to_string())); + } + + #[test] + fn test_refresh_token_credential() { + let space_id = Uuid::new_v4(); + let cred = Credential::refresh_token(space_id, "atlassian", "refresh_xyz", None); + + assert_eq!(cred.credential_type, CredentialType::RefreshToken); + assert!(!cred.is_expired()); // No expiry set + assert!(cred.is_oauth()); } #[test] fn test_expired_credential() { let space_id = Uuid::new_v4(); - let cred = Credential::oauth( + let cred = Credential::access_token( space_id, "atlassian", "access_token", - None, Some(Utc::now() - chrono::Duration::hours(1)), ); assert!(cred.is_expired()); - assert!(!cred.can_refresh()); + } + + #[test] + fn test_credential_type_roundtrip() { + for ct in [ + CredentialType::AccessToken, + CredentialType::RefreshToken, + CredentialType::ApiKey, + CredentialType::BasicAuthUser, + CredentialType::BasicAuthPass, + ] { + let s = ct.as_str(); + let parsed = CredentialType::parse(s).unwrap(); + assert_eq!(ct, parsed); + } } } diff --git a/crates/mcpmux-core/src/repository/mod.rs b/crates/mcpmux-core/src/repository/mod.rs index b49d97e1..95b1d838 100644 --- a/crates/mcpmux-core/src/repository/mod.rs +++ b/crates/mcpmux-core/src/repository/mod.rs @@ -7,7 +7,7 @@ use async_trait::async_trait; use uuid::Uuid; use crate::domain::{ - Client, Credential, FeatureSet, FeatureSetMember, InstalledServer, MemberMode, + Client, Credential, CredentialType, FeatureSet, FeatureSetMember, InstalledServer, MemberMode, OutboundOAuthRegistration, ServerFeature, Space, }; @@ -269,19 +269,38 @@ pub trait InboundMcpClientRepository: Send + Sync { } /// Credential repository trait (local-only, never synced) +/// +/// Each credential is a separate row per (space, server, type). +/// This allows independent lifecycle management for access tokens vs refresh tokens. #[async_trait] pub trait CredentialRepository: Send + Sync { - /// Get a credential for a (space, server) combination - async fn get(&self, space_id: &Uuid, server_id: &str) -> RepoResult>; + /// Get a specific credential by (space, server, type) + async fn get( + &self, + space_id: &Uuid, + server_id: &str, + credential_type: &CredentialType, + ) -> RepoResult>; + + /// Get all credentials for a (space, server) combination + async fn get_all(&self, space_id: &Uuid, server_id: &str) -> RepoResult>; - /// Save a credential + /// Save a credential (upsert by space_id + server_id + credential_type) async fn save(&self, credential: &Credential) -> RepoResult<()>; - /// Delete a credential completely - async fn delete(&self, space_id: &Uuid, server_id: &str) -> RepoResult<()>; + /// Delete a specific credential by type + async fn delete( + &self, + space_id: &Uuid, + server_id: &str, + credential_type: &CredentialType, + ) -> RepoResult<()>; + + /// Delete all credentials for a (space, server) combination + async fn delete_all(&self, space_id: &Uuid, server_id: &str) -> RepoResult<()>; - /// Clear OAuth tokens but preserve client registration (for logout) - /// Returns true if tokens were cleared, false if credential not found or not OAuth + /// Clear OAuth tokens (access + refresh) but preserve client registration (for logout) + /// Returns true if tokens were cleared async fn clear_tokens(&self, space_id: &Uuid, server_id: &str) -> RepoResult; /// List all credentials for a space diff --git a/crates/mcpmux-gateway/src/pool/credential_store.rs b/crates/mcpmux-gateway/src/pool/credential_store.rs index 7b18de35..88496d01 100644 --- a/crates/mcpmux-gateway/src/pool/credential_store.rs +++ b/crates/mcpmux-gateway/src/pool/credential_store.rs @@ -1,14 +1,15 @@ //! Database-backed CredentialStore adapter for rmcp SDK integration. //! -//! Bridges our split storage (OutboundOAuthRepository + CredentialRepository) -//! to rmcp's unified CredentialStore interface. +//! Bridges our typed credential rows (CredentialRepository) and +//! client registrations (OutboundOAuthRepository) to rmcp's unified +//! CredentialStore interface. use std::sync::Arc; use async_trait::async_trait; use chrono::{Duration, Utc}; use mcpmux_core::{ - Credential, CredentialRepository, CredentialValue, OutboundOAuthRegistration, + Credential, CredentialRepository, CredentialType, OutboundOAuthRegistration, OutboundOAuthRepository, }; use oauth2::{basic::BasicTokenType, AccessToken, RefreshToken, TokenResponse}; @@ -49,72 +50,70 @@ impl DatabaseCredentialStore { } } - /// Convert our Credential to SDK's OAuthTokenResponse - fn to_token_response(credential: &Credential) -> Option { - match &credential.value { - CredentialValue::OAuth { - access_token, - refresh_token, - expires_at, - .. - } => { - // Build a minimal token response - // The SDK uses oauth2 crate's StandardTokenResponse internally - let expires_in = expires_at.map(|exp| { - let duration = exp - Utc::now(); - std::time::Duration::from_secs(duration.num_seconds().max(0) as u64) - }); - - Some(build_token_response( - access_token.clone(), - refresh_token.clone(), - expires_in, - )) - } - _ => None, - } + /// Build an OAuthTokenResponse from separate access_token and refresh_token credentials. + fn build_token_response( + access_cred: &Credential, + refresh_cred: Option<&Credential>, + ) -> OAuthTokenResponse { + // Recalculate expires_in from stored expires_at + let expires_in = access_cred.expires_at.map(|exp| { + let duration = exp - Utc::now(); + std::time::Duration::from_secs(duration.num_seconds().max(0) as u64) + }); + + build_token_response( + access_cred.value.clone(), + refresh_cred.map(|r| r.value.clone()), + expires_in, + ) } - /// Convert SDK's StoredCredentials to our storage format + /// Save SDK's StoredCredentials to our typed credential rows. async fn save_to_database(&self, creds: &StoredCredentials) -> Result<(), AuthError> { - // Save token to credentials table + // Save tokens as separate rows if let Some(token_response) = &creds.token_response { - let access_token = token_response.access_token().secret().to_string(); - let refresh_token = token_response - .refresh_token() - .map(|t| t.secret().to_string()); + let access_token_str = token_response.access_token().secret().to_string(); let expires_at = token_response .expires_in() .map(|d| Utc::now() + Duration::seconds(d.as_secs() as i64)); - let credential = Credential { - space_id: self.space_id, - server_id: self.server_id.clone(), - value: CredentialValue::OAuth { - access_token, - refresh_token, - expires_at, - token_type: "Bearer".to_string(), - scope: None, - }, - created_at: Utc::now(), - updated_at: Utc::now(), - last_used: Some(Utc::now()), - }; - - self.credential_repo - .save(&credential) - .await - .map_err(|e| AuthError::InternalError(format!("Failed to save token: {}", e)))?; + // Save access_token row + let access_cred = Credential::access_token( + self.space_id, + &self.server_id, + access_token_str, + expires_at, + ); + self.credential_repo.save(&access_cred).await.map_err(|e| { + AuthError::InternalError(format!("Failed to save access token: {}", e)) + })?; + + // Save refresh_token row (if present in response). + // If the response doesn't include a refresh_token, preserve the existing one + // in the database — some providers (e.g. Atlassian) omit it during token rotation. + if let Some(refresh_token) = token_response.refresh_token() { + let refresh_cred = Credential::refresh_token( + self.space_id, + &self.server_id, + refresh_token.secret().to_string(), + None, // Refresh tokens typically don't have a fixed expiry + ); + self.credential_repo + .save(&refresh_cred) + .await + .map_err(|e| { + AuthError::InternalError(format!("Failed to save refresh token: {}", e)) + })?; + } + // If no refresh_token in response, existing refresh_token row stays untouched debug!( - "[CredentialStore] Saved token for {}/{}", + "[CredentialStore] Saved tokens for {}/{}", self.space_id, self.server_id ); } // Save/update client registration if we have a new client_id - // Note: client_id comes from DCR, we need to preserve redirect_uri from existing registration if !creds.client_id.is_empty() { let existing_reg = self .backend_oauth_repo @@ -123,14 +122,12 @@ impl DatabaseCredentialStore { .ok() .flatten(); - // Only save if new or client_id changed let should_save = match &existing_reg { None => true, Some(reg) => reg.client_id != creds.client_id, }; if should_save { - // Preserve redirect_uri from existing registration, or use empty if new let redirect_uri = existing_reg .as_ref() .and_then(|r| r.redirect_uri.clone()) @@ -170,40 +167,49 @@ impl CredentialStore for DatabaseCredentialStore { self.space_id, self.server_id ); - // NOTE: We intentionally DO NOT use a cache here because expires_in - // must be recalculated on every load() call. RMCP's AuthClient calls - // load() before each request to check if the token is expired. - // If we cache the StoredCredentials with the OAuthTokenResponse, - // the expires_in Duration becomes stale and RMCP won't refresh - // expired tokens properly. - - // Load from database + // Load from database — no caching, expires_in recalculated each time let registration = self .backend_oauth_repo .get(&self.space_id, &self.server_id) .await .map_err(|e| AuthError::InternalError(format!("Failed to load registration: {}", e)))?; - let credential = self + // Load access_token and refresh_token as separate rows + let access_cred = self .credential_repo - .get(&self.space_id, &self.server_id) + .get( + &self.space_id, + &self.server_id, + &CredentialType::AccessToken, + ) .await - .map_err(|e| AuthError::InternalError(format!("Failed to load credential: {}", e)))?; + .map_err(|e| AuthError::InternalError(format!("Failed to load access token: {}", e)))?; - let stored = match (registration, credential) { - (Some(reg), Some(cred)) => { + let refresh_cred = self + .credential_repo + .get( + &self.space_id, + &self.server_id, + &CredentialType::RefreshToken, + ) + .await + .map_err(|e| { + AuthError::InternalError(format!("Failed to load refresh token: {}", e)) + })?; + + let stored = match (registration, access_cred.as_ref()) { + (Some(reg), Some(access)) => { debug!( "[CredentialStore] Loaded registration + token for {}/{}, client_id={}", self.space_id, self.server_id, reg.client_id ); - let token_response = Self::to_token_response(&cred); + let token_response = Self::build_token_response(access, refresh_cred.as_ref()); Some(StoredCredentials { client_id: reg.client_id, - token_response, + token_response: Some(token_response), }) } (Some(reg), None) => { - // Have registration but no token yet debug!( "[CredentialStore] Loaded registration (no token) for {}/{}, client_id={} - will reuse for DCR", self.space_id, self.server_id, reg.client_id @@ -213,17 +219,15 @@ impl CredentialStore for DatabaseCredentialStore { token_response: None, }) } - (None, Some(cred)) => { - // Shouldn't happen normally, but handle gracefully + (None, Some(access)) => { warn!( "[CredentialStore] Token without registration for {}/{}", self.space_id, self.server_id ); - let token_response = Self::to_token_response(&cred); - // Use empty client_id - will need to re-register + let token_response = Self::build_token_response(access, refresh_cred.as_ref()); Some(StoredCredentials { client_id: String::new(), - token_response, + token_response: Some(token_response), }) } (None, None) => { @@ -239,14 +243,10 @@ impl CredentialStore for DatabaseCredentialStore { } async fn save(&self, credentials: StoredCredentials) -> Result<(), AuthError> { - // Save to database - self.save_to_database(&credentials).await?; - - Ok(()) + self.save_to_database(&credentials).await } async fn clear(&self) -> Result<(), AuthError> { - // Clear tokens only (keep registration for re-auth) self.credential_repo .clear_tokens(&self.space_id, &self.server_id) .await @@ -261,7 +261,6 @@ impl CredentialStore for DatabaseCredentialStore { } /// Build an OAuthTokenResponse from components. -/// This creates a response compatible with oauth2 crate's StandardTokenResponse. fn build_token_response( access_token: String, refresh_token: Option, @@ -294,49 +293,100 @@ mod tests { // Mock implementations for testing #[derive(Clone)] struct MockCredentialRepo { - credential: Arc>>, + credentials: Arc>>, } impl MockCredentialRepo { fn new() -> Self { Self { - credential: Arc::new(tokio::sync::RwLock::new(None)), + credentials: Arc::new(tokio::sync::RwLock::new(Vec::new())), } } - - async fn set(&self, cred: Credential) { - *self.credential.write().await = Some(cred); - } } #[async_trait] impl CredentialRepository for MockCredentialRepo { async fn get( &self, - _space_id: &Uuid, - _server_id: &str, + space_id: &Uuid, + server_id: &str, + credential_type: &CredentialType, ) -> anyhow::Result> { - Ok(self.credential.read().await.clone()) + let creds = self.credentials.read().await; + Ok(creds + .iter() + .find(|c| { + c.space_id == *space_id + && c.server_id == server_id + && c.credential_type == *credential_type + }) + .cloned()) + } + + async fn get_all( + &self, + space_id: &Uuid, + server_id: &str, + ) -> anyhow::Result> { + let creds = self.credentials.read().await; + Ok(creds + .iter() + .filter(|c| c.space_id == *space_id && c.server_id == server_id) + .cloned() + .collect()) } async fn save(&self, credential: &Credential) -> anyhow::Result<()> { - *self.credential.write().await = Some(credential.clone()); + let mut creds = self.credentials.write().await; + // Upsert: remove existing with same key, then insert + creds.retain(|c| { + !(c.space_id == credential.space_id + && c.server_id == credential.server_id + && c.credential_type == credential.credential_type) + }); + creds.push(credential.clone()); Ok(()) } - async fn delete(&self, _space_id: &Uuid, _server_id: &str) -> anyhow::Result<()> { - *self.credential.write().await = None; + async fn delete( + &self, + space_id: &Uuid, + server_id: &str, + credential_type: &CredentialType, + ) -> anyhow::Result<()> { + let mut creds = self.credentials.write().await; + creds.retain(|c| { + !(c.space_id == *space_id + && c.server_id == server_id + && c.credential_type == *credential_type) + }); Ok(()) } - async fn clear_tokens(&self, _space_id: &Uuid, _server_id: &str) -> anyhow::Result { - let had_token = self.credential.read().await.is_some(); - *self.credential.write().await = None; - Ok(had_token) + async fn delete_all(&self, space_id: &Uuid, server_id: &str) -> anyhow::Result<()> { + let mut creds = self.credentials.write().await; + creds.retain(|c| !(c.space_id == *space_id && c.server_id == server_id)); + Ok(()) } - async fn list_for_space(&self, _space_id: &Uuid) -> anyhow::Result> { - Ok(vec![]) + async fn clear_tokens(&self, space_id: &Uuid, server_id: &str) -> anyhow::Result { + let mut creds = self.credentials.write().await; + let before = creds.len(); + creds.retain(|c| { + !(c.space_id == *space_id + && c.server_id == server_id + && c.credential_type.is_oauth()) + }); + Ok(creds.len() < before) + } + + async fn list_for_space(&self, space_id: &Uuid) -> anyhow::Result> { + let creds = self.credentials.read().await; + Ok(creds + .iter() + .filter(|c| c.space_id == *space_id) + .cloned() + .collect()) } } @@ -402,8 +452,6 @@ mod tests { #[tokio::test] async fn test_expires_in_recalculated_on_each_load() { - // This test verifies the critical fix: expires_in must be recalculated - // on each load() call, not cached with stale values let space_id = Uuid::new_v4(); let server_id = "test-server"; let server_url = "https://test.example.com"; @@ -421,23 +469,17 @@ mod tests { ); oauth_repo.set(registration).await; - // Set up a credential that expires in 10 seconds - let expires_at = Utc::now() + Duration::seconds(10); - let credential = Credential { + // Set up access_token that expires in 10 seconds + let access_cred = Credential::access_token( space_id, - server_id: server_id.to_string(), - value: CredentialValue::OAuth { - access_token: "token123".to_string(), - refresh_token: Some("refresh123".to_string()), - expires_at: Some(expires_at), - token_type: "Bearer".to_string(), - scope: None, - }, - created_at: Utc::now(), - updated_at: Utc::now(), - last_used: Some(Utc::now()), - }; - cred_repo.set(credential).await; + server_id, + "token123", + Some(Utc::now() + Duration::seconds(10)), + ); + let refresh_cred = Credential::refresh_token(space_id, server_id, "refresh123", None); + + cred_repo.save(&access_cred).await.unwrap(); + cred_repo.save(&refresh_cred).await.unwrap(); let store = DatabaseCredentialStore::new(space_id, server_id, server_url, cred_repo, oauth_repo); @@ -446,7 +488,6 @@ mod tests { let stored1 = store.load().await.unwrap().unwrap(); let token1 = stored1.token_response.as_ref().unwrap(); let expires_in_1 = token1.expires_in().unwrap(); - assert!(expires_in_1.as_secs() >= 9 && expires_in_1.as_secs() <= 10); // Wait 2 seconds @@ -457,14 +498,12 @@ mod tests { let token2 = stored2.token_response.as_ref().unwrap(); let expires_in_2 = token2.expires_in().unwrap(); - // This is the critical assertion: expires_in should decrease because it's recalculated assert!( expires_in_2.as_secs() >= 7 && expires_in_2.as_secs() <= 8, "Expected expires_in to decrease from ~10s to ~8s, but got {} seconds", expires_in_2.as_secs() ); - // Verify it actually decreased assert!( expires_in_2 < expires_in_1, "expires_in should decrease on subsequent loads (was {}, now {})", @@ -475,7 +514,6 @@ mod tests { #[tokio::test] async fn test_expired_token_detected() { - // Verify that an expired token is properly detected let space_id = Uuid::new_v4(); let server_id = "test-server"; let server_url = "https://test.example.com"; @@ -483,7 +521,6 @@ mod tests { let cred_repo = Arc::new(MockCredentialRepo::new()); let oauth_repo = Arc::new(MockOAuthRepo::new()); - // Set up registration let registration = OutboundOAuthRegistration::new( space_id, server_id, @@ -493,28 +530,21 @@ mod tests { ); oauth_repo.set(registration).await; - // Set up a credential that already expired (5 seconds ago) - let expires_at = Utc::now() - Duration::seconds(5); - let credential = Credential { + // Set up access_token that already expired (5 seconds ago) + let access_cred = Credential::access_token( space_id, - server_id: server_id.to_string(), - value: CredentialValue::OAuth { - access_token: "expired_token".to_string(), - refresh_token: Some("refresh123".to_string()), - expires_at: Some(expires_at), - token_type: "Bearer".to_string(), - scope: None, - }, - created_at: Utc::now(), - updated_at: Utc::now(), - last_used: Some(Utc::now()), - }; - cred_repo.set(credential).await; + server_id, + "expired_token", + Some(Utc::now() - Duration::seconds(5)), + ); + let refresh_cred = Credential::refresh_token(space_id, server_id, "refresh123", None); + + cred_repo.save(&access_cred).await.unwrap(); + cred_repo.save(&refresh_cred).await.unwrap(); let store = DatabaseCredentialStore::new(space_id, server_id, server_url, cred_repo, oauth_repo); - // Load should return token with expires_in = 0 (expired) let stored = store.load().await.unwrap().unwrap(); let token = stored.token_response.as_ref().unwrap(); let expires_in = token.expires_in().unwrap(); @@ -529,7 +559,6 @@ mod tests { #[tokio::test] async fn test_save_updates_database() { - // Verify that save() writes to database, not just cache let space_id = Uuid::new_v4(); let server_id = "test-server"; let server_url = "https://test.example.com"; @@ -559,16 +588,77 @@ mod tests { store.save(credentials).await.unwrap(); - // Verify they were written to database by checking the mock repo directly - let saved_cred = cred_repo.get(&space_id, server_id).await.unwrap().unwrap(); - match saved_cred.value { - CredentialValue::OAuth { access_token, .. } => { - assert_eq!(access_token, "new_token"); - } - _ => panic!("Expected OAuth credential"), - } + // Verify access_token row + let saved_access = cred_repo + .get(&space_id, server_id, &CredentialType::AccessToken) + .await + .unwrap() + .unwrap(); + assert_eq!(saved_access.value, "new_token"); + + // Verify refresh_token row + let saved_refresh = cred_repo + .get(&space_id, server_id, &CredentialType::RefreshToken) + .await + .unwrap() + .unwrap(); + assert_eq!(saved_refresh.value, "new_refresh"); + // Verify registration let saved_reg = oauth_repo.get(&space_id, server_id).await.unwrap().unwrap(); assert_eq!(saved_reg.client_id, "new-client-id"); } + + #[tokio::test] + async fn test_refresh_token_preserved_when_not_in_response() { + let space_id = Uuid::new_v4(); + let server_id = "test-server"; + let server_url = "https://test.example.com"; + + let cred_repo = Arc::new(MockCredentialRepo::new()); + let oauth_repo = Arc::new(MockOAuthRepo::new()); + + // Pre-populate with existing refresh token + let existing_refresh = + Credential::refresh_token(space_id, server_id, "original_refresh", None); + cred_repo.save(&existing_refresh).await.unwrap(); + + let store = DatabaseCredentialStore::new( + space_id, + server_id, + server_url, + Arc::clone(&cred_repo) as Arc, + Arc::clone(&oauth_repo) as Arc, + ); + + // Save new token response WITHOUT refresh_token + let token_response = build_token_response( + "rotated_access".to_string(), + None, // No refresh token in response + Some(std::time::Duration::from_secs(3600)), + ); + + let credentials = StoredCredentials { + client_id: "client-id".to_string(), + token_response: Some(token_response), + }; + + store.save(credentials).await.unwrap(); + + // Access token should be updated + let saved_access = cred_repo + .get(&space_id, server_id, &CredentialType::AccessToken) + .await + .unwrap() + .unwrap(); + assert_eq!(saved_access.value, "rotated_access"); + + // Refresh token should still be the original (not overwritten) + let saved_refresh = cred_repo + .get(&space_id, server_id, &CredentialType::RefreshToken) + .await + .unwrap() + .unwrap(); + assert_eq!(saved_refresh.value, "original_refresh"); + } } diff --git a/crates/mcpmux-gateway/src/pool/oauth.rs b/crates/mcpmux-gateway/src/pool/oauth.rs index 42e29e84..e3d39c42 100644 --- a/crates/mcpmux-gateway/src/pool/oauth.rs +++ b/crates/mcpmux-gateway/src/pool/oauth.rs @@ -20,7 +20,7 @@ use anyhow::{Context, Result}; use chrono::{DateTime, Utc}; use dashmap::DashMap; use mcpmux_core::{ - branding, CredentialRepository, CredentialValue, LogLevel, LogSource, OutboundOAuthRepository, + branding, CredentialRepository, CredentialType, LogLevel, LogSource, OutboundOAuthRepository, ServerLog, ServerLogManager, }; use rmcp::transport::auth::{AuthError, AuthorizationManager, AuthorizationSession, OAuthState}; @@ -675,36 +675,40 @@ impl OutboundOAuthManager { space_id: Uuid, server_id: &str, ) -> Option { - match credential_repo.get(&space_id, server_id).await { - Ok(Some(credential)) => { - if let CredentialValue::OAuth { - access_token, - refresh_token, - expires_at, - token_type, - scope, - } = credential.value - { - Some(OAuthTokenInfo { - access_token, - refresh_token, - expires_at, - token_type, - scope, - }) - } else { - None - } - } - Ok(None) => None, + // Load access token row + let access_cred = match credential_repo + .get(&space_id, server_id, &CredentialType::AccessToken) + .await + { + Ok(Some(cred)) => cred, + Ok(None) => return None, Err(e) => { warn!( - "[OAuth] Failed to load credential for {}/{}: {}", + "[OAuth] Failed to load access token for {}/{}: {}", space_id, server_id, e ); - None + return None; } - } + }; + + // Load refresh token row (optional) + let refresh_token = match credential_repo + .get(&space_id, server_id, &CredentialType::RefreshToken) + .await + { + Ok(Some(cred)) => Some(cred.value), + _ => None, + }; + + Some(OAuthTokenInfo { + access_token: access_cred.value, + refresh_token, + expires_at: access_cred.expires_at, + token_type: access_cred + .token_type + .unwrap_or_else(|| "Bearer".to_string()), + scope: access_cred.scope, + }) } /// Create an AuthorizationManager with our database-backed credential store. diff --git a/crates/mcpmux-gateway/src/pool/routing.rs b/crates/mcpmux-gateway/src/pool/routing.rs index 85d83abe..180e09e6 100644 --- a/crates/mcpmux-gateway/src/pool/routing.rs +++ b/crates/mcpmux-gateway/src/pool/routing.rs @@ -17,6 +17,7 @@ use serde_json::Value; use tracing::{debug, info, warn}; use uuid::Uuid; +use super::connection::ConnectionResult; use super::features::FeatureService; use super::service::PoolService; @@ -332,32 +333,222 @@ impl RoutingService { Ok(result) => { let duration = call_start.elapsed(); if result.is_error { - warn!( - "[RoutingService] Tool execution error: {} (duration: {:?})", - actual_tool_name, duration - ); - self.log( - &space_id, - &server_id, - LogLevel::Error, - format!("Tool execution error: {}", actual_tool_name), - Some(serde_json::json!({ "result": result.content, "duration_ms": duration.as_millis() })) - ).await; + // Check if this is an auth error embedded in the tool result. + // Some servers (e.g., Atlassian) return 401 as tool results rather than + // HTTP errors. The SDK refreshes the token successfully, but the server's + // internal session may be stale. A fresh MCP connection fixes this. + if Self::content_has_auth_error(&result.content) { + warn!( + "[RoutingService] Auth error in tool result for {}/{}, attempting auto-reconnect", + server_id, actual_tool_name + ); + self.log( + &space_id, + &server_id, + LogLevel::Warn, + format!( + "Auth error in tool result for '{}' - auto-reconnecting", + actual_tool_name + ), + Some(serde_json::json!({ "result": result.content, "duration_ms": duration.as_millis() })), + ) + .await; + + match self + .pool_service + .reconnect_instance(space_id, &server_id) + .await + { + ConnectionResult::Connected { .. } => { + info!( + "[RoutingService] Reconnected {}, retrying tool call: {}", + server_id, actual_tool_name + ); + + let retry_start = std::time::Instant::now(); + match execute_call( + self.pool_service.clone(), + space_id, + server_id.clone(), + actual_tool_name.clone(), + arguments.clone(), + ) + .await + { + Ok(retry_result) => { + let retry_duration = retry_start.elapsed(); + if retry_result.is_error { + warn!( + "[RoutingService] Tool retry still has error: {} (duration: {:?})", + actual_tool_name, retry_duration + ); + } else { + info!( + "[RoutingService] Tool retry succeeded after reconnect: {} (duration: {:?})", + actual_tool_name, retry_duration + ); + } + self.log( + &space_id, + &server_id, + LogLevel::Info, + format!( + "Tool '{}' retried after auto-reconnect (is_error={})", + actual_tool_name, retry_result.is_error + ), + Some(serde_json::json!({ "retry_duration_ms": retry_duration.as_millis() })), + ) + .await; + Ok(retry_result) + } + Err(retry_err) => { + warn!( + "[RoutingService] Tool retry transport error: {} - {}", + actual_tool_name, retry_err + ); + self.log( + &space_id, + &server_id, + LogLevel::Error, + format!( + "Tool '{}' still failing after reconnect", + actual_tool_name + ), + Some(serde_json::json!({ "error": retry_err.to_string() })), + ) + .await; + // Return original tool result since it has the error details + Ok(result) + } + } + } + other => { + warn!( + "[RoutingService] Auto-reconnect failed for {}: {:?}", + server_id, other + ); + self.log( + &space_id, + &server_id, + LogLevel::Error, + format!( + "Auto-reconnect failed for tool '{}' - manual reconnection required", + actual_tool_name + ), + Some(serde_json::json!({ "reconnect_result": format!("{:?}", other) })), + ) + .await; + Ok(result) + } + } + } else { + warn!( + "[RoutingService] Tool execution error: {} (duration: {:?})", + actual_tool_name, duration + ); + self.log( + &space_id, + &server_id, + LogLevel::Error, + format!("Tool execution error: {}", actual_tool_name), + Some(serde_json::json!({ "result": result.content, "duration_ms": duration.as_millis() })) + ).await; + Ok(result) + } } else { - info!( - "[RoutingService] Tool executed successfully: {} (duration: {:?})", - actual_tool_name, duration - ); - self.log( - &space_id, - &server_id, - LogLevel::Info, - format!("Tool executed successfully: {}", actual_tool_name), - Some(serde_json::json!({ "duration_ms": duration.as_millis() })), - ) - .await; + // Even on "success" (is_error=false), some servers (e.g., Atlassian) + // return auth errors as plain text content like {"code":401,"message":"Unauthorized"}. + // Detect these and auto-reconnect + retry. + if Self::content_has_auth_error(&result.content) { + warn!( + "[RoutingService] Auth error in successful tool result for {}/{}, attempting auto-reconnect", + server_id, actual_tool_name + ); + self.log( + &space_id, + &server_id, + LogLevel::Warn, + format!( + "Auth error in tool result for '{}' (is_error=false) - auto-reconnecting", + actual_tool_name + ), + Some(serde_json::json!({ "result": result.content, "duration_ms": duration.as_millis() })), + ) + .await; + + match self + .pool_service + .reconnect_instance(space_id, &server_id) + .await + { + ConnectionResult::Connected { .. } => { + info!( + "[RoutingService] Reconnected {}, retrying tool call: {}", + server_id, actual_tool_name + ); + + let retry_start = std::time::Instant::now(); + match execute_call( + self.pool_service.clone(), + space_id, + server_id.clone(), + actual_tool_name.clone(), + arguments.clone(), + ) + .await + { + Ok(retry_result) => { + let retry_duration = retry_start.elapsed(); + info!( + "[RoutingService] Tool retry result: {} (is_error={}, duration: {:?})", + actual_tool_name, retry_result.is_error, retry_duration + ); + self.log( + &space_id, + &server_id, + LogLevel::Info, + format!( + "Tool '{}' retried after auto-reconnect (is_error={})", + actual_tool_name, retry_result.is_error + ), + Some(serde_json::json!({ "retry_duration_ms": retry_duration.as_millis() })), + ) + .await; + Ok(retry_result) + } + Err(retry_err) => { + warn!( + "[RoutingService] Tool retry transport error: {} - {}", + actual_tool_name, retry_err + ); + Ok(result) + } + } + } + other => { + warn!( + "[RoutingService] Auto-reconnect failed for {}: {:?}", + server_id, other + ); + Ok(result) + } + } + } else { + info!( + "[RoutingService] Tool executed successfully: {} (duration: {:?})", + actual_tool_name, duration + ); + self.log( + &space_id, + &server_id, + LogLevel::Info, + format!("Tool executed successfully: {}", actual_tool_name), + Some(serde_json::json!({ "duration_ms": duration.as_millis() })), + ) + .await; + Ok(result) + } } - Ok(result) } Err(e) => { let duration = call_start.elapsed(); @@ -368,26 +559,116 @@ impl RoutingService { actual_tool_name, server_id, e, duration ); - // Check if it's an auth error - // NOTE: With RMCP's AuthClient, token refresh happens automatically per-request. - // If we still get an auth error, it means the refresh token is invalid or expired. - // The user needs to reconnect to re-authorize. let is_auth = Self::is_auth_error(&err_str); - let is_timeout = err_str.contains("timed out"); - if is_auth || is_timeout { + if is_auth { + // Auth error detected - attempt auto-reconnect and retry once. + // This handles the case where RMCP's AuthClient failed to refresh + // the token (e.g., stale in-memory state after idle). + // Creating a fresh connection loads latest tokens from the database. warn!( - "[RoutingService] Auth/timeout error for {}/{} - RMCP auto-refresh likely failed, user needs to reconnect", + "[RoutingService] Auth error for {}/{}, attempting auto-reconnect", server_id, actual_tool_name ); self.log( &space_id, &server_id, - LogLevel::Error, - format!("Authentication failed for tool '{}' - reconnection required", actual_tool_name), - Some(serde_json::json!({ "error": e.to_string(), "duration_ms": duration.as_millis() })) - ).await; - Err(anyhow!("Server '{}' requires reconnection. Token may have expired. Please disconnect and connect again.", server_id)) + LogLevel::Warn, + format!( + "Auth error on tool '{}' - auto-reconnecting to refresh credentials", + actual_tool_name + ), + Some(serde_json::json!({ "error": e.to_string(), "duration_ms": duration.as_millis() })), + ) + .await; + + match self + .pool_service + .reconnect_instance(space_id, &server_id) + .await + { + ConnectionResult::Connected { .. } => { + info!( + "[RoutingService] Reconnected {}, retrying tool call: {}", + server_id, actual_tool_name + ); + + // Retry the call once with the fresh connection + let retry_start = std::time::Instant::now(); + match execute_call( + self.pool_service.clone(), + space_id, + server_id.clone(), + actual_tool_name.clone(), + arguments.clone(), + ) + .await + { + Ok(result) => { + let retry_duration = retry_start.elapsed(); + info!( + "[RoutingService] Tool retry succeeded: {} (duration: {:?})", + actual_tool_name, retry_duration + ); + self.log( + &space_id, + &server_id, + LogLevel::Info, + format!( + "Tool '{}' succeeded after auto-reconnect", + actual_tool_name + ), + Some(serde_json::json!({ "retry_duration_ms": retry_duration.as_millis() })), + ) + .await; + Ok(result) + } + Err(retry_err) => { + warn!( + "[RoutingService] Tool retry also failed: {} - {}", + actual_tool_name, retry_err + ); + self.log( + &space_id, + &server_id, + LogLevel::Error, + format!( + "Tool '{}' still failing after reconnect - manual reconnection required", + actual_tool_name + ), + Some(serde_json::json!({ "error": retry_err.to_string() })), + ) + .await; + Err(anyhow!( + "Server '{}' auth error persists after auto-reconnect. Please disconnect and connect again. Error: {}", + server_id, + retry_err + )) + } + } + } + other => { + warn!( + "[RoutingService] Auto-reconnect failed for {}: {:?}", + server_id, other + ); + self.log( + &space_id, + &server_id, + LogLevel::Error, + format!( + "Auto-reconnect failed for tool '{}' - manual reconnection required", + actual_tool_name + ), + Some(serde_json::json!({ "reconnect_result": format!("{:?}", other) })), + ) + .await; + Err(anyhow!( + "Server '{}' requires reconnection. Auto-reconnect failed. Please disconnect and connect again.", + server_id + )) + } + } } else { // Not an auth error, return original error self.log( @@ -395,8 +676,9 @@ impl RoutingService { &server_id, LogLevel::Error, format!("Tool call failed: {}", e), - Some(serde_json::json!({ "error": e.to_string(), "duration_ms": duration.as_millis() })) - ).await; + Some(serde_json::json!({ "error": e.to_string(), "duration_ms": duration.as_millis() })), + ) + .await; Err(e) } } @@ -426,7 +708,7 @@ impl RoutingService { } } - /// Check if an error indicates authentication is needed + /// Check if an error string indicates authentication is needed fn is_auth_error(error_str: &str) -> bool { let indicators = [ "401", @@ -437,4 +719,22 @@ impl RoutingService { ]; indicators.iter().any(|s| error_str.contains(s)) } + + /// Check if tool result content contains authentication error indicators. + /// + /// Some MCP servers (e.g., Atlassian) return auth errors as tool results + /// (`is_error: true` with 401 in the text) rather than HTTP-level errors. + /// The SDK may have already refreshed the token, but the server's internal + /// session can be stale. A fresh connection (reconnect) fixes this. + fn content_has_auth_error(content: &[Value]) -> bool { + for item in content { + if let Some(text) = item.get("text").and_then(|v| v.as_str()) { + let lower = text.to_lowercase(); + if Self::is_auth_error(&lower) { + return true; + } + } + } + false + } } diff --git a/crates/mcpmux-gateway/src/pool/service.rs b/crates/mcpmux-gateway/src/pool/service.rs index 1d4d6dd0..d8208a33 100644 --- a/crates/mcpmux-gateway/src/pool/service.rs +++ b/crates/mcpmux-gateway/src/pool/service.rs @@ -26,6 +26,19 @@ use super::oauth::OutboundOAuthManager; use super::token::TokenService; use super::transport::{ResolvedTransport, TransportType}; +/// Check if an error string indicates an authentication/authorization failure +fn is_auth_error(error_str: &str) -> bool { + let lower = error_str.to_lowercase(); + let indicators = [ + "401", + "unauthorized", + "invalid_token", + "token expired", + "access token", + ]; + indicators.iter().any(|s| lower.contains(s)) +} + /// Result of bulk reconnect operation #[derive(Debug, Default)] pub struct ReconnectResult { @@ -101,11 +114,45 @@ impl PoolService { } /// Read a resource from a backend server + /// + /// On auth errors, automatically reconnects the server and retries once. pub async fn read_resource( &self, space_id: Uuid, server_id: &str, uri: &str, + ) -> Result> { + match self.try_read_resource(space_id, server_id, uri).await { + Ok(content) => Ok(content), + Err(e) if is_auth_error(&e.to_string()) => { + warn!( + "[PoolService] Auth error on read_resource for {}/{}, attempting auto-reconnect", + server_id, uri + ); + match self.reconnect_instance(space_id, server_id).await { + ConnectionResult::Connected { .. } => { + info!( + "[PoolService] Reconnected {}, retrying read_resource", + server_id + ); + self.try_read_resource(space_id, server_id, uri).await + } + _ => Err(anyhow::anyhow!( + "Server '{}' auth error on read_resource. Auto-reconnect failed. Please disconnect and connect again.", + server_id + )), + } + } + Err(e) => Err(e), + } + } + + /// Internal: attempt to read a resource without retry logic + async fn try_read_resource( + &self, + space_id: Uuid, + server_id: &str, + uri: &str, ) -> Result> { let instance = self .get_instance(space_id, server_id) @@ -140,12 +187,51 @@ impl PoolService { } /// Get a prompt from a backend server + /// + /// On auth errors, automatically reconnects the server and retries once. pub async fn get_prompt( &self, space_id: Uuid, server_id: &str, prompt_name: &str, arguments: Option>, + ) -> Result { + match self + .try_get_prompt(space_id, server_id, prompt_name, arguments.clone()) + .await + { + Ok(value) => Ok(value), + Err(e) if is_auth_error(&e.to_string()) => { + warn!( + "[PoolService] Auth error on get_prompt for {}/{}, attempting auto-reconnect", + server_id, prompt_name + ); + match self.reconnect_instance(space_id, server_id).await { + ConnectionResult::Connected { .. } => { + info!( + "[PoolService] Reconnected {}, retrying get_prompt", + server_id + ); + self.try_get_prompt(space_id, server_id, prompt_name, arguments) + .await + } + _ => Err(anyhow::anyhow!( + "Server '{}' auth error on get_prompt. Auto-reconnect failed. Please disconnect and connect again.", + server_id + )), + } + } + Err(e) => Err(e), + } + } + + /// Internal: attempt to get a prompt without retry logic + async fn try_get_prompt( + &self, + space_id: Uuid, + server_id: &str, + prompt_name: &str, + arguments: Option>, ) -> Result { let instance = self .get_instance(space_id, server_id) diff --git a/crates/mcpmux-gateway/src/pool/transport/http.rs b/crates/mcpmux-gateway/src/pool/transport/http.rs index 7dc99aa7..c7d4bfbc 100644 --- a/crates/mcpmux-gateway/src/pool/transport/http.rs +++ b/crates/mcpmux-gateway/src/pool/transport/http.rs @@ -306,13 +306,17 @@ impl HttpTransport { "Connecting with manual token injection (RMCP metadata failed)" ); - // Load token from our database - let credential = match self + // Load access token from our database + let access_token = match self .credential_repo - .get(&self.space_id, &self.server_id) + .get( + &self.space_id, + &self.server_id, + &mcpmux_core::CredentialType::AccessToken, + ) .await { - Ok(Some(cred)) => cred, + Ok(Some(cred)) => cred.value, Ok(None) => { debug!(server_id = %self.server_id, "No stored token for manual injection"); return TransportConnectResult::OAuthRequired { @@ -326,16 +330,6 @@ impl HttpTransport { } }; - // Extract access token from credential value - let access_token = match &credential.value { - mcpmux_core::CredentialValue::OAuth { access_token, .. } => access_token.clone(), - _ => { - let err = "Credential is not OAuth type".to_string(); - error!(server_id = %self.server_id, "{}", err); - return TransportConnectResult::Failed(err); - } - }; - self.log( LogLevel::Info, LogSource::HttpRequest, @@ -516,7 +510,11 @@ impl Transport for HttpTransport { // Check if we have stored credentials for this server let has_credentials = self .credential_repo - .get(&self.space_id, &self.server_id) + .get( + &self.space_id, + &self.server_id, + &mcpmux_core::CredentialType::AccessToken, + ) .await .ok() .flatten() diff --git a/crates/mcpmux-gateway/src/server/handlers.rs b/crates/mcpmux-gateway/src/server/handlers.rs index 72800b2e..fa91f4f6 100644 --- a/crates/mcpmux-gateway/src/server/handlers.rs +++ b/crates/mcpmux-gateway/src/server/handlers.rs @@ -647,10 +647,32 @@ pub async fn oauth_token( )); }; - // Update last_seen in database + // Verify client still exists in DB before issuing new tokens. + // The JWT may be valid (same secret) but the client may have been + // removed (e.g., DB was reset). Without this check, the middleware + // would fail with "Client not found" after we issue a new token. if let Some(repo) = gateway_state.inbound_client_repository() { - if let Err(e) = repo.update_client_last_seen(&claims.client_id).await { - warn!("[OAuth] Failed to update last_seen: {}", e); + match repo.get_client(&claims.client_id).await { + Ok(Some(_)) => { + // Client exists, update last_seen + if let Err(e) = repo.update_client_last_seen(&claims.client_id).await { + warn!("[OAuth] Failed to update last_seen: {}", e); + } + } + Ok(None) => { + warn!( + "[OAuth] Client {} not found in DB during refresh", + claims.client_id + ); + return Err(token_error("invalid_grant", "Client no longer registered")); + } + Err(e) => { + warn!( + "[OAuth] Failed to look up client {}: {}", + claims.client_id, e + ); + return Err(token_error("server_error", "Database error")); + } } } diff --git a/crates/mcpmux-gateway/src/services/authorization.rs b/crates/mcpmux-gateway/src/services/authorization.rs index 4cc8dc46..83cf5c68 100644 --- a/crates/mcpmux-gateway/src/services/authorization.rs +++ b/crates/mcpmux-gateway/src/services/authorization.rs @@ -32,10 +32,15 @@ impl AuthorizationService { /// Get effective feature set grants for a client in a specific space. /// - /// Uses layered resolution: Returns explicit grants PLUS the default feature set - /// (deduplicated as a set). This ensures all clients always get default permissions. + /// Resolution strategy (least-privilege by default): + /// 1. Return explicit per-client grants from DB if any exist. + /// 2. Always include the Default feature set as a baseline. /// - /// Returns Vec of feature_set_ids that the client has access to. + /// Clients with no explicit grants only receive the Default feature set, + /// which starts empty (no features). The user must explicitly grant + /// additional feature sets (e.g. "All", "ServerAll", or custom sets) + /// through the UI to expose tools/prompts/resources to a client. + /// This avoids accidental exposure of all server capabilities. pub async fn get_client_grants(&self, client_id: &str, space_id: &Uuid) -> Result> { let space_id_str = space_id.to_string(); @@ -45,13 +50,13 @@ impl AuthorizationService { .get_grants_for_space(client_id, &space_id_str) .await?; - // Add default feature set (layered resolution) + // Always include the Default feature set as baseline permissions. + // Default starts empty — user must explicitly grant additional access. if let Some(default_fs) = self .feature_set_repo .get_default_for_space(&space_id_str) .await? { - // Add default if not already in grants (set semantics - no repetition) if !grants.contains(&default_fs.id) { grants.push(default_fs.id); } diff --git a/crates/mcpmux-storage/Cargo.toml b/crates/mcpmux-storage/Cargo.toml index 05c0654b..cfae207a 100644 --- a/crates/mcpmux-storage/Cargo.toml +++ b/crates/mcpmux-storage/Cargo.toml @@ -26,6 +26,9 @@ keyring.workspace = true zeroize.workspace = true sha2 = "0.10" +[target.'cfg(windows)'.dependencies] +windows-dpapi = "0.1" + [dev-dependencies] tokio = { workspace = true, features = ["test-util", "macros"] } tempfile = "3.14" diff --git a/crates/mcpmux-storage/src/database.rs b/crates/mcpmux-storage/src/database.rs index 7818696c..86f35d68 100644 --- a/crates/mcpmux-storage/src/database.rs +++ b/crates/mcpmux-storage/src/database.rs @@ -32,28 +32,11 @@ struct Migration { /// Note: Migrations have been consolidated into a single clean initial migration. /// The schema includes cached_definition for offline operation and excludes /// runtime fields (connection_status, last_connected_at, last_error). -const MIGRATIONS: &[Migration] = &[ - Migration { - version: 1, - name: "initial", - sql: include_str!("migrations/001_initial.sql"), - }, - Migration { - version: 2, - name: "add_client_approved", - sql: include_str!("migrations/002_add_client_approved.sql"), - }, - Migration { - version: 3, - name: "app_settings", - sql: include_str!("migrations/003_app_settings.sql"), - }, - Migration { - version: 4, - name: "add_installation_source", - sql: include_str!("migrations/004_add_installation_source.sql"), - }, -]; +const MIGRATIONS: &[Migration] = &[Migration { + version: 1, + name: "initial", + sql: include_str!("migrations/001_initial.sql"), +}]; /// SQLite database wrapper. pub struct Database { diff --git a/crates/mcpmux-storage/src/keychain_dpapi.rs b/crates/mcpmux-storage/src/keychain_dpapi.rs new file mode 100644 index 00000000..6d868a97 --- /dev/null +++ b/crates/mcpmux-storage/src/keychain_dpapi.rs @@ -0,0 +1,326 @@ +//! DPAPI-based key storage for Windows. +//! +//! Stores encryption keys as DPAPI-protected files instead of Windows Credential Manager. +//! This prevents the master key from being visible in the Credential Manager UI while +//! maintaining the same security guarantees (user-scope DPAPI protection). +//! +//! Key files are stored in `/keys/` as opaque encrypted blobs. + +use std::fs; +use std::path::{Path, PathBuf}; + +use anyhow::{Context, Result}; +use tracing::{debug, info, warn}; +use windows_dpapi::{decrypt_data, encrypt_data, Scope}; +use zeroize::Zeroizing; + +use crate::crypto::{generate_master_key, KEY_SIZE}; +use crate::keychain::{generate_jwt_secret, JwtSecretProvider, MasterKeyProvider, JWT_SECRET_SIZE}; + +/// File name for the DPAPI-protected master encryption key. +const MASTER_KEY_FILE: &str = "master.dpapi"; + +/// File name for the DPAPI-protected JWT signing secret. +const JWT_SECRET_FILE: &str = "jwt.dpapi"; + +/// DPAPI-based master key provider. +/// +/// Stores the master key in a DPAPI-protected file within the app's data directory. +/// The key is encrypted with user-scope DPAPI, meaning only the current Windows user +/// on this machine can decrypt it. +pub struct DpapiKeyProvider { + key_path: PathBuf, +} + +impl DpapiKeyProvider { + /// Create a new DPAPI key provider that stores keys in the given data directory. + pub fn new(data_dir: &Path) -> Result { + let keys_dir = data_dir.join("keys"); + fs::create_dir_all(&keys_dir) + .with_context(|| format!("Failed to create keys directory: {:?}", keys_dir))?; + + Ok(Self { + key_path: keys_dir.join(MASTER_KEY_FILE), + }) + } +} + +impl MasterKeyProvider for DpapiKeyProvider { + fn get_or_create_key(&self) -> Result> { + if self.key_path.exists() { + debug!( + "Reading DPAPI-protected master key from {:?}", + self.key_path + ); + let encrypted = fs::read(&self.key_path) + .with_context(|| format!("Failed to read key file: {:?}", self.key_path))?; + + let decrypted = decrypt_data(&encrypted, Scope::User) + .context("Failed to decrypt master key with DPAPI")?; + + if decrypted.len() != KEY_SIZE { + anyhow::bail!( + "Invalid key size in DPAPI file: expected {}, got {}", + KEY_SIZE, + decrypted.len() + ); + } + + let mut key = Zeroizing::new([0u8; KEY_SIZE]); + key.copy_from_slice(&decrypted); + debug!("Master key loaded from DPAPI-protected file"); + Ok(key) + } else { + info!("No master key found, generating new DPAPI-protected key"); + let key = generate_master_key()?; + + let encrypted = encrypt_data(&key, Scope::User) + .context("Failed to encrypt master key with DPAPI")?; + + fs::write(&self.key_path, &encrypted) + .with_context(|| format!("Failed to write key file: {:?}", self.key_path))?; + + info!("Master key generated and stored as DPAPI-protected file"); + Ok(Zeroizing::new(key)) + } + } + + fn key_exists(&self) -> bool { + self.key_path.exists() + } + + fn delete_key(&self) -> Result<()> { + if self.key_path.exists() { + fs::remove_file(&self.key_path) + .with_context(|| format!("Failed to delete key file: {:?}", self.key_path))?; + info!("Master key DPAPI file deleted"); + } else { + debug!("No DPAPI key file to delete"); + } + Ok(()) + } +} + +/// DPAPI-based JWT signing secret provider. +/// +/// Stores the JWT signing secret in a DPAPI-protected file. +pub struct DpapiJwtSecretProvider { + secret_path: PathBuf, +} + +impl DpapiJwtSecretProvider { + /// Create a new DPAPI JWT secret provider that stores secrets in the given data directory. + pub fn new(data_dir: &Path) -> Result { + let keys_dir = data_dir.join("keys"); + fs::create_dir_all(&keys_dir) + .with_context(|| format!("Failed to create keys directory: {:?}", keys_dir))?; + + Ok(Self { + secret_path: keys_dir.join(JWT_SECRET_FILE), + }) + } +} + +impl JwtSecretProvider for DpapiJwtSecretProvider { + fn get_or_create_secret(&self) -> Result> { + if self.secret_path.exists() { + debug!( + "Reading DPAPI-protected JWT secret from {:?}", + self.secret_path + ); + let encrypted = fs::read(&self.secret_path).with_context(|| { + format!("Failed to read JWT secret file: {:?}", self.secret_path) + })?; + + let decrypted = decrypt_data(&encrypted, Scope::User) + .context("Failed to decrypt JWT secret with DPAPI")?; + + if decrypted.len() != JWT_SECRET_SIZE { + anyhow::bail!( + "Invalid JWT secret size in DPAPI file: expected {}, got {}", + JWT_SECRET_SIZE, + decrypted.len() + ); + } + + let mut secret = Zeroizing::new([0u8; JWT_SECRET_SIZE]); + secret.copy_from_slice(&decrypted); + debug!("JWT secret loaded from DPAPI-protected file"); + Ok(secret) + } else { + info!("No JWT secret found, generating new DPAPI-protected secret"); + let secret = generate_jwt_secret()?; + + let encrypted = encrypt_data(&secret, Scope::User) + .context("Failed to encrypt JWT secret with DPAPI")?; + + fs::write(&self.secret_path, &encrypted).with_context(|| { + format!("Failed to write JWT secret file: {:?}", self.secret_path) + })?; + + info!("JWT secret generated and stored as DPAPI-protected file"); + Ok(Zeroizing::new(secret)) + } + } + + fn secret_exists(&self) -> bool { + self.secret_path.exists() + } + + fn delete_secret(&self) -> Result<()> { + if self.secret_path.exists() { + fs::remove_file(&self.secret_path).with_context(|| { + format!("Failed to delete JWT secret file: {:?}", self.secret_path) + })?; + info!("JWT secret DPAPI file deleted"); + } else { + debug!("No DPAPI JWT secret file to delete"); + } + Ok(()) + } +} + +/// Migrate existing keys from Windows Credential Manager to DPAPI files. +/// +/// If keys exist in Credential Manager but not as DPAPI files, this copies them +/// over and removes the Credential Manager entries. This is a one-time migration. +pub fn migrate_from_credential_manager(data_dir: &Path) -> Result<()> { + use crate::keychain::{KeychainJwtSecretProvider, KeychainKeyProvider}; + + let keys_dir = data_dir.join("keys"); + let master_dpapi_path = keys_dir.join(MASTER_KEY_FILE); + let jwt_dpapi_path = keys_dir.join(JWT_SECRET_FILE); + + // Migrate master key if DPAPI file doesn't exist yet + if !master_dpapi_path.exists() { + if let Ok(keychain_provider) = KeychainKeyProvider::new() { + if keychain_provider.key_exists() { + info!("Migrating master key from Credential Manager to DPAPI"); + match keychain_provider.get_or_create_key() { + Ok(key) => { + let dpapi_provider = DpapiKeyProvider::new(data_dir)?; + // Write through DPAPI provider to ensure proper encryption + let encrypted = encrypt_data(&*key, Scope::User) + .context("Failed to encrypt master key with DPAPI during migration")?; + fs::create_dir_all(&keys_dir)?; + fs::write(&master_dpapi_path, &encrypted)?; + + // Remove from Credential Manager + if let Err(e) = keychain_provider.delete_key() { + warn!("Failed to remove master key from Credential Manager after migration: {}", e); + } else { + info!("Master key migrated and removed from Credential Manager"); + } + drop(dpapi_provider); + } + Err(e) => { + warn!( + "Failed to read master key from Credential Manager for migration: {}", + e + ); + } + } + } + } + } + + // Migrate JWT secret if DPAPI file doesn't exist yet + if !jwt_dpapi_path.exists() { + if let Ok(keychain_provider) = KeychainJwtSecretProvider::new() { + if keychain_provider.secret_exists() { + info!("Migrating JWT secret from Credential Manager to DPAPI"); + match keychain_provider.get_or_create_secret() { + Ok(secret) => { + let encrypted = encrypt_data(&*secret, Scope::User) + .context("Failed to encrypt JWT secret with DPAPI during migration")?; + fs::create_dir_all(&keys_dir)?; + fs::write(&jwt_dpapi_path, &encrypted)?; + + // Remove from Credential Manager + if let Err(e) = keychain_provider.delete_secret() { + warn!("Failed to remove JWT secret from Credential Manager after migration: {}", e); + } else { + info!("JWT secret migrated and removed from Credential Manager"); + } + } + Err(e) => { + warn!( + "Failed to read JWT secret from Credential Manager for migration: {}", + e + ); + } + } + } + } + } + + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_dpapi_master_key_provider() { + let tmp = tempfile::tempdir().unwrap(); + let provider = DpapiKeyProvider::new(tmp.path()).unwrap(); + + // Initially no key + assert!(!provider.key_exists()); + + // Get or create generates a key + let key1 = provider.get_or_create_key().unwrap(); + assert!(provider.key_exists()); + + // Getting again returns the same key + let key2 = provider.get_or_create_key().unwrap(); + assert_eq!(&*key1, &*key2); + + // Delete removes the key + provider.delete_key().unwrap(); + assert!(!provider.key_exists()); + + // New key is generated after delete + let key3 = provider.get_or_create_key().unwrap(); + assert_ne!(&*key1, &*key3); + } + + #[test] + fn test_dpapi_jwt_secret_provider() { + let tmp = tempfile::tempdir().unwrap(); + let provider = DpapiJwtSecretProvider::new(tmp.path()).unwrap(); + + // Initially no secret + assert!(!provider.secret_exists()); + + // Get or create generates a secret + let secret1 = provider.get_or_create_secret().unwrap(); + assert!(provider.secret_exists()); + + // Getting again returns the same secret + let secret2 = provider.get_or_create_secret().unwrap(); + assert_eq!(&*secret1, &*secret2); + + // Delete removes the secret + provider.delete_secret().unwrap(); + assert!(!provider.secret_exists()); + } + + #[test] + fn test_dpapi_file_is_encrypted() { + let tmp = tempfile::tempdir().unwrap(); + let provider = DpapiKeyProvider::new(tmp.path()).unwrap(); + + let key = provider.get_or_create_key().unwrap(); + + // Read the raw file - it should NOT contain the key in plaintext + let file_contents = fs::read(tmp.path().join("keys").join(MASTER_KEY_FILE)).unwrap(); + + // DPAPI output is always larger than the input (has header + IV + tag) + assert!(file_contents.len() > KEY_SIZE); + + // The raw key bytes should not appear in the file + assert!(!file_contents.windows(KEY_SIZE).any(|w| w == &*key)); + } +} diff --git a/crates/mcpmux-storage/src/lib.rs b/crates/mcpmux-storage/src/lib.rs index 35a1a8c4..1c19037e 100644 --- a/crates/mcpmux-storage/src/lib.rs +++ b/crates/mcpmux-storage/src/lib.rs @@ -17,8 +17,8 @@ //! │ FieldEncryptor (AES-256-GCM) │ //! │ (Encrypts tokens/credentials) │ //! ├──────────────────────────────────────────────────────┤ -//! │ KeychainKeyProvider │ -//! │ (OS Keychain: Windows/macOS/Linux) │ +//! │ DpapiKeyProvider (Windows) / KeychainKeyProvider │ +//! │ (DPAPI file storage / OS Keychain) │ //! ├──────────────────────────────────────────────────────┤ //! │ Database │ //! │ (SQLite) │ @@ -30,13 +30,13 @@ //! ```rust,ignore //! use mcpmux_storage::{ //! Database, SqliteSpaceRepository, SqliteCredentialRepository, -//! FieldEncryptor, KeychainKeyProvider, MasterKeyProvider, +//! FieldEncryptor, MasterKeyProvider, //! }; //! use std::sync::Arc; //! use tokio::sync::Mutex; //! -//! // Get master key from OS keychain -//! let key_provider = KeychainKeyProvider::new()?; +//! // Get master key (DPAPI on Windows, OS Keychain elsewhere) +//! let key_provider = mcpmux_storage::create_key_provider(&data_dir)?; //! let master_key = key_provider.get_or_create_key()?; //! //! // Open database @@ -54,6 +54,8 @@ pub mod crypto; mod database; pub mod keychain; +#[cfg(windows)] +pub mod keychain_dpapi; mod repositories; pub use crypto::{generate_master_key, FieldEncryptor, KEY_SIZE}; @@ -62,6 +64,8 @@ pub use keychain::{ generate_jwt_secret, JwtSecretProvider, KeychainJwtSecretProvider, KeychainKeyProvider, MasterKeyProvider, JWT_SECRET_SIZE, }; +#[cfg(windows)] +pub use keychain_dpapi::{DpapiJwtSecretProvider, DpapiKeyProvider}; pub use repositories::*; /// Default database file name. @@ -71,3 +75,46 @@ pub const DATABASE_FILE: &str = "mcpmux.db"; pub fn default_database_path() -> Option { dirs::data_local_dir().map(|p| p.join("mcpmux").join(DATABASE_FILE)) } + +/// Create the platform-appropriate master key provider. +/// +/// - **Windows**: Uses DPAPI file-based storage (key not visible in Credential Manager UI). +/// Also migrates existing keys from Credential Manager on first use. +/// - **macOS/Linux**: Uses the OS keychain (Keychain / Secret Service). +pub fn create_key_provider( + data_dir: &std::path::Path, +) -> anyhow::Result> { + #[cfg(windows)] + { + // Migrate any existing keys from Credential Manager to DPAPI files + if let Err(e) = keychain_dpapi::migrate_from_credential_manager(data_dir) { + tracing::warn!("Credential Manager migration encountered an error: {}", e); + } + Ok(Box::new(DpapiKeyProvider::new(data_dir)?)) + } + + #[cfg(not(windows))] + { + let _ = data_dir; // suppress unused warning + Ok(Box::new(KeychainKeyProvider::new()?)) + } +} + +/// Create the platform-appropriate JWT secret provider. +/// +/// - **Windows**: Uses DPAPI file-based storage. +/// - **macOS/Linux**: Uses the OS keychain. +pub fn create_jwt_secret_provider( + data_dir: &std::path::Path, +) -> anyhow::Result> { + #[cfg(windows)] + { + Ok(Box::new(DpapiJwtSecretProvider::new(data_dir)?)) + } + + #[cfg(not(windows))] + { + let _ = data_dir; + Ok(Box::new(KeychainJwtSecretProvider::new()?)) + } +} diff --git a/crates/mcpmux-storage/src/migrations/001_initial.sql b/crates/mcpmux-storage/src/migrations/001_initial.sql index 508b76b8..6b747e57 100644 --- a/crates/mcpmux-storage/src/migrations/001_initial.sql +++ b/crates/mcpmux-storage/src/migrations/001_initial.sql @@ -71,40 +71,52 @@ CREATE TABLE IF NOT EXISTS installed_servers ( -- OAuth connection state (persistent) oauth_connected INTEGER NOT NULL DEFAULT 0, - + + -- Installation source tracking + source TEXT NOT NULL DEFAULT 'registry', -- 'registry', 'user_config:/path/to/file.json', 'manual_entry' + created_at TEXT NOT NULL, updated_at TEXT NOT NULL, - + UNIQUE(space_id, server_id), FOREIGN KEY (space_id) REFERENCES spaces(id) ON DELETE CASCADE ); CREATE INDEX IF NOT EXISTS idx_installed_servers_space ON installed_servers(space_id); CREATE INDEX IF NOT EXISTS idx_installed_servers_enabled ON installed_servers(space_id, enabled); +CREATE INDEX IF NOT EXISTS idx_installed_servers_source ON installed_servers(source); -- ============================================================================ --- CREDENTIALS (Per-Space) --- OAuth tokens, API keys per (space, server) - TOKENS ONLY, not client registration +-- CREDENTIALS (Per-Space, Typed Rows) +-- Each token/key is a separate row with its own type and expiry. +-- One row per (space, server, credential_type). -- ============================================================================ CREATE TABLE IF NOT EXISTS credentials ( id TEXT PRIMARY KEY, space_id TEXT NOT NULL, server_id TEXT NOT NULL, - credential_type TEXT NOT NULL, -- 'oauth', 'api_key', 'basic_auth' - - -- Encrypted credential data (AES-256-GCM) + credential_type TEXT NOT NULL, -- 'access_token', 'refresh_token', 'api_key', 'basic_auth_user', 'basic_auth_pass' + + -- Only the secret value is encrypted (AES-256-GCM). Not a JSON blob. credential_value TEXT NOT NULL, - + + -- Metadata stored as plaintext for queryability + expires_at TEXT, -- RFC3339, nullable (refresh tokens / API keys may not expire) + token_type TEXT, -- 'Bearer', etc. (only for access_token) + scope TEXT, -- OAuth scope (only for access_token) + last_used_at TEXT, created_at TEXT NOT NULL, updated_at TEXT NOT NULL, - - UNIQUE(space_id, server_id), + + UNIQUE(space_id, server_id, credential_type), FOREIGN KEY (space_id) REFERENCES spaces(id) ON DELETE CASCADE ); CREATE INDEX IF NOT EXISTS idx_credentials_space_server ON credentials(space_id, server_id); +CREATE INDEX IF NOT EXISTS idx_credentials_type ON credentials(space_id, server_id, credential_type); +CREATE INDEX IF NOT EXISTS idx_credentials_expiry ON credentials(credential_type, expires_at); -- ============================================================================ -- OUTBOUND OAUTH CLIENTS @@ -283,18 +295,22 @@ CREATE TABLE IF NOT EXISTS inbound_clients ( -- Permissions: JSON map of space_id -> [feature_set_ids] grants TEXT, - + + -- Approval status (user must explicitly approve clients) + approved INTEGER NOT NULL DEFAULT 0, + -- Timestamps last_seen TEXT, created_at TEXT NOT NULL, updated_at TEXT NOT NULL, - + FOREIGN KEY (locked_space_id) REFERENCES spaces(id) ON DELETE SET NULL ); CREATE INDEX IF NOT EXISTS idx_inbound_clients_type ON inbound_clients(registration_type); CREATE INDEX IF NOT EXISTS idx_inbound_clients_name ON inbound_clients(client_name); CREATE INDEX IF NOT EXISTS idx_inbound_clients_metadata_url ON inbound_clients(metadata_url) WHERE metadata_url IS NOT NULL; +CREATE INDEX IF NOT EXISTS idx_inbound_clients_approved ON inbound_clients(approved) WHERE approved = 1; -- ============================================================================ -- CLIENT GRANTS (Per-Space permissions for INBOUND clients) @@ -365,3 +381,19 @@ CREATE TABLE IF NOT EXISTS settings ( value TEXT NOT NULL, updated_at TEXT NOT NULL ); + +-- ============================================================================ +-- APP SETTINGS +-- Key-value store for application-wide settings. +-- ============================================================================ + +CREATE TABLE IF NOT EXISTS app_settings ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL, + updated_at TEXT NOT NULL +); + +-- Seed default settings +INSERT OR IGNORE INTO app_settings (key, value, updated_at) +VALUES + ('gateway.auto_start', 'true', datetime('now')); diff --git a/crates/mcpmux-storage/src/migrations/002_add_client_approved.sql b/crates/mcpmux-storage/src/migrations/002_add_client_approved.sql deleted file mode 100644 index d5324d76..00000000 --- a/crates/mcpmux-storage/src/migrations/002_add_client_approved.sql +++ /dev/null @@ -1,10 +0,0 @@ --- Migration 002: Add approved flag to inbound_clients --- --- This tracks whether a client has been explicitly approved by the user. --- DCR creates the client entry, but approval happens separately. --- Silent re-authentication only works for approved clients. - -ALTER TABLE inbound_clients ADD COLUMN approved INTEGER NOT NULL DEFAULT 0; - --- Index for quick approved client lookups -CREATE INDEX IF NOT EXISTS idx_inbound_clients_approved ON inbound_clients(approved) WHERE approved = 1; diff --git a/crates/mcpmux-storage/src/migrations/003_app_settings.sql b/crates/mcpmux-storage/src/migrations/003_app_settings.sql deleted file mode 100644 index a394429f..00000000 --- a/crates/mcpmux-storage/src/migrations/003_app_settings.sql +++ /dev/null @@ -1,21 +0,0 @@ --- App Settings Table --- --- Key-value store for application-wide settings. --- Replaces scattered config files (gateway-port.txt, etc.) with a unified store. --- --- Examples: --- - gateway.port = "45818" --- - gateway.auto_start = "true" --- - ui.theme = "dark" --- - ui.window_state = '{"x":100,"y":100,"width":800,"height":600}' - -CREATE TABLE IF NOT EXISTS app_settings ( - key TEXT PRIMARY KEY, - value TEXT NOT NULL, - updated_at TEXT NOT NULL -); - --- Seed default settings -INSERT OR IGNORE INTO app_settings (key, value, updated_at) -VALUES - ('gateway.auto_start', 'true', datetime('now')); diff --git a/crates/mcpmux-storage/src/migrations/004_add_installation_source.sql b/crates/mcpmux-storage/src/migrations/004_add_installation_source.sql deleted file mode 100644 index 7188b6cc..00000000 --- a/crates/mcpmux-storage/src/migrations/004_add_installation_source.sql +++ /dev/null @@ -1,12 +0,0 @@ --- Add source column to installed_servers table --- Tracks how the server was installed: registry, user_config, or manual_entry --- --- Values: --- - 'registry': Installed from the server registry (default) --- - 'user_config:/path/to/file.json': From a user config file --- - 'manual_entry': Manually added via UI - -ALTER TABLE installed_servers ADD COLUMN source TEXT NOT NULL DEFAULT 'registry'; - --- Index for querying servers by source (useful for listing all servers from a config file) -CREATE INDEX IF NOT EXISTS idx_installed_servers_source ON installed_servers(source); diff --git a/crates/mcpmux-storage/src/repositories/credential_repository.rs b/crates/mcpmux-storage/src/repositories/credential_repository.rs index df6d70f1..5b247fa4 100644 --- a/crates/mcpmux-storage/src/repositories/credential_repository.rs +++ b/crates/mcpmux-storage/src/repositories/credential_repository.rs @@ -1,11 +1,14 @@ -//! SQLite implementation of CredentialRepository with encryption. +//! SQLite implementation of CredentialRepository with typed rows and encryption. +//! +//! Each credential is stored as a separate row per (space, server, type). +//! Only the secret value is encrypted — metadata (type, expiry, scope) is plaintext. use std::sync::Arc; -use anyhow::{Context, Result}; +use anyhow::Result; use async_trait::async_trait; use chrono::{DateTime, Utc}; -use mcpmux_core::{Credential, CredentialRepository, CredentialValue}; +use mcpmux_core::{Credential, CredentialRepository, CredentialType}; use rusqlite::{params, OptionalExtension}; use tokio::sync::Mutex; use uuid::Uuid; @@ -13,10 +16,24 @@ use uuid::Uuid; use crate::crypto::FieldEncryptor; use crate::Database; +/// Raw row data extracted from SQLite before decryption. +struct RawCredentialRow { + space_id: String, + server_id: String, + credential_type: String, + credential_value: String, // Encrypted + expires_at: Option, + token_type: Option, + scope: Option, + last_used_at: Option, + created_at: String, + updated_at: String, +} + /// SQLite-backed credential repository with field-level encryption. /// -/// Sensitive fields (tokens, keys) are encrypted using AES-256-GCM -/// before being stored in the database. +/// Only the secret value (token, key, password) is encrypted using AES-256-GCM. +/// Metadata fields (type, expiry, scope) are stored as plaintext for queryability. pub struct SqliteCredentialRepository { db: Arc>, encryptor: Arc, @@ -28,25 +45,25 @@ impl SqliteCredentialRepository { Self { db, encryptor } } - /// Encrypt a credential value to JSON. - fn encrypt_value(&self, value: &CredentialValue) -> Result { - let json = serde_json::to_string(value)?; - self.encryptor.encrypt(&json) + /// Encrypt a credential value for storage. + fn encrypt_value(&self, value: &str) -> Result { + self.encryptor + .encrypt(value) + .map_err(|e| anyhow::anyhow!("Failed to encrypt credential value: {}", e)) } - /// Decrypt a credential value from encrypted JSON. - fn decrypt_value(&self, encrypted: &str) -> Result { - let json = self.encryptor.decrypt(encrypted)?; - serde_json::from_str(&json).context("Failed to parse credential value") + /// Decrypt a credential value from storage. + fn decrypt_value(&self, encrypted: &str) -> Result { + self.encryptor + .decrypt(encrypted) + .map_err(|e| anyhow::anyhow!("Failed to decrypt credential value: {}", e)) } /// Parse a datetime string to DateTime. fn parse_datetime(s: &str) -> DateTime { - // Try RFC3339 first if let Ok(dt) = DateTime::parse_from_rfc3339(s) { return dt.with_timezone(&Utc); } - // Try SQLite datetime format if let Ok(dt) = chrono::NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S") { return dt.and_utc(); } @@ -57,49 +74,93 @@ impl SqliteCredentialRepository { fn parse_optional_datetime(s: Option) -> Option> { s.map(|dt| Self::parse_datetime(&dt)) } + + /// Standard column list for SELECT queries. + const SELECT_COLUMNS: &'static str = + "space_id, server_id, credential_type, credential_value, expires_at, token_type, scope, last_used_at, created_at, updated_at"; + + /// Extract raw row data from a rusqlite Row. + fn extract_row(row: &rusqlite::Row) -> rusqlite::Result { + Ok(RawCredentialRow { + space_id: row.get(0)?, + server_id: row.get(1)?, + credential_type: row.get(2)?, + credential_value: row.get(3)?, + expires_at: row.get(4)?, + token_type: row.get(5)?, + scope: row.get(6)?, + last_used_at: row.get(7)?, + created_at: row.get(8)?, + updated_at: row.get(9)?, + }) + } + + /// Build a Credential from extracted row data (needs &self for decryption). + fn build_credential(&self, row: RawCredentialRow) -> Result { + let value = self.decrypt_value(&row.credential_value)?; + let credential_type = CredentialType::parse(&row.credential_type) + .ok_or_else(|| anyhow::anyhow!("Unknown credential type: {}", row.credential_type))?; + + Ok(Credential { + space_id: row.space_id.parse().unwrap_or_else(|_| Uuid::new_v4()), + server_id: row.server_id, + credential_type, + value, + expires_at: Self::parse_optional_datetime(row.expires_at), + token_type: row.token_type, + scope: row.scope, + created_at: Self::parse_datetime(&row.created_at), + updated_at: Self::parse_datetime(&row.updated_at), + last_used: Self::parse_optional_datetime(row.last_used_at), + }) + } } #[async_trait] impl CredentialRepository for SqliteCredentialRepository { - async fn get(&self, space_id: &Uuid, server_id: &str) -> Result> { + async fn get( + &self, + space_id: &Uuid, + server_id: &str, + credential_type: &CredentialType, + ) -> Result> { let db = self.db.lock().await; let conn = db.connection(); - let mut stmt = conn.prepare( - "SELECT space_id, server_id, credential_value, created_at, updated_at, last_used_at - FROM credentials - WHERE space_id = ? AND server_id = ?", - )?; + let mut stmt = conn.prepare(&format!( + "SELECT {} FROM credentials WHERE space_id = ?1 AND server_id = ?2 AND credential_type = ?3", + Self::SELECT_COLUMNS + ))?; let row = stmt - .query_row(params![space_id.to_string(), server_id], |row| { - Ok(( - row.get::<_, String>(0)?, - row.get::<_, String>(1)?, - row.get::<_, String>(2)?, - row.get::<_, String>(3)?, - row.get::<_, String>(4)?, - row.get::<_, Option>(5)?, - )) - }) + .query_row( + params![space_id.to_string(), server_id, credential_type.as_str()], + Self::extract_row, + ) .optional()?; match row { - Some((space_id_str, server_id, encrypted_value, created_at, updated_at, last_used)) => { - let value = self.decrypt_value(&encrypted_value)?; - Ok(Some(Credential { - space_id: space_id_str.parse().unwrap_or_else(|_| Uuid::new_v4()), - server_id, - value, - created_at: Self::parse_datetime(&created_at), - updated_at: Self::parse_datetime(&updated_at), - last_used: Self::parse_optional_datetime(last_used), - })) - } + Some(raw) => Ok(Some(self.build_credential(raw)?)), None => Ok(None), } } + async fn get_all(&self, space_id: &Uuid, server_id: &str) -> Result> { + let db = self.db.lock().await; + let conn = db.connection(); + + let mut stmt = conn.prepare(&format!( + "SELECT {} FROM credentials WHERE space_id = ?1 AND server_id = ?2 ORDER BY credential_type", + Self::SELECT_COLUMNS + ))?; + + let rows: Vec<_> = stmt + .query_map(params![space_id.to_string(), server_id], Self::extract_row)? + .collect::, _>>()?; + + rows.into_iter().map(|r| self.build_credential(r)).collect() + } + async fn save(&self, credential: &Credential) -> Result<()> { let db = self.db.lock().await; let conn = db.connection(); @@ -107,34 +168,56 @@ impl CredentialRepository for SqliteCredentialRepository { let encrypted_value = self.encrypt_value(&credential.value)?; conn.execute( - "INSERT INTO credentials (id, space_id, server_id, credential_type, credential_value, created_at, updated_at, last_used_at) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8) - ON CONFLICT(space_id, server_id) DO UPDATE SET - credential_type = excluded.credential_type, + "INSERT INTO credentials (id, space_id, server_id, credential_type, credential_value, expires_at, token_type, scope, last_used_at, created_at, updated_at) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11) + ON CONFLICT(space_id, server_id, credential_type) DO UPDATE SET credential_value = excluded.credential_value, + expires_at = excluded.expires_at, + token_type = excluded.token_type, + scope = excluded.scope, updated_at = excluded.updated_at, last_used_at = excluded.last_used_at", params![ Uuid::new_v4().to_string(), credential.space_id.to_string(), credential.server_id, - credential_type_name(&credential.value), + credential.credential_type.as_str(), encrypted_value, + credential.expires_at.map(|dt| dt.to_rfc3339()), + credential.token_type, + credential.scope, + credential.last_used.map(|dt| dt.to_rfc3339()), credential.created_at.to_rfc3339(), credential.updated_at.to_rfc3339(), - credential.last_used.map(|dt| dt.to_rfc3339()), ], )?; Ok(()) } - async fn delete(&self, space_id: &Uuid, server_id: &str) -> Result<()> { + async fn delete( + &self, + space_id: &Uuid, + server_id: &str, + credential_type: &CredentialType, + ) -> Result<()> { + let db = self.db.lock().await; + let conn = db.connection(); + + conn.execute( + "DELETE FROM credentials WHERE space_id = ?1 AND server_id = ?2 AND credential_type = ?3", + params![space_id.to_string(), server_id, credential_type.as_str()], + )?; + + Ok(()) + } + + async fn delete_all(&self, space_id: &Uuid, server_id: &str) -> Result<()> { let db = self.db.lock().await; let conn = db.connection(); conn.execute( - "DELETE FROM credentials WHERE space_id = ? AND server_id = ?", + "DELETE FROM credentials WHERE space_id = ?1 AND server_id = ?2", params![space_id.to_string(), server_id], )?; @@ -142,70 +225,32 @@ impl CredentialRepository for SqliteCredentialRepository { } async fn clear_tokens(&self, space_id: &Uuid, server_id: &str) -> Result { - // For OAuth, client registration is in oauth_clients table, so just delete tokens - let existing = self.get(space_id, server_id).await?; - - match existing { - Some(credential) if credential.is_oauth() => { - // Delete OAuth tokens - client registration is preserved in oauth_clients - self.delete(space_id, server_id).await?; - Ok(true) - } - Some(_) => { - // Non-OAuth credentials (API keys) - don't clear on logout - Ok(false) - } - None => Ok(false), - } + let db = self.db.lock().await; + let conn = db.connection(); + + // Delete only OAuth tokens (access_token + refresh_token), preserve API keys etc. + let deleted = conn.execute( + "DELETE FROM credentials WHERE space_id = ?1 AND server_id = ?2 AND credential_type IN ('access_token', 'refresh_token')", + params![space_id.to_string(), server_id], + )?; + + Ok(deleted > 0) } async fn list_for_space(&self, space_id: &Uuid) -> Result> { let db = self.db.lock().await; let conn = db.connection(); - let mut stmt = conn.prepare( - "SELECT space_id, server_id, credential_value, created_at, updated_at, last_used_at - FROM credentials - WHERE space_id = ? - ORDER BY server_id", - )?; - - let rows = stmt.query_map(params![space_id.to_string()], |row| { - Ok(( - row.get::<_, String>(0)?, - row.get::<_, String>(1)?, - row.get::<_, String>(2)?, - row.get::<_, String>(3)?, - row.get::<_, String>(4)?, - row.get::<_, Option>(5)?, - )) - })?; - - let mut credentials = Vec::new(); - for row in rows { - let (space_id_str, server_id, encrypted_value, created_at, updated_at, last_used) = - row?; - let value = self.decrypt_value(&encrypted_value)?; - credentials.push(Credential { - space_id: space_id_str.parse().unwrap_or_else(|_| Uuid::new_v4()), - server_id, - value, - created_at: Self::parse_datetime(&created_at), - updated_at: Self::parse_datetime(&updated_at), - last_used: Self::parse_optional_datetime(last_used), - }); - } + let mut stmt = conn.prepare(&format!( + "SELECT {} FROM credentials WHERE space_id = ?1 ORDER BY server_id, credential_type", + Self::SELECT_COLUMNS + ))?; - Ok(credentials) - } -} + let rows: Vec<_> = stmt + .query_map(params![space_id.to_string()], Self::extract_row)? + .collect::, _>>()?; -/// Get the type name for a credential value. -fn credential_type_name(value: &CredentialValue) -> &'static str { - match value { - CredentialValue::ApiKey { .. } => "api_key", - CredentialValue::OAuth { .. } => "oauth", - CredentialValue::BasicAuth { .. } => "basic_auth", + rows.into_iter().map(|r| self.build_credential(r)).collect() } } @@ -224,7 +269,6 @@ mod tests { #[tokio::test] async fn test_credential_crud() { - // Create in-memory database and encryptor let db = Arc::new(Mutex::new(Database::open_in_memory().unwrap())); let key = crate::crypto::generate_master_key().unwrap(); let encryptor = Arc::new(FieldEncryptor::new(&key).unwrap()); @@ -238,33 +282,40 @@ mod tests { repo.save(&cred).await.unwrap(); // Retrieve - let found = repo.get(&space_id, "github").await.unwrap(); + let found = repo + .get(&space_id, "github", &CredentialType::ApiKey) + .await + .unwrap(); assert!(found.is_some()); let found = found.unwrap(); assert_eq!(found.server_id, "github"); - match found.value { - CredentialValue::ApiKey { key } => assert_eq!(key, "ghp_test_token_12345"), - _ => panic!("Wrong credential type"), - } + assert_eq!(found.credential_type, CredentialType::ApiKey); + assert_eq!(found.value, "ghp_test_token_12345"); // Update let updated_cred = Credential::api_key(space_id, "github", "ghp_new_token"); repo.save(&updated_cred).await.unwrap(); - let found = repo.get(&space_id, "github").await.unwrap().unwrap(); - match found.value { - CredentialValue::ApiKey { key } => assert_eq!(key, "ghp_new_token"), - _ => panic!("Wrong credential type"), - } + let found = repo + .get(&space_id, "github", &CredentialType::ApiKey) + .await + .unwrap() + .unwrap(); + assert_eq!(found.value, "ghp_new_token"); // Delete - repo.delete(&space_id, "github").await.unwrap(); - let found = repo.get(&space_id, "github").await.unwrap(); + repo.delete(&space_id, "github", &CredentialType::ApiKey) + .await + .unwrap(); + let found = repo + .get(&space_id, "github", &CredentialType::ApiKey) + .await + .unwrap(); assert!(found.is_none()); } #[tokio::test] - async fn test_oauth_credential() { + async fn test_access_token_credential() { let db = Arc::new(Mutex::new(Database::open_in_memory().unwrap())); let key = crate::crypto::generate_master_key().unwrap(); let encryptor = Arc::new(FieldEncryptor::new(&key).unwrap()); @@ -274,27 +325,115 @@ mod tests { create_test_space(&db, &space_id).await; let expires = Utc::now() + chrono::Duration::hours(1); - let cred = Credential::oauth( + let cred = + Credential::access_token(space_id, "atlassian", "access_token_xyz", Some(expires)); + repo.save(&cred).await.unwrap(); + + let found = repo + .get(&space_id, "atlassian", &CredentialType::AccessToken) + .await + .unwrap() + .unwrap(); + + assert_eq!(found.credential_type, CredentialType::AccessToken); + assert_eq!(found.value, "access_token_xyz"); + assert_eq!(found.token_type, Some("Bearer".to_string())); + assert!(!found.is_expired()); + } + + #[tokio::test] + async fn test_separate_access_and_refresh_tokens() { + let db = Arc::new(Mutex::new(Database::open_in_memory().unwrap())); + let key = crate::crypto::generate_master_key().unwrap(); + let encryptor = Arc::new(FieldEncryptor::new(&key).unwrap()); + let repo = SqliteCredentialRepository::new(db.clone(), encryptor); + + let space_id = Uuid::new_v4(); + create_test_space(&db, &space_id).await; + + // Save access token and refresh token as separate rows + let access = Credential::access_token( space_id, "atlassian", - "access_token_xyz", - Some("refresh_token_abc".to_string()), - Some(expires), + "access_xyz", + Some(Utc::now() + chrono::Duration::hours(1)), ); - repo.save(&cred).await.unwrap(); + let refresh = Credential::refresh_token(space_id, "atlassian", "refresh_abc", None); - let found = repo.get(&space_id, "atlassian").await.unwrap().unwrap(); - match found.value { - CredentialValue::OAuth { - access_token, - refresh_token, - .. - } => { - assert_eq!(access_token, "access_token_xyz"); - assert_eq!(refresh_token, Some("refresh_token_abc".to_string())); - } - _ => panic!("Wrong credential type"), - } + repo.save(&access).await.unwrap(); + repo.save(&refresh).await.unwrap(); + + // Get all for server — should return 2 rows + let all = repo.get_all(&space_id, "atlassian").await.unwrap(); + assert_eq!(all.len(), 2); + + // Get each individually + let found_access = repo + .get(&space_id, "atlassian", &CredentialType::AccessToken) + .await + .unwrap() + .unwrap(); + assert_eq!(found_access.value, "access_xyz"); + + let found_refresh = repo + .get(&space_id, "atlassian", &CredentialType::RefreshToken) + .await + .unwrap() + .unwrap(); + assert_eq!(found_refresh.value, "refresh_abc"); + } + + #[tokio::test] + async fn test_clear_tokens_only_removes_oauth() { + let db = Arc::new(Mutex::new(Database::open_in_memory().unwrap())); + let key = crate::crypto::generate_master_key().unwrap(); + let encryptor = Arc::new(FieldEncryptor::new(&key).unwrap()); + let repo = SqliteCredentialRepository::new(db.clone(), encryptor); + + let space_id = Uuid::new_v4(); + create_test_space(&db, &space_id).await; + + // Save access_token, refresh_token, and api_key for same server + let access = Credential::access_token(space_id, "server", "access", None); + let refresh = Credential::refresh_token(space_id, "server", "refresh", None); + let api_key = Credential::api_key(space_id, "server", "key123"); + + repo.save(&access).await.unwrap(); + repo.save(&refresh).await.unwrap(); + repo.save(&api_key).await.unwrap(); + + // clear_tokens should remove access + refresh but keep api_key + let cleared = repo.clear_tokens(&space_id, "server").await.unwrap(); + assert!(cleared); + + let all = repo.get_all(&space_id, "server").await.unwrap(); + assert_eq!(all.len(), 1); + assert_eq!(all[0].credential_type, CredentialType::ApiKey); + } + + #[tokio::test] + async fn test_delete_all() { + let db = Arc::new(Mutex::new(Database::open_in_memory().unwrap())); + let key = crate::crypto::generate_master_key().unwrap(); + let encryptor = Arc::new(FieldEncryptor::new(&key).unwrap()); + let repo = SqliteCredentialRepository::new(db.clone(), encryptor); + + let space_id = Uuid::new_v4(); + create_test_space(&db, &space_id).await; + + repo.save(&Credential::access_token( + space_id, "server", "access", None, + )) + .await + .unwrap(); + repo.save(&Credential::api_key(space_id, "server", "key")) + .await + .unwrap(); + + repo.delete_all(&space_id, "server").await.unwrap(); + + let all = repo.get_all(&space_id, "server").await.unwrap(); + assert!(all.is_empty()); } #[tokio::test] @@ -322,13 +461,8 @@ mod tests { .await .unwrap(); - // List space1 - let creds = repo.list_for_space(&space1).await.unwrap(); - assert_eq!(creds.len(), 2); - - // List space2 - let creds = repo.list_for_space(&space2).await.unwrap(); - assert_eq!(creds.len(), 1); + assert_eq!(repo.list_for_space(&space1).await.unwrap().len(), 2); + assert_eq!(repo.list_for_space(&space2).await.unwrap().len(), 1); } #[tokio::test] @@ -362,5 +496,16 @@ mod tests { // Raw value should be hex-encoded (encrypted) assert!(hex::decode(&raw_value).is_ok()); + + // But expires_at and token_type should be plaintext (queryable) + let (cred_type, expires_at): (String, Option) = conn + .query_row( + "SELECT credential_type, expires_at FROM credentials WHERE server_id = 'test'", + [], + |row| Ok((row.get(0)?, row.get(1)?)), + ) + .unwrap(); + assert_eq!(cred_type, "api_key"); + assert!(expires_at.is_none()); // API keys don't expire } } diff --git a/crates/mcpmux-storage/src/repositories/installed_server_repository.rs b/crates/mcpmux-storage/src/repositories/installed_server_repository.rs index 63e60e9e..3ddef0cc 100644 --- a/crates/mcpmux-storage/src/repositories/installed_server_repository.rs +++ b/crates/mcpmux-storage/src/repositories/installed_server_repository.rs @@ -12,17 +12,58 @@ use rusqlite::{params, OptionalExtension}; use tokio::sync::Mutex; use uuid::Uuid; -use crate::Database; +use crate::{crypto::FieldEncryptor, Database}; + +/// Raw row data extracted from SQLite before decryption. +struct RawServerRow { + id: String, + space_id: String, + server_id: String, + server_name: Option, + cached_definition: Option, + input_values: Option, + enabled: bool, + env_overrides: Option, + args_append: Option, + extra_headers: Option, + oauth_connected: bool, + created_at: String, + updated_at: String, + source: Option, +} /// SQLite-backed implementation of InstalledServerRepository. pub struct SqliteInstalledServerRepository { db: Arc>, + encryptor: Arc, } impl SqliteInstalledServerRepository { /// Create a new SQLite installed server repository. - pub fn new(db: Arc>) -> Self { - Self { db } + pub fn new(db: Arc>, encryptor: Arc) -> Self { + Self { db, encryptor } + } + + /// Encrypt input values for storage. + fn encrypt_input_values(&self, values: &HashMap) -> Result { + let json = Self::serialize_json_map(values); + self.encryptor + .encrypt(&json) + .map_err(|e| anyhow::anyhow!("Failed to encrypt input values: {}", e)) + } + + /// Decrypt input values from storage. + /// Falls back to plaintext JSON for backward compatibility with unencrypted data. + fn decrypt_input_values(&self, stored: Option) -> HashMap { + let Some(data) = stored else { + return HashMap::new(); + }; + // Try decrypting first (new encrypted format) + if let Ok(json) = self.encryptor.decrypt(&data) { + return serde_json::from_str(&json).unwrap_or_default(); + } + // Fallback: try parsing as plaintext JSON (backward compat) + serde_json::from_str(&data).unwrap_or_default() } /// Parse a datetime string to DateTime. @@ -92,40 +133,45 @@ impl SqliteInstalledServerRepository { "id, space_id, server_id, server_name, cached_definition, input_values, enabled, env_overrides, args_append, extra_headers, oauth_connected, created_at, updated_at, source"; - /// Map a row to InstalledServer - fn map_row(row: &rusqlite::Row) -> rusqlite::Result { - let id: String = row.get(0)?; - let space_id: String = row.get(1)?; - let server_id: String = row.get(2)?; - let server_name: Option = row.get(3)?; - let cached_definition: Option = row.get(4)?; - let input_values: Option = row.get(5)?; - let enabled: bool = row.get(6)?; - let env_overrides: Option = row.get(7)?; - let args_append: Option = row.get(8)?; - let extra_headers: Option = row.get(9)?; - let oauth_connected: bool = row.get(10)?; - let created_at: String = row.get(11)?; - let updated_at: String = row.get(12)?; - let source: Option = row.get(13)?; - - Ok(InstalledServer { - id: Uuid::parse_str(&id).unwrap_or_else(|_| Uuid::new_v4()), - space_id, - server_id, - server_name, - cached_definition, - input_values: Self::parse_json_map(input_values), - enabled, - env_overrides: Self::parse_json_map(env_overrides), - args_append: Self::parse_json_vec(args_append), - extra_headers: Self::parse_json_map(extra_headers), - oauth_connected, - source: Self::parse_source(source), - created_at: Self::parse_datetime(&created_at), - updated_at: Self::parse_datetime(&updated_at), + /// Extract raw row data (used in the closure passed to rusqlite). + fn extract_row(row: &rusqlite::Row) -> rusqlite::Result { + Ok(RawServerRow { + id: row.get(0)?, + space_id: row.get(1)?, + server_id: row.get(2)?, + server_name: row.get(3)?, + cached_definition: row.get(4)?, + input_values: row.get(5)?, + enabled: row.get(6)?, + env_overrides: row.get(7)?, + args_append: row.get(8)?, + extra_headers: row.get(9)?, + oauth_connected: row.get(10)?, + created_at: row.get(11)?, + updated_at: row.get(12)?, + source: row.get(13)?, }) } + + /// Build InstalledServer from extracted row data (needs &self for decryption). + fn build_server(&self, row: RawServerRow) -> InstalledServer { + InstalledServer { + id: Uuid::parse_str(&row.id).unwrap_or_else(|_| Uuid::new_v4()), + space_id: row.space_id, + server_id: row.server_id, + server_name: row.server_name, + cached_definition: row.cached_definition, + input_values: self.decrypt_input_values(row.input_values), + enabled: row.enabled, + env_overrides: Self::parse_json_map(row.env_overrides), + args_append: Self::parse_json_vec(row.args_append), + extra_headers: Self::parse_json_map(row.extra_headers), + oauth_connected: row.oauth_connected, + source: Self::parse_source(row.source), + created_at: Self::parse_datetime(&row.created_at), + updated_at: Self::parse_datetime(&row.updated_at), + } + } } #[async_trait] @@ -139,11 +185,11 @@ impl InstalledServerRepository for SqliteInstalledServerRepository { Self::SELECT_COLUMNS ))?; - let servers = stmt - .query_map([], Self::map_row)? + let rows: Vec<_> = stmt + .query_map([], Self::extract_row)? .collect::, _>>()?; - Ok(servers) + Ok(rows.into_iter().map(|r| self.build_server(r)).collect()) } async fn list_for_space(&self, space_id: &str) -> Result> { @@ -155,11 +201,11 @@ impl InstalledServerRepository for SqliteInstalledServerRepository { Self::SELECT_COLUMNS ))?; - let servers = stmt - .query_map([space_id], Self::map_row)? + let rows: Vec<_> = stmt + .query_map([space_id], Self::extract_row)? .collect::, _>>()?; - Ok(servers) + Ok(rows.into_iter().map(|r| self.build_server(r)).collect()) } async fn list_by_source_file( @@ -177,11 +223,11 @@ impl InstalledServerRepository for SqliteInstalledServerRepository { Self::SELECT_COLUMNS ))?; - let servers = stmt - .query_map([&source_prefix], Self::map_row)? + let rows: Vec<_> = stmt + .query_map([&source_prefix], Self::extract_row)? .collect::, _>>()?; - Ok(servers) + Ok(rows.into_iter().map(|r| self.build_server(r)).collect()) } async fn get(&self, id: &Uuid) -> Result> { @@ -193,9 +239,11 @@ impl InstalledServerRepository for SqliteInstalledServerRepository { Self::SELECT_COLUMNS ))?; - let server = stmt.query_row([id.to_string()], Self::map_row).optional()?; + let row = stmt + .query_row([id.to_string()], Self::extract_row) + .optional()?; - Ok(server) + Ok(row.map(|r| self.build_server(r))) } async fn get_by_server_id( @@ -211,17 +259,19 @@ impl InstalledServerRepository for SqliteInstalledServerRepository { Self::SELECT_COLUMNS ))?; - let server = stmt - .query_row([space_id, server_id], Self::map_row) + let row = stmt + .query_row([space_id, server_id], Self::extract_row) .optional()?; - Ok(server) + Ok(row.map(|r| self.build_server(r))) } async fn install(&self, server: &InstalledServer) -> Result<()> { let db = self.db.lock().await; let conn = db.connection(); + let encrypted_inputs = self.encrypt_input_values(&server.input_values)?; + conn.execute( "INSERT INTO installed_servers (id, space_id, server_id, server_name, cached_definition, input_values, enabled, env_overrides, @@ -233,7 +283,7 @@ impl InstalledServerRepository for SqliteInstalledServerRepository { server.server_id, server.server_name, server.cached_definition, - Self::serialize_json_map(&server.input_values), + encrypted_inputs, server.enabled, Self::serialize_json_map(&server.env_overrides), Self::serialize_json_vec(&server.args_append), @@ -251,6 +301,8 @@ impl InstalledServerRepository for SqliteInstalledServerRepository { let db = self.db.lock().await; let conn = db.connection(); + let encrypted_inputs = self.encrypt_input_values(&server.input_values)?; + conn.execute( "UPDATE installed_servers SET server_name = ?2, cached_definition = ?3, input_values = ?4, enabled = ?5, @@ -261,7 +313,7 @@ impl InstalledServerRepository for SqliteInstalledServerRepository { server.id.to_string(), server.server_name, server.cached_definition, - Self::serialize_json_map(&server.input_values), + encrypted_inputs, server.enabled, Self::serialize_json_map(&server.env_overrides), Self::serialize_json_vec(&server.args_append), @@ -294,11 +346,11 @@ impl InstalledServerRepository for SqliteInstalledServerRepository { Self::SELECT_COLUMNS ))?; - let servers = stmt - .query_map([space_id], Self::map_row)? + let rows: Vec<_> = stmt + .query_map([space_id], Self::extract_row)? .collect::, _>>()?; - Ok(servers) + Ok(rows.into_iter().map(|r| self.build_server(r)).collect()) } async fn list_enabled_all(&self) -> Result> { @@ -310,11 +362,11 @@ impl InstalledServerRepository for SqliteInstalledServerRepository { Self::SELECT_COLUMNS ))?; - let servers = stmt - .query_map([], Self::map_row)? + let rows: Vec<_> = stmt + .query_map([], Self::extract_row)? .collect::, _>>()?; - Ok(servers) + Ok(rows.into_iter().map(|r| self.build_server(r)).collect()) } async fn set_enabled(&self, id: &Uuid, enabled: bool) -> Result<()> { @@ -347,18 +399,17 @@ impl InstalledServerRepository for SqliteInstalledServerRepository { let db = self.db.lock().await; let conn = db.connection(); - let input_json = Self::serialize_json_map(&input_values); + let encrypted_inputs = self.encrypt_input_values(&input_values)?; tracing::debug!( - "[InstalledServerRepo] Updating inputs for {}: {} values, JSON: {:?}", + "[InstalledServerRepo] Updating inputs for {}: {} values (encrypted)", id, input_values.len(), - input_json ); conn.execute( "UPDATE installed_servers SET input_values = ?2, updated_at = ?3 WHERE id = ?1", - params![id.to_string(), input_json, Utc::now().to_rfc3339()], + params![id.to_string(), encrypted_inputs, Utc::now().to_rfc3339()], )?; tracing::debug!( diff --git a/tests/rust/src/mocks.rs b/tests/rust/src/mocks.rs index 50b0d268..227cc13f 100644 --- a/tests/rust/src/mocks.rs +++ b/tests/rust/src/mocks.rs @@ -9,8 +9,8 @@ use uuid::Uuid; use mcpmux_core::{ domain::{ - Client, Credential, FeatureSet, FeatureSetMember, FeatureSetType, InstalledServer, - MemberMode, MemberType, OutboundOAuthRegistration, ServerFeature, Space, + Client, Credential, CredentialType, FeatureSet, FeatureSetMember, FeatureSetType, + InstalledServer, MemberMode, MemberType, OutboundOAuthRegistration, ServerFeature, Space, }, repository::{ AppSettingsRepository, CredentialRepository, FeatureSetRepository, @@ -685,7 +685,7 @@ impl InboundMcpClientRepository for MockInboundMcpClientRepository { #[derive(Default)] pub struct MockCredentialRepository { - credentials: RwLock>, + credentials: RwLock>, } impl MockCredentialRepository { @@ -694,54 +694,90 @@ impl MockCredentialRepository { } pub fn with_credential(self, cred: Credential) -> Self { - self.credentials - .write() - .unwrap() - .insert((cred.space_id, cred.server_id.clone()), cred); + self.credentials.write().unwrap().insert( + ( + cred.space_id, + cred.server_id.clone(), + cred.credential_type.as_str().to_string(), + ), + cred, + ); self } } #[async_trait] impl CredentialRepository for MockCredentialRepository { - async fn get(&self, space_id: &Uuid, server_id: &str) -> RepoResult> { + async fn get( + &self, + space_id: &Uuid, + server_id: &str, + credential_type: &CredentialType, + ) -> RepoResult> { Ok(self .credentials .read() .unwrap() - .get(&(*space_id, server_id.to_string())) + .get(&( + *space_id, + server_id.to_string(), + credential_type.as_str().to_string(), + )) .cloned()) } + async fn get_all(&self, space_id: &Uuid, server_id: &str) -> RepoResult> { + Ok(self + .credentials + .read() + .unwrap() + .values() + .filter(|c| c.space_id == *space_id && c.server_id == server_id) + .cloned() + .collect()) + } + async fn save(&self, credential: &Credential) -> RepoResult<()> { self.credentials.write().unwrap().insert( - (credential.space_id, credential.server_id.clone()), + ( + credential.space_id, + credential.server_id.clone(), + credential.credential_type.as_str().to_string(), + ), credential.clone(), ); Ok(()) } - async fn delete(&self, space_id: &Uuid, server_id: &str) -> RepoResult<()> { + async fn delete( + &self, + space_id: &Uuid, + server_id: &str, + credential_type: &CredentialType, + ) -> RepoResult<()> { + self.credentials.write().unwrap().remove(&( + *space_id, + server_id.to_string(), + credential_type.as_str().to_string(), + )); + Ok(()) + } + + async fn delete_all(&self, space_id: &Uuid, server_id: &str) -> RepoResult<()> { self.credentials .write() .unwrap() - .remove(&(*space_id, server_id.to_string())); + .retain(|k, _| !(k.0 == *space_id && k.1 == server_id)); Ok(()) } async fn clear_tokens(&self, space_id: &Uuid, server_id: &str) -> RepoResult { - // Clearing tokens means removing the credential entirely for simplicity in mock - // In production, this preserves client_id but clears access/refresh tokens - if self - .credentials - .write() - .unwrap() - .remove(&(*space_id, server_id.to_string())) - .is_some() - { - return Ok(true); - } - Ok(false) + let mut creds = self.credentials.write().unwrap(); + let before = creds.len(); + creds.retain(|_, c| { + !(c.space_id == *space_id && c.server_id == server_id && c.credential_type.is_oauth()) + }); + Ok(creds.len() < before) } async fn list_for_space(&self, space_id: &Uuid) -> RepoResult> { diff --git a/tests/rust/tests/database/installed_server.rs b/tests/rust/tests/database/installed_server.rs index 63245ee6..6f8b12c5 100644 --- a/tests/rust/tests/database/installed_server.rs +++ b/tests/rust/tests/database/installed_server.rs @@ -1,18 +1,25 @@ //! InstalledServerRepository integration tests use mcpmux_core::repository::{InstalledServerRepository, SpaceRepository}; -use mcpmux_storage::{SqliteInstalledServerRepository, SqliteSpaceRepository}; +use mcpmux_storage::{ + generate_master_key, FieldEncryptor, SqliteInstalledServerRepository, SqliteSpaceRepository, +}; use pretty_assertions::assert_eq; use std::collections::HashMap; use std::sync::Arc; use tests::{db::TestDatabase, fixtures}; use tokio::sync::Mutex; +fn test_encryptor() -> Arc { + let key = generate_master_key().expect("Failed to generate key"); + Arc::new(FieldEncryptor::new(&key).expect("Failed to create encryptor")) +} + #[tokio::test] async fn test_installed_server_install_and_get() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); // First create a space (needed for foreign key) @@ -41,7 +48,7 @@ async fn test_installed_server_install_and_get() { async fn test_installed_server_get_by_server_id() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); let space = fixtures::test_space("Test Space"); @@ -78,7 +85,7 @@ async fn test_installed_server_get_by_server_id() { async fn test_installed_server_list_for_space() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); // Create two spaces @@ -121,7 +128,7 @@ async fn test_installed_server_list_for_space() { async fn test_installed_server_uninstall() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); let space = fixtures::test_space("Test Space"); @@ -149,7 +156,7 @@ async fn test_installed_server_uninstall() { async fn test_installed_server_set_enabled() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); let space = fixtures::test_space("Test Space"); @@ -193,7 +200,7 @@ async fn test_installed_server_set_enabled() { async fn test_installed_server_list_enabled() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); let space = fixtures::test_space("Test Space"); @@ -227,7 +234,7 @@ async fn test_installed_server_list_enabled() { async fn test_installed_server_set_oauth_connected() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); let space = fixtures::test_space("Test Space"); @@ -261,7 +268,7 @@ async fn test_installed_server_set_oauth_connected() { async fn test_installed_server_update_inputs() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); let space = fixtures::test_space("Test Space"); @@ -301,7 +308,7 @@ async fn test_installed_server_update_inputs() { async fn test_installed_server_update_cached_definition() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); let space = fixtures::test_space("Test Space"); @@ -337,7 +344,7 @@ async fn test_installed_server_update_cached_definition() { async fn test_installed_server_list_all() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); // Create two spaces with servers @@ -366,7 +373,7 @@ async fn test_installed_server_list_all() { async fn test_installed_server_list_enabled_all() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); let space1 = fixtures::test_space("Space 1"); @@ -402,7 +409,7 @@ async fn test_installed_server_list_enabled_all() { async fn test_installed_server_env_overrides_persist() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); let space = fixtures::test_space("Test Space"); @@ -438,7 +445,7 @@ async fn test_installed_server_env_overrides_persist() { async fn test_installed_server_args_append_persist() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); let space = fixtures::test_space("Test Space"); @@ -471,7 +478,7 @@ async fn test_installed_server_args_append_persist() { async fn test_installed_server_extra_headers_persist() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); let space = fixtures::test_space("Test Space"); @@ -510,7 +517,7 @@ async fn test_installed_server_extra_headers_persist() { async fn test_installed_server_update_preserves_custom_fields() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); let space = fixtures::test_space("Test Space"); @@ -580,7 +587,7 @@ async fn test_installed_server_update_preserves_custom_fields() { async fn test_installed_server_empty_custom_fields_by_default() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); let space = fixtures::test_space("Test Space"); @@ -607,7 +614,7 @@ async fn test_installed_server_empty_custom_fields_by_default() { async fn test_installed_server_clear_custom_fields_via_update() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); let space = fixtures::test_space("Test Space"); @@ -670,7 +677,7 @@ async fn test_installed_server_clear_custom_fields_via_update() { async fn test_installed_server_special_characters_persist() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); - let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db)); + let server_repo = SqliteInstalledServerRepository::new(Arc::clone(&db), test_encryptor()); let space_repo = SqliteSpaceRepository::new(db); let space = fixtures::test_space("Test Space"); diff --git a/tests/rust/tests/database/outbound_oauth.rs b/tests/rust/tests/database/outbound_oauth.rs index 2ab41361..01c3969b 100644 --- a/tests/rust/tests/database/outbound_oauth.rs +++ b/tests/rust/tests/database/outbound_oauth.rs @@ -4,7 +4,7 @@ //! Handles OAuth client registrations (DCR with servers) and token storage. use chrono::{Duration, Utc}; -use mcpmux_core::domain::{Credential, CredentialValue, OutboundOAuthRegistration}; +use mcpmux_core::domain::{Credential, CredentialType, OutboundOAuthRegistration}; use mcpmux_core::repository::{CredentialRepository, OutboundOAuthRepository, SpaceRepository}; use mcpmux_storage::{ generate_master_key, FieldEncryptor, SqliteCredentialRepository, SqliteOutboundOAuthRepository, @@ -182,18 +182,20 @@ async fn test_registrations_isolated_by_space() { } // ============================================================================= -// CredentialRepository Tests +// CredentialRepository Tests (Typed Rows) // ============================================================================= -fn create_oauth_credential(space_id: Uuid, server_id: &str) -> Credential { +/// Helper: save access_token + refresh_token as separate rows +async fn save_oauth_credentials( + repo: &SqliteCredentialRepository, + space_id: Uuid, + server_id: &str, +) { let expires_at = Some(Utc::now() + Duration::hours(1)); - Credential::oauth( - space_id, - server_id, - "access_token_xyz", - Some("refresh_token_abc".to_string()), - expires_at, - ) + let access = Credential::access_token(space_id, server_id, "access_token_xyz", expires_at); + let refresh = Credential::refresh_token(space_id, server_id, "refresh_token_abc", None); + CredentialRepository::save(repo, &access).await.unwrap(); + CredentialRepository::save(repo, &refresh).await.unwrap(); } fn create_api_key_credential(space_id: Uuid, server_id: &str, api_key: &str) -> Credential { @@ -211,28 +213,36 @@ async fn test_save_and_get_credential() { let space = fixtures::test_space("Test Space"); SpaceRepository::create(&space_repo, &space).await.unwrap(); - let cred = create_oauth_credential(space.id, "server-oauth"); - CredentialRepository::save(&cred_repo, &cred) - .await - .expect("Failed to save"); + save_oauth_credentials(&cred_repo, space.id, "server-oauth").await; - let loaded = CredentialRepository::get(&cred_repo, &space.id, "server-oauth") - .await - .expect("Failed to get"); - assert!(loaded.is_some()); - let loaded = loaded.unwrap(); + // Load access_token + let loaded = CredentialRepository::get( + &cred_repo, + &space.id, + "server-oauth", + &CredentialType::AccessToken, + ) + .await + .expect("Failed to get") + .expect("Should find access token"); + + assert_eq!(loaded.value, "access_token_xyz"); + assert_eq!(loaded.credential_type, CredentialType::AccessToken); + assert_eq!(loaded.token_type, Some("Bearer".to_string())); + + // Load refresh_token + let refresh = CredentialRepository::get( + &cred_repo, + &space.id, + "server-oauth", + &CredentialType::RefreshToken, + ) + .await + .unwrap() + .expect("Should find refresh token"); - match &loaded.value { - CredentialValue::OAuth { - access_token, - refresh_token, - .. - } => { - assert_eq!(access_token, "access_token_xyz"); - assert_eq!(refresh_token.as_deref(), Some("refresh_token_abc")); - } - _ => panic!("Expected OAuth credential"), - } + assert_eq!(refresh.value, "refresh_token_abc"); + assert_eq!(refresh.credential_type, CredentialType::RefreshToken); } #[tokio::test] @@ -242,9 +252,14 @@ async fn test_credential_not_found() { let encryptor = test_encryptor(); let cred_repo = SqliteCredentialRepository::new(db, encryptor); - let loaded = CredentialRepository::get(&cred_repo, &Uuid::new_v4(), "nonexistent") - .await - .unwrap(); + let loaded = CredentialRepository::get( + &cred_repo, + &Uuid::new_v4(), + "nonexistent", + &CredentialType::AccessToken, + ) + .await + .unwrap(); assert!(loaded.is_none()); } @@ -262,17 +277,14 @@ async fn test_save_api_key_credential() { let cred = create_api_key_credential(space.id, "api-server", "my_secret_api_key"); CredentialRepository::save(&cred_repo, &cred).await.unwrap(); - let loaded = CredentialRepository::get(&cred_repo, &space.id, "api-server") - .await - .unwrap() - .unwrap(); + let loaded = + CredentialRepository::get(&cred_repo, &space.id, "api-server", &CredentialType::ApiKey) + .await + .unwrap() + .unwrap(); - match &loaded.value { - CredentialValue::ApiKey { key } => { - assert_eq!(key, "my_secret_api_key"); - } - _ => panic!("Expected ApiKey"), - } + assert_eq!(loaded.value, "my_secret_api_key"); + assert_eq!(loaded.credential_type, CredentialType::ApiKey); } #[tokio::test] @@ -286,30 +298,27 @@ async fn test_update_credential() { let space = fixtures::test_space("Test Space"); SpaceRepository::create(&space_repo, &space).await.unwrap(); - let mut cred = create_oauth_credential(space.id, "server-1"); + // Save initial access token + let cred = Credential::access_token(space.id, "server-1", "old_token", None); CredentialRepository::save(&cred_repo, &cred).await.unwrap(); - // Update with new tokens - cred.value = CredentialValue::OAuth { - access_token: "new_access_token".to_string(), - refresh_token: Some("new_refresh_token".to_string()), - expires_at: None, - token_type: "Bearer".to_string(), - scope: None, - }; - CredentialRepository::save(&cred_repo, &cred).await.unwrap(); - - let loaded = CredentialRepository::get(&cred_repo, &space.id, "server-1") + // Update with new token + let updated = Credential::access_token(space.id, "server-1", "new_access_token", None); + CredentialRepository::save(&cred_repo, &updated) .await - .unwrap() .unwrap(); - match &loaded.value { - CredentialValue::OAuth { access_token, .. } => { - assert_eq!(access_token, "new_access_token"); - } - _ => panic!("Expected OAuth credential"), - } + let loaded = CredentialRepository::get( + &cred_repo, + &space.id, + "server-1", + &CredentialType::AccessToken, + ) + .await + .unwrap() + .unwrap(); + + assert_eq!(loaded.value, "new_access_token"); } #[tokio::test] @@ -323,17 +332,43 @@ async fn test_delete_credential() { let space = fixtures::test_space("Test Space"); SpaceRepository::create(&space_repo, &space).await.unwrap(); - let cred = create_oauth_credential(space.id, "to-delete"); - CredentialRepository::save(&cred_repo, &cred).await.unwrap(); + save_oauth_credentials(&cred_repo, space.id, "to-delete").await; - CredentialRepository::delete(&cred_repo, &space.id, "to-delete") - .await - .unwrap(); - - let loaded = CredentialRepository::get(&cred_repo, &space.id, "to-delete") - .await - .unwrap(); - assert!(loaded.is_none()); + // Delete just the access token + CredentialRepository::delete( + &cred_repo, + &space.id, + "to-delete", + &CredentialType::AccessToken, + ) + .await + .unwrap(); + + // Access token gone, refresh token still there + let access = CredentialRepository::get( + &cred_repo, + &space.id, + "to-delete", + &CredentialType::AccessToken, + ) + .await + .unwrap(); + assert!(access.is_none()); + + let refresh = CredentialRepository::get( + &cred_repo, + &space.id, + "to-delete", + &CredentialType::RefreshToken, + ) + .await + .unwrap(); + assert!(refresh.is_some()); + + // Delete all remaining + cred_repo.delete_all(&space.id, "to-delete").await.unwrap(); + let all = cred_repo.get_all(&space.id, "to-delete").await.unwrap(); + assert!(all.is_empty()); } #[tokio::test] @@ -347,17 +382,12 @@ async fn test_list_credentials_for_space() { let space = fixtures::test_space("Test Space"); SpaceRepository::create(&space_repo, &space).await.unwrap(); - let cred1 = create_oauth_credential(space.id, "server-1"); - let cred2 = create_oauth_credential(space.id, "server-2"); - CredentialRepository::save(&cred_repo, &cred1) - .await - .unwrap(); - CredentialRepository::save(&cred_repo, &cred2) - .await - .unwrap(); + // 2 rows for server-1 (access + refresh), 2 rows for server-2 + save_oauth_credentials(&cred_repo, space.id, "server-1").await; + save_oauth_credentials(&cred_repo, space.id, "server-2").await; let list = cred_repo.list_for_space(&space.id).await.unwrap(); - assert_eq!(list.len(), 2); + assert_eq!(list.len(), 4); // 2 servers × 2 types each } #[tokio::test] @@ -377,19 +407,13 @@ async fn test_credentials_isolated_by_space() { .await .unwrap(); - let cred_a = create_oauth_credential(space_a.id, "shared-server"); - let cred_b = create_oauth_credential(space_b.id, "shared-server"); - CredentialRepository::save(&cred_repo, &cred_a) - .await - .unwrap(); - CredentialRepository::save(&cred_repo, &cred_b) - .await - .unwrap(); + save_oauth_credentials(&cred_repo, space_a.id, "shared-server").await; + save_oauth_credentials(&cred_repo, space_b.id, "shared-server").await; let list_a = cred_repo.list_for_space(&space_a.id).await.unwrap(); let list_b = cred_repo.list_for_space(&space_b.id).await.unwrap(); - assert_eq!(list_a.len(), 1); - assert_eq!(list_b.len(), 1); + assert_eq!(list_a.len(), 2); // access + refresh + assert_eq!(list_b.len(), 2); } // ============================================================================= @@ -407,29 +431,29 @@ async fn test_credential_expiration() { let space = fixtures::test_space("Test Space"); SpaceRepository::create(&space_repo, &space).await.unwrap(); - // Create credential that expires in the future + // Create access token that expires in the future let future_expiry = Some(Utc::now() + Duration::hours(2)); - let not_expired = Credential::oauth( - space.id, - "valid-server", - "access_token", - Some("refresh".to_string()), - future_expiry, - ); + let not_expired = + Credential::access_token(space.id, "valid-server", "access_token", future_expiry); CredentialRepository::save(&cred_repo, ¬_expired) .await .unwrap(); - let loaded = CredentialRepository::get(&cred_repo, &space.id, "valid-server") - .await - .unwrap() - .unwrap(); + let loaded = CredentialRepository::get( + &cred_repo, + &space.id, + "valid-server", + &CredentialType::AccessToken, + ) + .await + .unwrap() + .unwrap(); assert!(!loaded.is_expired()); - assert!(loaded.can_refresh()); + assert!(loaded.is_oauth()); } #[tokio::test] -async fn test_credential_without_refresh_token() { +async fn test_clear_tokens_preserves_api_keys() { let test_db = TestDatabase::new(); let db = Arc::new(Mutex::new(test_db.db)); let encryptor = test_encryptor(); @@ -439,17 +463,20 @@ async fn test_credential_without_refresh_token() { let space = fixtures::test_space("Test Space"); SpaceRepository::create(&space_repo, &space).await.unwrap(); - // Create credential without refresh token - let no_refresh = Credential::oauth(space.id, "no-refresh-server", "access_token", None, None); - CredentialRepository::save(&cred_repo, &no_refresh) + // Save access_token, refresh_token, and api_key + save_oauth_credentials(&cred_repo, space.id, "server").await; + let api_key = Credential::api_key(space.id, "server", "key123"); + CredentialRepository::save(&cred_repo, &api_key) .await .unwrap(); - let loaded = CredentialRepository::get(&cred_repo, &space.id, "no-refresh-server") - .await - .unwrap() - .unwrap(); - assert!(!loaded.can_refresh()); + // clear_tokens should remove access + refresh, keep api_key + let cleared = cred_repo.clear_tokens(&space.id, "server").await.unwrap(); + assert!(cleared); + + let all = cred_repo.get_all(&space.id, "server").await.unwrap(); + assert_eq!(all.len(), 1); + assert_eq!(all[0].credential_type, CredentialType::ApiKey); } // ============================================================================= @@ -477,7 +504,13 @@ async fn test_different_encryptors_cannot_read_each_others_data() { let cred_repo2 = SqliteCredentialRepository::new(db, encryptor2); // Reading with wrong key should fail - let result = CredentialRepository::get(&cred_repo2, &space.id, "encrypted-server").await; + let result = CredentialRepository::get( + &cred_repo2, + &space.id, + "encrypted-server", + &CredentialType::ApiKey, + ) + .await; assert!( result.is_err() || result.unwrap().is_none(), "Should fail to decrypt with wrong key"