diff --git a/src-tauri/src/commands/mcp_bridge.rs b/src-tauri/src/commands/mcp_bridge.rs index b20f51ac5..fe08a90d8 100644 --- a/src-tauri/src/commands/mcp_bridge.rs +++ b/src-tauri/src/commands/mcp_bridge.rs @@ -279,7 +279,7 @@ fn find_config_by_name<'a>( #[cfg(test)] mod tests { - use super::write_port_file; + use super::{resolve_mongo_database, resolve_mongo_target_values, write_port_file}; #[test] fn writes_bridge_port_file_to_resolved_data_dir() { @@ -300,6 +300,32 @@ mod tests { let _ = std::fs::remove_dir_all(root); } + + #[test] + fn mongo_database_uses_configured_default_for_missing_or_blank_request() { + let configured = Some("sample_db".to_string()); + + assert_eq!(resolve_mongo_database(None, configured.clone()), "sample_db"); + assert_eq!(resolve_mongo_database(Some(String::new()), configured.clone()), "sample_db"); + assert_eq!(resolve_mongo_database(Some(" ".to_string()), configured), "sample_db"); + } + + #[test] + fn mongo_database_preserves_explicit_target() { + assert_eq!(resolve_mongo_database(Some("admin".to_string()), Some("sample_db".to_string())), "admin"); + } + + #[test] + fn mongo_target_keeps_connection_id_separate_from_database() { + assert_eq!( + resolve_mongo_target_values( + "connection-id".to_string(), + Some("sample_db".to_string()), + Some("default_db".to_string()), + ), + ("connection-id".to_string(), "sample_db".to_string()) + ); + } } async fn respond(stream: &mut tokio::net::TcpStream, status: &str, body: &str) { @@ -349,13 +375,25 @@ fn check_visible_database(config: &crate::models::connection::ConnectionConfig, Ok(()) } -async fn resolve_mongo_pool_key( +fn resolve_mongo_database(requested: Option, configured: Option) -> String { + requested.filter(|database| !database.trim().is_empty()).or(configured).unwrap_or_default() +} + +fn resolve_mongo_target_values( + connection_id: String, + requested_database: Option, + configured_database: Option, +) -> (String, String) { + (connection_id, resolve_mongo_database(requested_database, configured_database)) +} + +async fn resolve_mongo_target( state: &Arc, connection_id: Option<&str>, connection_name: &str, database: Option, stream: &mut tokio::net::TcpStream, -) -> Option<(String, String, String)> { +) -> Option<(String, String)> { let config = match resolve_connection(state, connection_id, connection_name).await { Ok(c) => c, Err(e) => { @@ -363,16 +401,12 @@ async fn resolve_mongo_pool_key( return None; } }; - let connection_id = config.id.clone(); - let database = database.unwrap_or_else(|| config.database.clone().unwrap_or_default()); - let pool_key = match state.get_or_create_pool(&config.id, Some(&database)).await { - Ok(key) => key, - Err(e) => { - respond_error(stream, "500 Internal Server Error", &e).await; - return None; - } - }; - Some((pool_key, database, connection_id)) + let (connection_id, database) = resolve_mongo_target_values(config.id.clone(), database, config.database.clone()); + if let Err(e) = check_visible_database(&config, &database) { + respond_error(stream, "403 Forbidden", &e).await; + return None; + } + Some((connection_id, database)) } async fn handle_open_table(app: &AppHandle, state: &Arc, body: &str, stream: &mut tokio::net::TcpStream) { @@ -488,12 +522,12 @@ async fn handle_mongo_list_collections_data(state: &Arc, body: &str, s return; } }; - let Some((pool_key, database, _connection_id)) = - resolve_mongo_pool_key(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await + let Some((connection_id, database)) = + resolve_mongo_target(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await else { return; }; - match dbx_core::mongo_ops::mongo_list_collections_core(state, &pool_key, &database).await { + match dbx_core::mongo_ops::mongo_list_collections_core(state, &connection_id, &database).await { Ok(collections) => respond_json(stream, &collections).await, Err(e) => respond_error(stream, "500 Internal Server Error", &e).await, } @@ -507,14 +541,14 @@ async fn handle_mongo_find_documents_data(state: &Arc, body: &str, str return; } }; - let Some((pool_key, database, _connection_id)) = - resolve_mongo_pool_key(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await + let Some((connection_id, database)) = + resolve_mongo_target(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await else { return; }; match dbx_core::mongo_ops::mongo_find_documents_core( state, - &pool_key, + &connection_id, &database, &req.collection, req.skip.unwrap_or(0), @@ -538,14 +572,14 @@ async fn handle_mongo_count_documents_data(state: &Arc, body: &str, st return; } }; - let Some((pool_key, database, _connection_id)) = - resolve_mongo_pool_key(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await + let Some((connection_id, database)) = + resolve_mongo_target(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await else { return; }; match dbx_core::mongo_ops::mongo_count_documents_core( state, - &pool_key, + &connection_id, &database, &req.collection, req.filter.as_deref(), @@ -566,12 +600,12 @@ async fn handle_mongo_server_version_data(state: &Arc, body: &str, str return; } }; - let Some((pool_key, database, _connection_id)) = - resolve_mongo_pool_key(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await + let Some((connection_id, database)) = + resolve_mongo_target(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await else { return; }; - match dbx_core::mongo_ops::mongo_server_version_core(state, &pool_key, &database).await { + match dbx_core::mongo_ops::mongo_server_version_core(state, &connection_id, &database).await { Ok(version) => respond_json(stream, &version).await, Err(e) => respond_error(stream, "500 Internal Server Error", &e).await, } @@ -585,12 +619,12 @@ async fn handle_mongo_collection_stats_data(state: &Arc, body: &str, s return; } }; - let Some((pool_key, database, _connection_id)) = - resolve_mongo_pool_key(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await + let Some((connection_id, database)) = + resolve_mongo_target(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await else { return; }; - match dbx_core::mongo_ops::mongo_collection_stats_core(state, &pool_key, &database, &req.collection, req.scale) + match dbx_core::mongo_ops::mongo_collection_stats_core(state, &connection_id, &database, &req.collection, req.scale) .await { Ok(result) => respond_json(stream, &result).await, @@ -606,14 +640,14 @@ async fn handle_mongo_aggregate_documents_data(state: &Arc, body: &str return; } }; - let Some((pool_key, database, _connection_id)) = - resolve_mongo_pool_key(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await + let Some((connection_id, database)) = + resolve_mongo_target(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await else { return; }; match dbx_core::mongo_ops::mongo_aggregate_documents_core( state, - &pool_key, + &connection_id, &database, &req.collection, &req.pipeline_json, @@ -634,8 +668,8 @@ async fn handle_mongo_create_index_data(state: &Arc, body: &str, strea return; } }; - let Some((pool_key, database, connection_id)) = - resolve_mongo_pool_key(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await + let Some((connection_id, database)) = + resolve_mongo_target(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await else { return; }; @@ -645,7 +679,7 @@ async fn handle_mongo_create_index_data(state: &Arc, body: &str, strea } match dbx_core::mongo_ops::mongo_create_index_core( state, - &pool_key, + &connection_id, &database, &req.collection, &req.keys_json, @@ -666,8 +700,8 @@ async fn handle_mongo_drop_indexes_data(state: &Arc, body: &str, strea return; } }; - let Some((pool_key, database, connection_id)) = - resolve_mongo_pool_key(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await + let Some((connection_id, database)) = + resolve_mongo_target(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await else { return; }; @@ -677,7 +711,7 @@ async fn handle_mongo_drop_indexes_data(state: &Arc, body: &str, strea } match dbx_core::mongo_ops::mongo_drop_indexes_core( state, - &pool_key, + &connection_id, &database, &req.collection, req.indexes_json.as_deref(), @@ -698,8 +732,8 @@ async fn handle_mongo_drop_collection_data(state: &Arc, body: &str, st return; } }; - let Some((pool_key, database, connection_id)) = - resolve_mongo_pool_key(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await + let Some((connection_id, database)) = + resolve_mongo_target(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await else { return; }; @@ -707,7 +741,7 @@ async fn handle_mongo_drop_collection_data(state: &Arc, body: &str, st respond_error(stream, "403 Forbidden", &e).await; return; } - match dbx_core::mongo_ops::mongo_drop_collection_core(state, &pool_key, &database, &req.collection).await { + match dbx_core::mongo_ops::mongo_drop_collection_core(state, &connection_id, &database, &req.collection).await { Ok(()) => respond_json(stream, &serde_json::json!({ "ok": true })).await, Err(e) => respond_error(stream, "500 Internal Server Error", &e).await, } @@ -721,8 +755,8 @@ async fn handle_mongo_insert_documents_data(state: &Arc, body: &str, s return; } }; - let Some((pool_key, database, connection_id)) = - resolve_mongo_pool_key(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await + let Some((connection_id, database)) = + resolve_mongo_target(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await else { return; }; @@ -730,8 +764,14 @@ async fn handle_mongo_insert_documents_data(state: &Arc, body: &str, s respond_error(stream, "403 Forbidden", &e).await; return; } - match dbx_core::mongo_ops::mongo_insert_documents_core(state, &pool_key, &database, &req.collection, &req.docs_json) - .await + match dbx_core::mongo_ops::mongo_insert_documents_core( + state, + &connection_id, + &database, + &req.collection, + &req.docs_json, + ) + .await { Ok(inserted) => respond_json(stream, &serde_json::json!({ "affected_rows": inserted })).await, Err(e) => respond_error(stream, "500 Internal Server Error", &e).await, @@ -746,8 +786,8 @@ async fn handle_mongo_update_documents_data(state: &Arc, body: &str, s return; } }; - let Some((pool_key, database, connection_id)) = - resolve_mongo_pool_key(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await + let Some((connection_id, database)) = + resolve_mongo_target(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await else { return; }; @@ -757,7 +797,7 @@ async fn handle_mongo_update_documents_data(state: &Arc, body: &str, s } match dbx_core::mongo_ops::mongo_update_documents_core( state, - &pool_key, + &connection_id, &database, &req.collection, &req.filter_json, @@ -780,8 +820,8 @@ async fn handle_mongo_delete_documents_data(state: &Arc, body: &str, s return; } }; - let Some((pool_key, database, connection_id)) = - resolve_mongo_pool_key(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await + let Some((connection_id, database)) = + resolve_mongo_target(state, req.connection_id.as_deref(), &req.connection_name, req.database, stream).await else { return; }; @@ -791,7 +831,7 @@ async fn handle_mongo_delete_documents_data(state: &Arc, body: &str, s } match dbx_core::mongo_ops::mongo_delete_documents_core( state, - &pool_key, + &connection_id, &database, &req.collection, &req.filter_json,