From ee26b2a5809c026f85fbd5631b14e363a41eedc8 Mon Sep 17 00:00:00 2001 From: t8y2 <1156263951@qq.com> Date: Fri, 31 Jul 2026 21:14:48 +0800 Subject: [PATCH] fix(mongodb): support legacy update commands --- crates/dbx-core/src/db/mongo_driver.rs | 45 ++++++++------ crates/dbx-core/src/mongo_shell.rs | 84 ++++++++++++++++++++++++++ 2 files changed, 110 insertions(+), 19 deletions(-) diff --git a/crates/dbx-core/src/db/mongo_driver.rs b/crates/dbx-core/src/db/mongo_driver.rs index c28593796..ef6aa99e3 100644 --- a/crates/dbx-core/src/db/mongo_driver.rs +++ b/crates/dbx-core/src/db/mongo_driver.rs @@ -1230,16 +1230,22 @@ pub async fn update_documents( serde_json::from_str(update_json).map_err(|e| format!("Invalid update JSON: {e}"))?; let filter = json_filter_to_document(&filter_value).map_err(|e| format!("Invalid filter: {e}"))?; let update = json_update_to_modifications(&update_value).map_err(|e| format!("Invalid update: {e}"))?; - let array_filters = parse_update_array_filters(options_json)?; + let ParsedMongoUpdateOptions { upsert, array_filters } = parse_update_options(options_json)?; let col = client.database(database).collection::(collection); let result = if many { let mut action = col.update_many(filter, update); + if let Some(upsert) = upsert { + action = action.upsert(upsert); + } if let Some(filters) = array_filters { action = action.array_filters(filters); } action.await.map_err(|e| e.to_string())? } else { let mut action = col.update_one(filter, update); + if let Some(upsert) = upsert { + action = action.upsert(upsert); + } if let Some(filters) = array_filters { action = action.array_filters(filters); } @@ -1249,17 +1255,24 @@ pub async fn update_documents( } #[derive(Default, Deserialize)] -#[serde(rename_all = "camelCase", deny_unknown_fields)] +#[serde(rename_all = "camelCase")] struct MongoUpdateOptions { + upsert: Option, array_filters: Option>, } -fn parse_update_array_filters(options_json: Option<&str>) -> Result>, String> { +#[derive(Debug, Default)] +struct ParsedMongoUpdateOptions { + upsert: Option, + array_filters: Option>, +} + +fn parse_update_options(options_json: Option<&str>) -> Result { let Some(raw) = options_json.filter(|value| !value.trim().is_empty()) else { - return Ok(None); + return Ok(ParsedMongoUpdateOptions::default()); }; let options: MongoUpdateOptions = serde_json::from_str(raw).map_err(|e| format!("Invalid update options: {e}"))?; - options + let array_filters = options .array_filters .map(|filters| { filters @@ -1268,7 +1281,8 @@ fn parse_update_array_filters(options_json: Option<&str>) -> Result, _>>() .map_err(|e| format!("Invalid arrayFilters: {e}")) }) - .transpose() + .transpose()?; + Ok(ParsedMongoUpdateOptions { upsert: options.upsert, array_filters }) } #[derive(Debug, Default, Deserialize)] @@ -2021,20 +2035,13 @@ mod tests { } #[test] - fn update_options_parse_array_filters() { - let filters = parse_update_array_filters(Some(r#"{"arrayFilters":[{"item.id":322678},{"item.active":true}]}"#)) - .unwrap() - .unwrap(); + fn update_options_parse_upsert_and_array_filters() { + let options = + parse_update_options(Some(r#"{"upsert":true,"arrayFilters":[{"item.id":322678},{"item.active":true}]}"#)) + .unwrap(); - assert_eq!(filters, vec![doc! { "item.id": 322678_i64 }, doc! { "item.active": true }]); - } - - #[test] - fn update_options_reject_unsupported_fields() { - let error = parse_update_array_filters(Some(r#"{"upsert":true}"#)).unwrap_err(); - - assert!(error.starts_with("Invalid update options:")); - assert!(error.contains("unknown field `upsert`")); + assert_eq!(options.upsert, Some(true)); + assert_eq!(options.array_filters.unwrap(), vec![doc! { "item.id": 322678_i64 }, doc! { "item.active": true }]); } #[test] diff --git a/crates/dbx-core/src/mongo_shell.rs b/crates/dbx-core/src/mongo_shell.rs index 351f96312..48c51803b 100644 --- a/crates/dbx-core/src/mongo_shell.rs +++ b/crates/dbx-core/src/mongo_shell.rs @@ -501,6 +501,19 @@ pub fn parse(input: &str) -> Result { }); } } + if let Some((args, tail)) = method_call(source, prefix_end, "update") { + if !tail.is_empty() || !(2..=3).contains(&args.len()) { + return Err("Invalid MongoDB update() command.".to_string()); + } + let (options, many) = legacy_update_options(args.get(2))?; + return Ok(MongoCommand::Update { + collection, + filter: normalized_json(&args[0])?, + update: normalized_json(&args[1])?, + options, + many, + }); + } for (method, many) in [("deleteOne", false), ("deleteMany", true)] { if let Some((args, tail)) = method_call(source, prefix_end, method) { if !tail.is_empty() || args.len() != 1 { @@ -672,6 +685,28 @@ fn optional_json_argument(value: Option<&String>) -> Result, Stri value.filter(|value| !value.trim().is_empty()).map(|value| normalized_json(value)).transpose() } +fn legacy_update_options(value: Option<&String>) -> Result<(Option, bool), String> { + let Some(value) = value.filter(|value| !value.trim().is_empty()) else { + return Ok((None, false)); + }; + let normalized = normalized_json(value)?; + let value = parse_json_value(&normalized).ok_or("Invalid MongoDB update() options.")?; + let Value::Object(mut options) = value else { + return Err("MongoDB update() options must be a document.".to_string()); + }; + let many = match options.remove("multi") { + Some(Value::Bool(many)) => many, + Some(_) => return Err("MongoDB update() multi option must be a boolean.".to_string()), + None => false, + }; + let options = if options.is_empty() { + None + } else { + Some(serde_json::to_string(&Value::Object(options)).map_err(|error| error.to_string())?) + }; + Ok((options, many)) +} + fn parse_use_database(source: &str) -> Option { let mut parts = source.split_whitespace(); if !parts.next()?.eq_ignore_ascii_case("use") { @@ -808,6 +843,9 @@ mod tests { assert!(aggregate.is_dangerous()); let update = parse("db.projects.updateMany({}, {$set: {active: false}})").unwrap(); assert!(update.has_empty_filter()); + let legacy_update = parse("db.projects.update({}, {$set: {active: false}}, {multi: true})").unwrap(); + assert!(legacy_update.has_empty_filter()); + assert_eq!(validate_safety(&legacy_update, true, false, false), Err(MongoSafetyError::EmptyFilter)); } #[test] @@ -856,6 +894,52 @@ mod tests { assert!(matches!(update, MongoCommand::Update { many: true, options: Some(_), .. })); } + #[test] + fn parses_legacy_update_with_single_and_multi_semantics() { + assert_eq!( + parse("db.projects.update({_id: 1}, {$set: {active: true}})").unwrap(), + MongoCommand::Update { + collection: "projects".to_string(), + filter: r#"{"_id":1}"#.to_string(), + update: r#"{"$set":{"active":true}}"#.to_string(), + options: None, + many: false, + } + ); + + let command = + parse(r#"db.getCollection("xxx").update({tenantId: 7}, {$set: {active: true}}, {upsert: true})"#).unwrap(); + let MongoCommand::Update { collection, update, options, many, .. } = command else { + panic!("expected legacy update command"); + }; + assert_eq!(collection, "xxx"); + assert!(!many); + assert_eq!(parse_json_value(&update).unwrap(), serde_json::json!({ "$set": { "active": true } })); + assert_eq!(parse_json_value(options.as_deref().unwrap()).unwrap(), serde_json::json!({ "upsert": true })); + + let command = parse( + r#"db.projects.update({tenantId: 7}, [{$set: {active: true}}], {multi: true, arrayFilters: [{"item.id": 1}]})"#, + ) + .unwrap(); + let MongoCommand::Update { update, options, many, .. } = command else { + panic!("expected legacy multi update command"); + }; + assert!(many); + assert_eq!(parse_json_value(&update).unwrap(), serde_json::json!([{ "$set": { "active": true } }])); + assert_eq!( + parse_json_value(options.as_deref().unwrap()).unwrap(), + serde_json::json!({ "arrayFilters": [{ "item.id": 1 }] }) + ); + } + + #[test] + fn rejects_invalid_legacy_update_arguments() { + assert!(parse("db.projects.update({_id: 1})").is_err()); + assert!(parse("db.projects.update({_id: 1}, {$set: {active: true}}, true)").is_err()); + assert!(parse("db.projects.update({_id: 1}, {$set: {active: true}}, {multi: 'yes'})").is_err()); + assert!(parse("db.projects.update({_id: 1}, {$set: {active: true}}, {}, false)").is_err()); + } + #[test] fn accepts_legacy_insert_and_rejects_unsupported_options() { assert_eq!(