fix(mcp): resolve MongoDB connection targets

This commit is contained in:
t8y2 2026-07-15 13:44:51 +08:00
parent ebebc4fccd
commit fb2a04660b
1 changed files with 90 additions and 50 deletions

View File

@ -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<String>, configured: Option<String>) -> String {
requested.filter(|database| !database.trim().is_empty()).or(configured).unwrap_or_default()
}
fn resolve_mongo_target_values(
connection_id: String,
requested_database: Option<String>,
configured_database: Option<String>,
) -> (String, String) {
(connection_id, resolve_mongo_database(requested_database, configured_database))
}
async fn resolve_mongo_target(
state: &Arc<AppState>,
connection_id: Option<&str>,
connection_name: &str,
database: Option<String>,
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<AppState>, body: &str, stream: &mut tokio::net::TcpStream) {
@ -488,12 +522,12 @@ async fn handle_mongo_list_collections_data(state: &Arc<AppState>, 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<AppState>, 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<AppState>, 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<AppState>, 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<AppState>, 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<AppState>, 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<AppState>, 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<AppState>, 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<AppState>, 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<AppState>, 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<AppState>, 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<AppState>, 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<AppState>, 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<AppState>, 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<AppState>, 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<AppState>, 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<AppState>, 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<AppState>, body: &str, s
}
match dbx_core::mongo_ops::mongo_delete_documents_core(
state,
&pool_key,
&connection_id,
&database,
&req.collection,
&req.filter_json,