diff --git a/sqlx-core/src/any/connection/mod.rs b/sqlx-core/src/any/connection/mod.rs index c8c8b45809..1f1165806f 100644 --- a/sqlx-core/src/any/connection/mod.rs +++ b/sqlx-core/src/any/connection/mod.rs @@ -200,9 +200,8 @@ impl Connection for AnyConnection { #[cfg(feature = "mssql")] AnyConnectionKind::Mssql(_) => 0, - // no cache #[cfg(feature = "odbc")] - AnyConnectionKind::Odbc(_) => 0, + AnyConnectionKind::Odbc(conn) => conn.cached_statements_size(), } } @@ -221,9 +220,8 @@ impl Connection for AnyConnection { #[cfg(feature = "mssql")] AnyConnectionKind::Mssql(_) => Box::pin(futures_util::future::ok(())), - // no cache #[cfg(feature = "odbc")] - AnyConnectionKind::Odbc(_) => Box::pin(futures_util::future::ok(())), + AnyConnectionKind::Odbc(conn) => Box::pin(conn.clear_cached_statements()), } } diff --git a/sqlx-core/src/odbc/connection/executor.rs b/sqlx-core/src/odbc/connection/executor.rs index 12570979ae..61cf6611b9 100644 --- a/sqlx-core/src/odbc/connection/executor.rs +++ b/sqlx-core/src/odbc/connection/executor.rs @@ -5,7 +5,7 @@ use crate::odbc::{Odbc, OdbcConnection, OdbcQueryResult, OdbcRow, OdbcStatement, use either::Either; use futures_core::future::BoxFuture; use futures_core::stream::BoxStream; -use futures_util::TryStreamExt; +use futures_util::{TryFutureExt, TryStreamExt}; impl<'c> Executor<'c> for &'c mut OdbcConnection { type Database = Odbc; @@ -18,8 +18,15 @@ impl<'c> Executor<'c> for &'c mut OdbcConnection { 'c: 'e, E: Execute<'q, Self::Database> + 'q, { + let sql = query.sql(); let args = query.take_arguments(); - Box::pin(self.execute_stream(query.sql(), args).into_stream()) + let persistent = query.persistent(); + + Box::pin( + self.execute_stream(sql, args, persistent) + .map_ok(flume::Receiver::into_stream) + .try_flatten_stream(), + ) } fn fetch_optional<'e, 'q: 'e, E>( diff --git a/sqlx-core/src/odbc/connection/mod.rs b/sqlx-core/src/odbc/connection/mod.rs index 0a58fbe391..0f5922d6a4 100644 --- a/sqlx-core/src/odbc/connection/mod.rs +++ b/sqlx-core/src/odbc/connection/mod.rs @@ -216,20 +216,21 @@ impl OdbcConnection { } /// Launches a background task to execute the SQL statement and send the results to the returned channel. - pub(crate) fn execute_stream( + pub(crate) async fn execute_stream( &mut self, sql: &str, args: Option, - ) -> flume::Receiver, Error>> { - let (tx, rx) = flume::bounded(64); - - let sql_owned = sql.to_string(); - let maybe_prepared = if let Some(prepared) = self.stmt_cache.get_mut(sql) { - MaybePrepared::Prepared(Arc::clone(prepared)) + persistent: bool, + ) -> Result, Error>>, Error> { + let statement = if args.is_some() { + MaybePrepared::Prepared(self.prepared_statement(sql, persistent).await?) } else { - MaybePrepared::NotPrepared(sql_owned.clone()) + MaybePrepared::NotPrepared(sql.to_string()) }; + let (tx, rx) = flume::bounded(64); + + let sql_owned = sql.to_string(); let conn = Arc::clone(&self.conn); let buffer_settings = self.buffer_settings; let log_settings = self.log_settings.clone(); @@ -241,7 +242,7 @@ impl OdbcConnection { .and_then(|mut conn| { execute_sql( &mut conn, - maybe_prepared, + statement, args, &tx, buffer_settings, @@ -254,7 +255,29 @@ impl OdbcConnection { } })); - rx + Ok(rx) + } + + async fn prepared_statement( + &mut self, + sql: &str, + store_to_cache: bool, + ) -> Result { + if let Some(cached) = self.stmt_cache.get_mut(sql) { + return Ok(Arc::clone(cached)); + } + + let conn = Arc::clone(&self.conn); + let sql_owned = sql.to_string(); + let prepared = + spawn_blocking(move || conn.into_prepared(&sql_owned).map_err(Error::from)).await?; + let prepared = Arc::new(Mutex::new(prepared)); + + if store_to_cache && self.stmt_cache.is_enabled() { + self.stmt_cache.insert(sql, Arc::clone(&prepared)); + } + + Ok(prepared) } pub(crate) async fn clear_cached_statements(&mut self) -> Result<(), Error> { @@ -268,35 +291,16 @@ impl OdbcConnection { store_to_cache: bool, allow_deferred_result_columns: bool, ) -> Result, Error> { - let sql_owned = sql.to_string(); - let cached = self - .stmt_cache - .get_mut(sql) - .map(|prepared| Arc::clone(prepared)); + let prepared = self.prepared_statement(sql, false).await?; - if let Some(prepared) = cached { - let metadata = spawn_blocking(move || { + let (metadata, metadata_complete) = spawn_blocking({ + let prepared = Arc::clone(&prepared); + move || { let mut prepared = prepared.lock().map_err(|_| { Error::Protocol("ODBC prepare: failed to lock prepared statement".into()) })?; collect_statement_metadata(&mut prepared, allow_deferred_result_columns) - .map(|(metadata, _)| metadata) - }) - .await?; - - return Ok(OdbcStatement { - sql: Cow::Borrowed(sql), - metadata, - }); - } - - let conn = Arc::clone(&self.conn); - let sql_clone = sql_owned.clone(); - let (prepared, metadata, metadata_complete) = spawn_blocking(move || { - let mut prepared = conn.into_prepared(&sql_clone)?; - let metadata = - collect_statement_metadata(&mut prepared, allow_deferred_result_columns)?; - Ok::<_, Error>((prepared, metadata.0, metadata.1)) + } }) .await?; @@ -307,8 +311,7 @@ impl OdbcConnection { } if store_to_cache && metadata_complete && self.stmt_cache.is_enabled() { - self.stmt_cache - .insert(&sql_owned, Arc::new(Mutex::new(prepared))); + self.stmt_cache.insert(sql, prepared); } Ok(OdbcStatement { diff --git a/tests/any/odbc.rs b/tests/any/odbc.rs index 8bc3779722..8858684903 100644 --- a/tests/any/odbc.rs +++ b/tests/any/odbc.rs @@ -432,3 +432,25 @@ async fn it_accepts_standard_odbc_connection_strings() -> anyhow::Result<()> { Ok(()) } + +#[cfg(feature = "odbc")] +#[sqlx_macros::test] +async fn it_reports_the_statement_cache_via_any_odbc() -> anyhow::Result<()> { + let mut conn = odbc_conn().await?; + + for _ in 0..3 { + let row: AnyRow = sqlx_oldapi::query("SELECT ? AS value") + .bind(42i32) + .fetch_one(&mut conn) + .await?; + assert_eq!(row.try_get::("value")?, 42); + } + + assert_eq!(conn.cached_statements_size(), 1); + + conn.clear_cached_statements().await?; + assert_eq!(conn.cached_statements_size(), 0); + + conn.close().await?; + Ok(()) +} diff --git a/tests/odbc/odbc.rs b/tests/odbc/odbc.rs index 723f5baeb6..6d93a41061 100644 --- a/tests/odbc/odbc.rs +++ b/tests/odbc/odbc.rs @@ -80,6 +80,81 @@ async fn it_bounds_statement_cache() -> anyhow::Result<()> { Ok(()) } +#[tokio::test] +async fn it_caches_statements_across_repeated_queries() -> anyhow::Result<()> { + let mut conn = new::().await?; + + for _ in 0..10 { + let row = sqlx_oldapi::query(PARAMETERIZED_SELECT_WITH_COLUMN) + .bind(42_i32) + .fetch_one(&mut conn) + .await?; + assert_eq!(row.try_get_raw(0)?.to_owned().decode::(), 42); + } + + assert_eq!(conn.cached_statements_size(), 1); + + Ok(()) +} + +#[tokio::test] +async fn it_does_not_cache_non_persistent_queries() -> anyhow::Result<()> { + let mut conn = new::().await?; + + let row = sqlx_oldapi::query(PARAMETERIZED_SELECT_WITH_COLUMN) + .bind(42_i32) + .persistent(false) + .fetch_one(&mut conn) + .await?; + + assert_eq!(row.try_get_raw(0)?.to_owned().decode::(), 42); + assert_eq!(conn.cached_statements_size(), 0); + + Ok(()) +} + +#[tokio::test] +async fn it_does_not_cache_statements_without_arguments() -> anyhow::Result<()> { + let mut conn = new::().await?; + + conn.execute("SELECT 1").await?; + conn.execute("CREATE TEMPORARY TABLE sqlx_odbc_uncached (id INTEGER NOT NULL)") + .await?; + + assert_eq!(conn.cached_statements_size(), 0); + + Ok(()) +} + +#[tokio::test] +async fn it_reuses_cached_statements_after_partial_reads() -> anyhow::Result<()> { + let mut conn = new::().await?; + + conn.execute("CREATE TEMPORARY TABLE sqlx_odbc_partial_reads (id INTEGER NOT NULL)") + .await?; + let rows = (1..=200) + .map(|id| format!("({id})")) + .collect::>() + .join(", "); + conn.execute(&*format!( + "INSERT INTO sqlx_odbc_partial_reads (id) VALUES {rows}" + )) + .await?; + + for _ in 0..3 { + let row = + sqlx_oldapi::query("SELECT id FROM sqlx_odbc_partial_reads WHERE id > ? ORDER BY id") + .bind(0_i32) + .fetch_one(&mut conn) + .await?; + assert_eq!(row.try_get_raw(0)?.to_owned().decode::(), 1); + } + + assert_eq!(conn.cached_statements_size(), 1); + + Ok(()) +} + #[tokio::test] async fn it_rolls_back_dropped_transaction() -> anyhow::Result<()> { let mut conn = new::().await?;