Skip to content

Commit eed1a47

Browse files
committed
fix(quaint): add url decode for MSSQL database name
1 parent c6be8e6 commit eed1a47

4 files changed

Lines changed: 60 additions & 23 deletions

File tree

quaint/src/connector/connection_info.rs

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -125,7 +125,7 @@ impl ConnectionInfo {
125125
#[cfg(feature = "mysql-native")]
126126
NativeConnectionInfo::Mysql(url) => url.dbname().map(Cow::Borrowed),
127127
#[cfg(feature = "mssql-native")]
128-
NativeConnectionInfo::Mssql(url) => Some(Cow::Borrowed(url.dbname())),
128+
NativeConnectionInfo::Mssql(url) => Some(url.dbname()),
129129
#[cfg(feature = "sqlite-native")]
130130
NativeConnectionInfo::Sqlite { .. } | NativeConnectionInfo::InMemorySqlite { .. } => None,
131131
},

quaint/src/connector/mssql/native/mod.rs

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -66,7 +66,8 @@ pub struct Mssql {
6666
impl Mssql {
6767
/// Creates a new connection to SQL Server.
6868
pub async fn new(url: MssqlUrl) -> crate::Result<Self> {
69-
let config = Config::from_jdbc_string(&url.connection_string)?;
69+
let mut config = Config::from_jdbc_string(&url.connection_string)?;
70+
config.database(url.dbname());
7071
let tcp = TcpStream::connect_named(&config).await?;
7172
let socket_timeout = url.socket_timeout();
7273

@@ -75,6 +76,7 @@ impl Mssql {
7576
Ok(client) => Ok(client),
7677
Err(tiberius::error::Error::Routing { host, port }) => {
7778
let mut config = Config::from_jdbc_string(&url.connection_string)?;
79+
config.database(url.dbname());
7880
config.host(host);
7981
config.port(port);
8082

quaint/src/connector/mssql/url.rs

Lines changed: 51 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,8 @@ use crate::{
55
error::{Error, ErrorKind},
66
};
77
use connection_string::JdbcString;
8-
use std::{fmt, str::FromStr, time::Duration};
8+
use percent_encoding::percent_decode;
9+
use std::{borrow::Cow, fmt, str::FromStr, time::Duration};
910

1011
/// Wraps a connection url and exposes the parsing logic used by Quaint,
1112
/// including default values.
@@ -97,9 +98,16 @@ impl MssqlUrl {
9798
self.query_params.transaction_isolation_level
9899
}
99100

100-
/// Name of the database.
101-
pub fn dbname(&self) -> &str {
102-
self.query_params.database()
101+
/// Decoded database name. Defaults to `master`.
102+
pub fn dbname(&self) -> Cow<'_, str> {
103+
let db = self.query_params.database();
104+
match percent_decode(db.as_bytes()).decode_utf8() {
105+
Ok(decoded) => decoded,
106+
Err(_) => {
107+
tracing::warn!("Couldn't decode dbname to UTF-8, using the non-decoded version.");
108+
Cow::Borrowed(db)
109+
}
110+
}
103111
}
104112

105113
/// The prefix which to use when querying database.
@@ -368,6 +376,7 @@ impl MssqlUrl {
368376

369377
#[cfg(test)]
370378
mod tests {
379+
use super::*;
371380
use crate::tests::test_api::mssql::CONN_STR;
372381
use crate::{error::*, single::Quaint};
373382

@@ -381,4 +390,42 @@ mod tests {
381390
let err = res.unwrap_err();
382391
assert!(matches!(err.kind(), ErrorKind::AuthenticationFailed { user } if user == &Name::available("WRONG")));
383392
}
393+
394+
#[test]
395+
fn should_decode_percent_encoded_dbname() {
396+
// Chinese characters: 测试库 (test database)
397+
let url = MssqlUrl::new("sqlserver://localhost:1433;database=%E6%B5%8B%E8%AF%95%E5%BA%93;user=SA;password=pass;trustServerCertificate=true").unwrap();
398+
assert_eq!("测试库", url.dbname());
399+
}
400+
401+
#[test]
402+
fn should_decode_dbname_with_spaces() {
403+
let url = MssqlUrl::new("sqlserver://localhost:1433;database=my%20database;user=SA;password=pass;trustServerCertificate=true").unwrap();
404+
assert_eq!("my database", url.dbname());
405+
}
406+
407+
#[test]
408+
fn should_decode_dbname_with_special_characters() {
409+
// test-db_name
410+
let url = MssqlUrl::new("sqlserver://localhost:1433;database=test%2Ddb%5Fname;user=SA;password=pass;trustServerCertificate=true").unwrap();
411+
assert_eq!("test-db_name", url.dbname());
412+
}
413+
414+
#[test]
415+
fn should_return_master_as_default_dbname() {
416+
let url = MssqlUrl::new("sqlserver://localhost:1433;user=SA;password=pass;trustServerCertificate=true").unwrap();
417+
assert_eq!("master", url.dbname());
418+
}
419+
420+
#[tokio::test]
421+
async fn should_connect_to_percent_encoded_chinese_dbname() {
422+
// 测试库 = %E6%B5%8B%E8%AF%95%E5%BA%93
423+
let url = CONN_STR.replace("database=master", "database=%E6%B5%8B%E8%AF%95%E5%BA%93");
424+
let conn = Quaint::new(&url).await.unwrap();
425+
426+
use crate::prelude::Queryable;
427+
let result = conn.query_raw("SELECT DB_NAME() AS db_name", &[]).await.unwrap();
428+
let db_name = result.first().unwrap().get("db_name").unwrap().to_string().unwrap();
429+
assert_eq!("测试库", db_name);
430+
}
384431
}

schema-engine/connectors/sql-schema-connector/src/flavour/mssql.rs

Lines changed: 5 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -153,11 +153,13 @@ impl MssqlConnector {
153153

154154
/// Get the url as a JDBC string, extract the database name, and re-encode the string.
155155
fn master_url(input: &str) -> ConnectorResult<(String, String)> {
156+
let url = MssqlUrl::new(input).map_err(ConnectorError::url_parse_error)?;
157+
let db_name = url.dbname().into_owned();
158+
156159
let mut conn = JdbcString::from_str(&format!("jdbc:{input}"))
157160
.map_err(|e| ConnectorError::from_source(e, "JDBC string parse error"))?;
158-
let params = conn.properties_mut();
161+
conn.properties_mut().remove("database");
159162

160-
let db_name = params.remove("database").unwrap_or_else(|| String::from("master"));
161163
Ok((db_name, conn.to_string()))
162164
}
163165
}
@@ -255,22 +257,8 @@ impl SqlConnector for MssqlConnector {
255257
fn drop_database(&mut self) -> BoxFuture<'_, ConnectorResult<()>> {
256258
Box::pin(async {
257259
let params = self.state.get_unwrapped_params();
258-
let connection_string = &params.connector_params.connection_string;
259-
{
260-
let conn_str: JdbcString = format!("jdbc:{connection_string}")
261-
.parse()
262-
.map_err(ConnectorError::url_parse_error)?;
263-
264-
let db_name = conn_str
265-
.properties()
266-
.get("database")
267-
.map(|s| s.to_owned())
268-
.unwrap_or_else(|| "master".to_owned());
269-
270-
assert!(db_name != "master", "Cannot drop the `master` database.");
271-
}
272-
273260
let (db_name, master_uri) = Self::master_url(&params.connector_params.connection_string)?;
261+
assert!(db_name != "master", "Cannot drop the `master` database.");
274262
let mut conn = Connection::new(&master_uri.to_string()).await?;
275263

276264
let query = format!("DROP DATABASE IF EXISTS [{db_name}]");

0 commit comments

Comments
 (0)