Skip to content
12 changes: 12 additions & 0 deletions connectorx-python/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,18 @@ fn connectorx(_: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_wrapped(wrap_pyfunction!(get_meta))?;
#[cfg(feature = "srcs")]
{
let version = m
.py()
.import("importlib.metadata")?
.call_method1("version", ("connectorx",))?
.extract::<String>()?;
let python_version = m.py().import("sys")?.getattr("version_info")?;
let major = python_version.getattr("major")?.extract::<u8>()?;
let minor = python_version.getattr("minor")?.extract::<u8>()?;
let micro = python_version.getattr("micro")?.extract::<u8>()?;
let runtime = format!("Python {major}.{minor}.{micro}");
::connectorx::sources::mssql::set_user_agent_info(version, runtime)
.map_err(PyRuntimeError::new_err)?;
m.add_wrapped(wrap_pyfunction!(get_mssql_driver))?;
m.add_wrapped(wrap_pyfunction!(set_mssql_driver))?;
}
Expand Down
3 changes: 3 additions & 0 deletions connectorx/src/sources/mssql/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -60,3 +60,6 @@ pub use self::dual_impl::{MsSQLSource, MsSQLSourceParser, MsSQLSourcePartition};

#[cfg(feature = "src_mssql_tds")]
pub(crate) use self::tds_impl::tds_get_partition_range;

#[cfg(feature = "src_mssql_tds")]
pub use self::tds_impl::set_user_agent_info;
75 changes: 74 additions & 1 deletion connectorx/src/sources/mssql/tds_impl.rs
Original file line number Diff line number Diff line change
Expand Up @@ -59,13 +59,38 @@ use sqlparser::dialect::MsSqlDialect;
use std::collections::HashMap;
use std::convert::TryFrom;
use std::ops::{Deref, DerefMut};
use std::sync::Arc;
use std::sync::{Arc, OnceLock};
use std::time::Duration;
use tokio::runtime::Runtime;
use url::Url;
use urlencoding::decode;
use uuid_old::Uuid;

#[derive(Clone, PartialEq)]
struct UserAgentInfo {
version: String,
runtime: String,
}

static USER_AGENT_INFO: OnceLock<UserAgentInfo> = OnceLock::new();

/// Sets the wrapper's package version and runtime for the TDS User-Agent, leaving LOGIN7 unchanged.
/// Language bindings should call this before creating connections. Reinitializing
/// with the same values is allowed; different values are rejected.
pub fn set_user_agent_info(version: String, runtime: String) -> Result<(), &'static str> {
if version.is_empty() {
return Err("MSSQL User-Agent version must not be empty");
}
if runtime.is_empty() {
return Err("MSSQL User-Agent runtime must not be empty");
}
let info = UserAgentInfo { version, runtime };
if USER_AGENT_INFO.get_or_init(|| info.clone()) != &info {
return Err("MSSQL User-Agent is already initialized with different values");
}
Ok(())
}

/// Builds the `mssql-tds` datasource string and [`ClientContext`] from a
/// ConnectorX MSSQL connection URL. Mirrors
/// [`super::tiberius_impl::mssql_config`]'s parsing, but expressed against
Expand All @@ -84,6 +109,15 @@ fn build_client_context(url: &Url) -> (String, ClientContext) {
};

let mut context = ClientContext::with_data_source(&datasource);
// The User-Agent feature carries a separate driver name from LOGIN7.
context
.user_agent
.set_library_name("connectorx".to_string());
if let Some(info) = USER_AGENT_INFO.get() {
context.user_agent.set_driver_version(info.version.clone());
context.user_agent.set_runtime(info.runtime.clone());
}
context.application_name = "ConnectorX".to_string();
context.database = decode(&url.path()[1..])?.into_owned();

let params: HashMap<String, String> = url.query_pairs().into_owned().collect();
Expand Down Expand Up @@ -137,6 +171,45 @@ fn build_client_context(url: &Url) -> (String, ClientContext) {
mod configuration_tests {
use super::*;

#[test]
fn python_user_agent_info_preserves_login_version() {
let version = "0.4.7a1";
let runtime = "Python 3.12.3";
assert!(set_user_agent_info(String::new(), runtime.to_string()).is_err());
assert!(set_user_agent_info(version.to_string(), String::new()).is_err());
set_user_agent_info(version.to_string(), runtime.to_string()).unwrap();
set_user_agent_info(version.to_string(), runtime.to_string()).unwrap();
assert!(set_user_agent_info("0.4.8".to_string(), runtime.to_string()).is_err());
assert!(set_user_agent_info(version.to_string(), "Python 3.13.0".to_string()).is_err());

let url = Url::parse("mssql://localhost/db").unwrap();
let (_, context) = build_client_context(&url).unwrap();
assert_eq!(context.user_agent.driver_version, version);
assert_eq!(context.user_agent.runtime, runtime);
assert_eq!(
context.driver_version,
ClientContext::with_data_source("tcp:localhost,1433").driver_version
);
}

#[test]
fn driver_identity_and_application_name() {
for (query, expected_appname) in [
("", "ConnectorX"),
("?appname=Custom%20Application", "Custom Application"),
("?appname=", ""),
] {
let url = Url::parse(&format!("mssql://localhost/db{query}")).unwrap();
let (_, context) = build_client_context(&url).unwrap();
assert_eq!(
context.library_name,
ClientContext::with_data_source("tcp:localhost,1433").library_name
);
assert_eq!(context.user_agent.library_name, "connectorx");
assert_eq!(context.application_name, expected_appname);
}
}

#[test]
fn integrated_auth_respects_tiberius_platform_and_feature_gates() {
for trusted in ["true", "false", "TrUe"] {
Expand Down
44 changes: 44 additions & 0 deletions connectorx/tests/test_mssql.rs
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,50 @@ use tokio::runtime::Runtime;

mod test_db;

#[cfg(feature = "src_mssql_tds")]
#[test]
fn test_mssql_tds_session_identity() {
let default_library_name =
mssql_tds::connection::client_context::ClientContext::with_data_source(
"tcp:localhost,1433",
)
.library_name;
let rt = Arc::new(Runtime::new().unwrap());
let mut url = url::Url::parse(&test_db::mssql_url()).unwrap();
let params: Vec<_> = url
.query_pairs()
.into_owned()
.filter(|(key, _)| key != "appname")
.collect();
url.set_query(None);
url.query_pairs_mut().extend_pairs(params);

for appname in [None, Some("ConnectorX identity override")] {
let mut conn = url.clone();
if let Some(appname) = appname {
conn.query_pairs_mut().append_pair("appname", appname);
}
let mut source = MsSQLSource::new(rt.clone(), conn.as_str(), 1).unwrap();
source.set_queries(&[CXQuery::naked(
"SELECT CAST(client_interface_name AS varchar(128)) AS driver_name, \
CAST(program_name AS varchar(128)) AS application_name \
FROM sys.dm_exec_sessions WHERE session_id = @@SPID",
)]);
source.fetch_metadata().unwrap();
let mut partitions = source.partition().unwrap();
let mut parser = partitions[0].parser().unwrap();
assert_eq!(parser.fetch_next().unwrap(), (1, true));
assert_eq!(
parser.parse::<Option<&str>>().unwrap(),
Some(default_library_name.as_str())
);
assert_eq!(
parser.parse::<Option<&str>>().unwrap(),
Some(appname.unwrap_or("ConnectorX"))
);
}
}

#[cfg(feature = "src_mssql_tds")]
mod tds_pool_tests {
use super::*;
Expand Down
Loading