diff --git a/.cargo-husky/hooks/pre-commit b/.cargo-husky/hooks/pre-commit index b770d69..8312ad6 100755 --- a/.cargo-husky/hooks/pre-commit +++ b/.cargo-husky/hooks/pre-commit @@ -1,4 +1,6 @@ #!/bin/sh +# +# This hook was set by cargo-husky v1.5.0: https://github.com/rhysd/cargo-husky#readme set -eux diff --git a/benches/connection.rs b/benches/connection.rs index 775f08b..16a60c0 100644 --- a/benches/connection.rs +++ b/benches/connection.rs @@ -137,6 +137,7 @@ fn write_to_file(file_name: &str, data: &[u8]) { }; let mut file = OpenOptions::new() .create(true) + .truncate(true) .write(true) .open(path) .expect("unable to write to file"); diff --git a/src/bindings/mod.rs b/src/bindings.rs similarity index 100% rename from src/bindings/mod.rs rename to src/bindings.rs diff --git a/src/connection/mod.rs b/src/connection.rs similarity index 82% rename from src/connection/mod.rs rename to src/connection.rs index ef5a3e1..2d2e5b2 100644 --- a/src/connection/mod.rs +++ b/src/connection.rs @@ -181,15 +181,25 @@ fn read_error_message(mg_session: *mut bindings::mg_session) -> String { unsafe { c_string_to_string(c_error_message, None) } } +/// Decrements the global connection count, finalizing mgclient once the last connection is gone. +fn release_connection() { + if CONNECTION_COUNT.fetch_sub(1, Ordering::SeqCst) == 1 { + // This was the last connection, safe to finalize. + Connection::finalize(); + } +} + +/// Converts an optional string into an optional `CString`, mapping interior null bytes to an error. +fn optional_cstring(value: Option<&String>, field: &str) -> Result, MgError> { + value + .map(|s| CString::new(s.as_str()).map_err(|_| MgError::null_byte(field))) + .transpose() +} + impl Drop for Connection { fn drop(&mut self) { unsafe { bindings::mg_session_destroy(self.mg_session) }; - - // Decrement the connection counter and finalize only if this was the last connection - if CONNECTION_COUNT.fetch_sub(1, Ordering::SeqCst) == 1 { - // This was the last connection, safe to finalize - Connection::finalize(); - } + release_connection(); } } @@ -294,6 +304,26 @@ impl Connection { self.arraysize = arraysize; } + /// Builds a query error from the session's last error message, without changing status. + fn query_error(&self) -> MgError { + MgError::query(read_error_message(self.mg_session)) + } + + /// Marks the connection as bad and returns a query error from the session's last message. + fn fail_query(&mut self) -> MgError { + self.status = ConnectionStatus::Bad; + self.query_error() + } + + /// The status to settle into once a query's results have been fully consumed. + fn settled_status(&self) -> ConnectionStatus { + if self.autocommit { + ConnectionStatus::Ready + } else { + ConnectionStatus::InTransaction + } + } + /// Creates a connection to database using provided connection parameters. /// /// Returns `Connection` if connection to database is successfully established, otherwise @@ -324,10 +354,8 @@ impl Connection { let mg_session_params = unsafe { bindings::mg_session_params_make() }; if mg_session_params.is_null() { - // Connection failed, decrement the counter and finalize if needed - if CONNECTION_COUNT.fetch_sub(1, Ordering::SeqCst) == 1 { - Connection::finalize(); - } + // Connection failed, release the counter and finalize if needed + release_connection(); return Err(MgError::ffi( "Failed to allocate mg_session_params".to_string(), )); @@ -335,32 +363,14 @@ impl Connection { let mut trust_callback_box: Option> = None; // Create CStrings and keep them alive for the duration of mg_connect - let c_host = match param_struct.host.as_ref() { - Some(s) => Some(CString::new(s.as_str()).map_err(|_| MgError::null_byte("host"))?), - None => None, - }; - let c_address = match param_struct.address.as_ref() { - Some(s) => Some(CString::new(s.as_str()).map_err(|_| MgError::null_byte("address"))?), - None => None, - }; - let c_username = match param_struct.username.as_ref() { - Some(s) => Some(CString::new(s.as_str()).map_err(|_| MgError::null_byte("username"))?), - None => None, - }; - let c_password = match param_struct.password.as_ref() { - Some(s) => Some(CString::new(s.as_str()).map_err(|_| MgError::null_byte("password"))?), - None => None, - }; + let c_host = optional_cstring(param_struct.host.as_ref(), "host")?; + let c_address = optional_cstring(param_struct.address.as_ref(), "address")?; + let c_username = optional_cstring(param_struct.username.as_ref(), "username")?; + let c_password = optional_cstring(param_struct.password.as_ref(), "password")?; let c_client_name = CString::new(param_struct.client_name.as_str()) .map_err(|_| MgError::null_byte("client_name"))?; - let c_sslcert = match param_struct.sslcert.as_ref() { - Some(s) => Some(CString::new(s.as_str()).map_err(|_| MgError::null_byte("sslcert"))?), - None => None, - }; - let c_sslkey = match param_struct.sslkey.as_ref() { - Some(s) => Some(CString::new(s.as_str()).map_err(|_| MgError::null_byte("sslkey"))?), - None => None, - }; + let c_sslcert = optional_cstring(param_struct.sslcert.as_ref(), "sslcert")?; + let c_sslkey = optional_cstring(param_struct.sslkey.as_ref(), "sslkey")?; unsafe { if let Some(ref x) = c_host { @@ -418,10 +428,8 @@ impl Connection { }; if status != 0 { - // Connection failed, decrement the counter and finalize if needed - if CONNECTION_COUNT.fetch_sub(1, Ordering::SeqCst) == 1 { - Connection::finalize(); - } + // Connection failed, release the counter and finalize if needed + release_connection(); return Err(MgError::connection(read_error_message(mg_session))); } @@ -455,20 +463,14 @@ impl Connection { 0 => { self.status = ConnectionStatus::Executing; } - _ => { - self.status = ConnectionStatus::Bad; - return Err(MgError::query(read_error_message(self.mg_session))); - } + _ => return Err(self.fail_query()), } match unsafe { bindings::mg_session_pull(self.mg_session, std::ptr::null_mut()) } { 0 => { self.status = ConnectionStatus::Fetching; } - _ => { - self.status = ConnectionStatus::Bad; - return Err(MgError::query(read_error_message(self.mg_session))); - } + _ => return Err(self.fail_query()), } loop { @@ -481,10 +483,7 @@ impl Connection { self.status = ConnectionStatus::Ready; return Ok(()); } - _ => { - self.status = ConnectionStatus::Bad; - return Err(MgError::query(read_error_message(self.mg_session))); - } + _ => return Err(self.fail_query()), }; } } @@ -520,10 +519,8 @@ impl Connection { } if !self.autocommit && self.status == ConnectionStatus::Ready { - match self.execute_without_results("BEGIN") { - Ok(()) => self.status = ConnectionStatus::InTransaction, - Err(err) => return Err(err), - } + self.execute_without_results("BEGIN")?; + self.status = ConnectionStatus::InTransaction; } self.summary = None; @@ -551,8 +548,7 @@ impl Connection { } if status != 0 { - self.status = ConnectionStatus::Bad; - return Err(MgError::query(read_error_message(self.mg_session))); + return Err(self.fail_query()); } self.status = ConnectionStatus::Executing; @@ -594,33 +590,22 @@ impl Connection { match self.lazy { true => { - if self.status == ConnectionStatus::Executing { - match self.pull(1) { - Ok(_) => { - // The state update is already done in the pull. - } - Err(err) => { - self.status = ConnectionStatus::Bad; - return Err(err); - } - } + if self.status == ConnectionStatus::Executing + && let Err(err) = self.pull(1) + { + self.status = ConnectionStatus::Bad; + return Err(err); } + // On success pull() has already updated the status. // Fetch the record or summary match self.fetch() { Ok((Some(x), None)) => { - // Got a record, fetch summary to check has_more - match self.fetch()? { - (None, Some(has_more)) => { - if has_more { - self.status = ConnectionStatus::Executing; - } - // If has_more is false, leave status as Fetching - } - _ => { - // If we don't get a summary, stay in Fetching state - } + // Got a record; peek at the summary to see if more records remain. + if let (None, Some(true)) = self.fetch()? { + self.status = ConnectionStatus::Executing; } + // Otherwise stay in Fetching state. Ok(Some(x)) } Ok((None, Some(has_more))) => { @@ -628,31 +613,14 @@ impl Connection { if has_more { self.status = ConnectionStatus::Executing; } else { - self.status = if self.autocommit { - ConnectionStatus::Ready - } else { - ConnectionStatus::InTransaction - }; + self.status = self.settled_status(); } Ok(None) } - Ok(_) => { - // Unexpected case - self.status = if self.autocommit { - ConnectionStatus::Ready - } else { - ConnectionStatus::InTransaction - }; - Ok(None) - } - Err(_) => { - // If fetch fails (e.g., "called fetch while not executing"), - // it means no more records, finalize the transaction - self.status = if self.autocommit { - ConnectionStatus::Ready - } else { - ConnectionStatus::InTransaction - }; + // No more records: either an unexpected fetch result, or a fetch error such as + // "called fetch while not executing". Either way, finalize the transaction. + Ok(_) | Err(_) => { + self.status = self.settled_status(); Ok(None) } } @@ -660,11 +628,7 @@ impl Connection { false => match self.next_record() { Some(x) => Ok(Some(x)), None => { - self.status = if self.autocommit { - ConnectionStatus::Ready - } else { - ConnectionStatus::InTransaction - }; + self.status = self.settled_status(); Ok(None) } }, @@ -687,22 +651,25 @@ impl Connection { /// Returns error if connection is not in `Executing` status or if there was an error while /// pulling record from database. pub fn fetchmany(&mut self, size: Option) -> Result, MgError> { - let size = match size { - Some(x) => x, - None => self.arraysize, - }; + let size = size.unwrap_or(self.arraysize); + self.fetch_records(Some(size)) + } + /// Fetches records one at a time until exhausted, or until `limit` records have been + /// collected when a limit is provided. + fn fetch_records(&mut self, limit: Option) -> Result, MgError> { let mut vec = Vec::new(); - for _i in 0..size { - match self.fetchone() { - Ok(record) => match record { - Some(x) => vec.push(x), - None => break, - }, - Err(err) => return Err(err), + let mut remaining = limit; + loop { + if remaining == Some(0) { + break; + } + match self.fetchone()? { + Some(x) => vec.push(x), + None => break, } + remaining = remaining.map(|n| n - 1); } - Ok(vec) } @@ -711,17 +678,7 @@ impl Connection { /// Returns error if connection is not in `Executing` status or if there was an error while /// pulling record from database. pub fn fetchall(&mut self) -> Result, MgError> { - let mut vec = Vec::new(); - loop { - match self.fetchone() { - Ok(record) => match record { - Some(x) => vec.push(x), - None => break, - }, - Err(err) => return Err(err), - } - } - Ok(vec) + self.fetch_records(None) } fn pull(&mut self, n: i64) -> Result<(), MgError> { @@ -781,10 +738,7 @@ impl Connection { self.status = ConnectionStatus::Fetching; Ok(()) } - _ => { - self.status = ConnectionStatus::Bad; - Err(MgError::query(read_error_message(self.mg_session))) - } + _ => Err(self.fail_query()), } } @@ -830,21 +784,15 @@ impl Connection { self.summary = Some(mg_map_to_hash_map(mg_summary)); Ok((None, Some(has_more))) }, - _ => Err(MgError::query(read_error_message(self.mg_session))), + _ => Err(self.query_error()), } } fn pull_and_fetch_all(&mut self) -> Result, MgError> { + self.pull(0)?; let mut res = Vec::new(); - match self.pull(0) { - Ok(_) => loop { - let x = self.fetch()?; - match x { - (Some(x), _) => res.push(x), - (None, _) => break, - } - }, - Err(err) => return Err(err), + while let (Some(x), _) = self.fetch()? { + res.push(x); } Ok(res) } @@ -876,13 +824,9 @@ impl Connection { return Ok(()); } - match self.execute_without_results("COMMIT") { - Ok(()) => { - self.status = ConnectionStatus::Ready; - Ok(()) - } - Err(err) => Err(err), - } + self.execute_without_results("COMMIT")?; + self.status = ConnectionStatus::Ready; + Ok(()) } /// Rollback any pending transaction to the database. @@ -914,13 +858,9 @@ impl Connection { return Ok(()); } - match self.execute_without_results("ROLLBACK") { - Ok(()) => { - self.status = ConnectionStatus::Ready; - Ok(()) - } - Err(err) => Err(err), - } + self.execute_without_results("ROLLBACK")?; + self.status = ConnectionStatus::Ready; + Ok(()) } /// Closes the connection. diff --git a/src/connection/tests.rs b/src/connection/tests.rs index 9c2dc46..efe6df9 100644 --- a/src/connection/tests.rs +++ b/src/connection/tests.rs @@ -2,6 +2,21 @@ use super::*; use crate::{Node, Value}; use serial_test::serial; +/// Base connection params for tests, overridable via `MGHOST`/`MGPORT`/`MGUSER`/`MGPASSWORD` +/// (host/port default to `127.0.0.1:7687`, credentials unset). +fn test_params() -> ConnectParams { + ConnectParams { + host: Some(std::env::var("MGHOST").unwrap_or_else(|_| "127.0.0.1".to_string())), + port: std::env::var("MGPORT") + .ok() + .and_then(|p| p.parse().ok()) + .unwrap_or(7687), + username: std::env::var("MGUSER").ok(), + password: std::env::var("MGPASSWORD").ok(), + ..Default::default() + } +} + fn get_connection(prms: &ConnectParams) -> Connection { match Connection::connect(prms) { Ok(c) => c, @@ -18,9 +33,8 @@ fn execute_query(connection: &mut Connection, query: &str) -> Vec { fn execute_query_and_fetchall(query: &str) -> Vec { let connect_prms = ConnectParams { - address: Some(String::from("127.0.0.1")), autocommit: true, - ..Default::default() + ..test_params() }; let mut connection = get_connection(&connect_prms); assert_eq!(connection.status, ConnectionStatus::Ready); @@ -41,10 +55,7 @@ fn execute_query_and_fetchall(query: &str) -> Vec { } fn initialize() -> Connection { - let connect_prms = ConnectParams { - address: Some(String::from("127.0.0.1")), - ..Default::default() - }; + let connect_prms = test_params(); let mut connection = get_connection(&connect_prms); assert_eq!(connection.status, ConnectionStatus::Ready); @@ -107,11 +118,10 @@ fn my_callback(host: &String, ip_address: &String, key_type: &String, fingerprin fn panic_sslcert() { initialize(); let connect_prms = ConnectParams { - address: Some(String::from("127.0.0.1")), trust_callback: Some(&my_callback), lazy: false, sslcert: Some(String::from("test_sslcert")), - ..Default::default() + ..test_params() }; get_connection(&connect_prms); } @@ -122,11 +132,10 @@ fn panic_sslcert() { fn panic_sslkey() { initialize(); let connect_prms = ConnectParams { - address: Some(String::from("127.0.0.1")), trust_callback: Some(&my_callback), lazy: false, sslkey: Some(String::from("test_sslkey")), - ..Default::default() + ..test_params() }; let _connection = get_connection(&connect_prms); } diff --git a/src/error/mod.rs b/src/error.rs similarity index 100% rename from src/error/mod.rs rename to src/error.rs diff --git a/src/value/mod.rs b/src/value.rs similarity index 80% rename from src/value/mod.rs rename to src/value.rs index 63cc3f8..a525c32 100644 --- a/src/value/mod.rs +++ b/src/value.rs @@ -80,143 +80,94 @@ pub enum QueryParam { Map(HashMap), } +/// Returns `ptr` unless it is null, in which case a fresh mg null value is returned instead. +fn value_or_null(ptr: *mut bindings::mg_value) -> *mut bindings::mg_value { + if ptr.is_null() { + unsafe { bindings::mg_value_make_null() } + } else { + ptr + } +} + impl QueryParam { fn to_c_mg_value(&self) -> *mut bindings::mg_value { + // Wraps an intermediate mgclient handle in an mg_value, destroying the handle and + // falling back to an mg null value if either allocation fails. + macro_rules! wrap_or_null { + ($intermediate:expr, $make:path, $destroy:path) => {{ + let handle = $intermediate; + if handle.is_null() { + return bindings::mg_value_make_null(); + } + let ptr = $make(handle); + if ptr.is_null() { + $destroy(handle); + return bindings::mg_value_make_null(); + } + ptr + }}; + } + unsafe { match self { QueryParam::Null => bindings::mg_value_make_null(), QueryParam::Bool(x) => { - let ptr = bindings::mg_value_make_bool(match *x { + let val = match *x { false => 0, true => 1, - }); - if ptr.is_null() { - return bindings::mg_value_make_null(); - } - ptr - } - QueryParam::Int(x) => { - let ptr = bindings::mg_value_make_integer(*x); - if ptr.is_null() { - return bindings::mg_value_make_null(); - } - ptr - } - QueryParam::Float(x) => { - let ptr = bindings::mg_value_make_float(*x); - if ptr.is_null() { - return bindings::mg_value_make_null(); - } - ptr + }; + value_or_null(bindings::mg_value_make_bool(val)) } + QueryParam::Int(x) => value_or_null(bindings::mg_value_make_integer(*x)), + QueryParam::Float(x) => value_or_null(bindings::mg_value_make_float(*x)), QueryParam::String(x) => { // String parameter may contain null bytes - return null on error let c_string = match CString::new(x.as_str()) { Ok(s) => s, Err(_) => return bindings::mg_value_make_null(), }; - let ptr = bindings::mg_value_make_string(c_string.as_ptr()); - if ptr.is_null() { - return bindings::mg_value_make_null(); - } - ptr - } - QueryParam::Date(x) => { - let mg_date = naive_date_to_mg_date(x); - if mg_date.is_null() { - return bindings::mg_value_make_null(); - } - let ptr = bindings::mg_value_make_date(mg_date); - if ptr.is_null() { - bindings::mg_date_destroy(mg_date); - return bindings::mg_value_make_null(); - } - ptr - } - QueryParam::LocalTime(x) => { - let mg_local_time = naive_local_time_to_mg_local_time(x); - if mg_local_time.is_null() { - return bindings::mg_value_make_null(); - } - let ptr = bindings::mg_value_make_local_time(mg_local_time); - if ptr.is_null() { - bindings::mg_local_time_destroy(mg_local_time); - return bindings::mg_value_make_null(); - } - ptr - } - QueryParam::LocalDateTime(x) => { - let mg_local_date_time = naive_local_date_time_to_mg_local_date_time(x); - if mg_local_date_time.is_null() { - return bindings::mg_value_make_null(); - } - let ptr = bindings::mg_value_make_local_date_time(mg_local_date_time); - if ptr.is_null() { - bindings::mg_local_date_time_destroy(mg_local_date_time); - return bindings::mg_value_make_null(); - } - ptr - } - QueryParam::Duration(x) => { - let mg_duration = duration_to_mg_duration(x); - if mg_duration.is_null() { - return bindings::mg_value_make_null(); - } - let ptr = bindings::mg_value_make_duration(mg_duration); - if ptr.is_null() { - bindings::mg_duration_destroy(mg_duration); - return bindings::mg_value_make_null(); - } - ptr - } - QueryParam::Point2D(x) => { - let mg_point_2d = point2d_to_mg_point_2d(x); - if mg_point_2d.is_null() { - return bindings::mg_value_make_null(); - } - let ptr = bindings::mg_value_make_point_2d(mg_point_2d); - if ptr.is_null() { - bindings::mg_point_2d_destroy(mg_point_2d); - return bindings::mg_value_make_null(); - } - ptr - } - QueryParam::Point3D(x) => { - let mg_point_3d = point3d_to_mg_point_3d(x); - if mg_point_3d.is_null() { - return bindings::mg_value_make_null(); - } - let ptr = bindings::mg_value_make_point_3d(mg_point_3d); - if ptr.is_null() { - bindings::mg_point_3d_destroy(mg_point_3d); - return bindings::mg_value_make_null(); - } - ptr - } - QueryParam::List(x) => { - let mg_list = vector_to_mg_list(x); - if mg_list.is_null() { - return bindings::mg_value_make_null(); - } - let ptr = bindings::mg_value_make_list(mg_list); - if ptr.is_null() { - bindings::mg_list_destroy(mg_list); - return bindings::mg_value_make_null(); - } - ptr - } - QueryParam::Map(x) => { - let mg_map = hash_map_to_mg_map(x); - if mg_map.is_null() { - return bindings::mg_value_make_null(); - } - let ptr = bindings::mg_value_make_map(mg_map); - if ptr.is_null() { - bindings::mg_map_destroy(mg_map); - return bindings::mg_value_make_null(); - } - ptr + value_or_null(bindings::mg_value_make_string(c_string.as_ptr())) } + QueryParam::Date(x) => wrap_or_null!( + naive_date_to_mg_date(x), + bindings::mg_value_make_date, + bindings::mg_date_destroy + ), + QueryParam::LocalTime(x) => wrap_or_null!( + naive_local_time_to_mg_local_time(x), + bindings::mg_value_make_local_time, + bindings::mg_local_time_destroy + ), + QueryParam::LocalDateTime(x) => wrap_or_null!( + naive_local_date_time_to_mg_local_date_time(x), + bindings::mg_value_make_local_date_time, + bindings::mg_local_date_time_destroy + ), + QueryParam::Duration(x) => wrap_or_null!( + duration_to_mg_duration(x), + bindings::mg_value_make_duration, + bindings::mg_duration_destroy + ), + QueryParam::Point2D(x) => wrap_or_null!( + point2d_to_mg_point_2d(x), + bindings::mg_value_make_point_2d, + bindings::mg_point_2d_destroy + ), + QueryParam::Point3D(x) => wrap_or_null!( + point3d_to_mg_point_3d(x), + bindings::mg_value_make_point_3d, + bindings::mg_point_3d_destroy + ), + QueryParam::List(x) => wrap_or_null!( + vector_to_mg_list(x), + bindings::mg_value_make_list, + bindings::mg_list_destroy + ), + QueryParam::Map(x) => wrap_or_null!( + hash_map_to_mg_map(x), + bindings::mg_value_make_map, + bindings::mg_map_destroy + ), } } } @@ -662,12 +613,8 @@ pub(crate) fn naive_date_to_mg_date(input: &NaiveDate) -> *mut bindings::mg_date let unix_epoch = NaiveDate::from_ymd_opt(1970, 1, 1) .expect("Unix epoch is a valid date") .num_days_from_ce(); - let ptr = unsafe { bindings::mg_date_make((input.num_days_from_ce() - unix_epoch) as i64) }; - // mg_date_make can return NULL on OOM - if ptr.is_null() { - return std::ptr::null_mut(); - } - ptr + // mg_date_make returns NULL on OOM, which we propagate to the caller as-is. + unsafe { bindings::mg_date_make((input.num_days_from_ce() - unix_epoch) as i64) } } pub(crate) fn naive_local_time_to_mg_local_time(input: &NaiveTime) -> *mut bindings::mg_local_time { @@ -675,13 +622,8 @@ pub(crate) fn naive_local_time_to_mg_local_time(input: &NaiveTime) -> *mut bindi let minutes_ns = minutes_as_seconds(input.minute() as i64) * NSEC_IN_SEC; let seconds_ns = (input.second() as i64) * NSEC_IN_SEC; let nanoseconds = input.nanosecond() as i64; - let ptr = - unsafe { bindings::mg_local_time_make(hours_ns + minutes_ns + seconds_ns + nanoseconds) }; - // mg_local_time_make can return NULL on OOM - if ptr.is_null() { - return std::ptr::null_mut(); - } - ptr + // mg_local_time_make returns NULL on OOM, which we propagate to the caller as-is. + unsafe { bindings::mg_local_time_make(hours_ns + minutes_ns + seconds_ns + nanoseconds) } } pub(crate) fn naive_local_date_time_to_mg_local_date_time( @@ -696,14 +638,10 @@ pub(crate) fn naive_local_date_time_to_mg_local_date_time( let minutes_s = minutes_as_seconds(input.minute() as i64); let seconds_s = input.second() as i64; let nanoseconds = input.nanosecond() as i64; - let ptr = unsafe { + // mg_local_date_time_make returns NULL on OOM, which we propagate to the caller as-is. + unsafe { bindings::mg_local_date_time_make(days_s + hours_s + minutes_s + seconds_s, nanoseconds) - }; - // mg_local_date_time_make can return NULL on OOM - if ptr.is_null() { - return std::ptr::null_mut(); } - ptr } pub(crate) fn duration_to_mg_duration(input: &Duration) -> *mut bindings::mg_duration { @@ -718,38 +656,25 @@ pub(crate) fn duration_to_mg_duration(input: &Duration) -> *mut bindings::mg_dur duration -= Duration::seconds(seconds); // After subtracting days and seconds, remaining nanoseconds should always fit in i64 let nanoseconds = duration.num_nanoseconds().unwrap_or(0); - let ptr = unsafe { bindings::mg_duration_make(0, days, seconds, nanoseconds) }; - // mg_duration_make can return NULL on OOM - if ptr.is_null() { - return std::ptr::null_mut(); - } - ptr + // mg_duration_make returns NULL on OOM, which we propagate to the caller as-is. + unsafe { bindings::mg_duration_make(0, days, seconds, nanoseconds) } } pub(crate) fn point2d_to_mg_point_2d(input: &Point2D) -> *mut bindings::mg_point_2d { - let ptr = - unsafe { bindings::mg_point_2d_make(input.srid, input.x_longitude, input.y_latitude) }; - // mg_point_2d_make can return NULL on OOM - if ptr.is_null() { - return std::ptr::null_mut(); - } - ptr + // mg_point_2d_make returns NULL on OOM, which we propagate to the caller as-is. + unsafe { bindings::mg_point_2d_make(input.srid, input.x_longitude, input.y_latitude) } } pub(crate) fn point3d_to_mg_point_3d(input: &Point3D) -> *mut bindings::mg_point_3d { - let ptr = unsafe { + // mg_point_3d_make returns NULL on OOM, which we propagate to the caller as-is. + unsafe { bindings::mg_point_3d_make( input.srid, input.x_longitude, input.y_latitude, input.z_height, ) - }; - // mg_point_3d_make can return NULL on OOM - if ptr.is_null() { - return std::ptr::null_mut(); } - ptr } pub(crate) fn vector_to_mg_list(vector: &[QueryParam]) -> *mut bindings::mg_list { diff --git a/src/value/tests.rs b/src/value/tests.rs index 83296d9..c374a01 100644 --- a/src/value/tests.rs +++ b/src/value/tests.rs @@ -50,7 +50,7 @@ unsafe fn to_array_of_unbound_relationships( } unsafe fn to_c_int_array(vec: &[i64]) -> *mut i64 { - let size = vec.len() * mem::size_of::(); + let size = mem::size_of_val(vec); let ptr = unsafe { libc::malloc(size) as *mut i64 }; for (i, el) in vec.iter().enumerate() { unsafe { *ptr.add(i) = *el }; diff --git a/tests/datetime_test.rs b/tests/datetime_test.rs index 3984de5..4b50e5a 100644 --- a/tests/datetime_test.rs +++ b/tests/datetime_test.rs @@ -4,7 +4,13 @@ use rsmgclient::{ConnectParams, Connection, Value}; fn test_datetime_with_timezone() { // Setup: Create connection parameters and connect to the database let params = ConnectParams { - host: Some(String::from("localhost")), + host: Some(std::env::var("MGHOST").unwrap_or_else(|_| "localhost".to_string())), + port: std::env::var("MGPORT") + .ok() + .and_then(|p| p.parse().ok()) + .unwrap_or(7687), + username: std::env::var("MGUSER").ok(), + password: std::env::var("MGPASSWORD").ok(), ..ConnectParams::default() }; let mut connection = Connection::connect(¶ms).unwrap(); @@ -21,7 +27,7 @@ fn test_datetime_with_timezone() { // Extract the datetime value from the result if let Some(record) = records.first() { - if let Some(Value::DateTime(datetime)) = record.values.get(0) { + if let Some(Value::DateTime(datetime)) = record.values.first() { // Assert the datetime fields assert_eq!(datetime.year, 2024); assert_eq!(datetime.month, 4); @@ -37,7 +43,7 @@ fn test_datetime_with_timezone() { || datetime .time_zone_id .as_ref() - .map_or(false, |id| id.starts_with("TZ_")), + .is_some_and(|id| id.starts_with("TZ_")), "Expected timezone ID to be 'Etc/UTC' or start with 'TZ_', got {:?}", datetime.time_zone_id );