diff --git a/backend/windmill-worker/src/mssql_executor.rs b/backend/windmill-worker/src/mssql_executor.rs index 4770e3456d..65dc022b74 100644 --- a/backend/windmill-worker/src/mssql_executor.rs +++ b/backend/windmill-worker/src/mssql_executor.rs @@ -28,6 +28,7 @@ struct MssqlDatabase { port: Option, dbname: String, instance_name: Option, + #[serde(deserialize_with = "deserialize_aad_token")] aad_token: Option, } @@ -110,7 +111,6 @@ pub async fn do_mssql( "Neither AAD token nor username/password credentials are set".to_string(), )); } - config.trust_cert(); // on production, it is not a good idea to do this let tcp = if use_instance_name { TcpStream::connect_named(&config).await.map_err(to_anyhow)? // named instance @@ -349,3 +349,61 @@ fn sql_to_json_value(val: ColumnData) -> Result { ), } } + +fn deserialize_aad_token<'de, D>(deserializer: D) -> Result, D::Error> +where + D: serde::Deserializer<'de>, +{ + // Use a custom visitor to handle both string and object cases + struct AadTokenVisitor; + impl<'de> serde::de::Visitor<'de> for AadTokenVisitor { + type Value = Option; + + fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result { + formatter.write_str("a string, null, or a JSON object with a token field") + } + + fn visit_str(self, _v: &str) -> Result + where + E: serde::de::Error, + { + // Any string (including empty) is treated as None + Ok(None) + } + + fn visit_none(self) -> Result + where + E: serde::de::Error, + { + // Explicit null is treated as None + Ok(None) + } + + fn visit_unit(self) -> Result + where + E: serde::de::Error, + { + // Unit (null in JSON) is treated as None + Ok(None) + } + + fn visit_map(self, map: A) -> Result + where + A: serde::de::MapAccess<'de>, + { + // Deserialize as a JSON value + let value = serde_json::Value::deserialize(serde::de::value::MapAccessDeserializer::new(map))?; + + // Check if it has an empty token + if let Some(token) = value.get("token").and_then(|t| t.as_str()) { + if token.is_empty() { + return Ok(None); + } + } + + Ok(Some(value)) + } + } + + deserializer.deserialize_any(AadTokenVisitor) +}