Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 2 additions & 4 deletions sqlx-core/src/any/connection/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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(),
}
}

Expand All @@ -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()),
}
}

Expand Down
11 changes: 9 additions & 2 deletions sqlx-core/src/odbc/connection/executor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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>(
Expand Down
75 changes: 39 additions & 36 deletions sqlx-core/src/odbc/connection/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<OdbcArguments>,
) -> flume::Receiver<Result<Either<OdbcQueryResult, OdbcRow>, 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<flume::Receiver<Result<Either<OdbcQueryResult, OdbcRow>, 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();
Expand All @@ -241,7 +242,7 @@ impl OdbcConnection {
.and_then(|mut conn| {
execute_sql(
&mut conn,
maybe_prepared,
statement,
args,
&tx,
buffer_settings,
Expand All @@ -254,7 +255,29 @@ impl OdbcConnection {
}
}));

rx
Ok(rx)
}

async fn prepared_statement(
&mut self,
sql: &str,
store_to_cache: bool,
) -> Result<SharedPreparedStatement, Error> {
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> {
Expand All @@ -268,35 +291,16 @@ impl OdbcConnection {
store_to_cache: bool,
allow_deferred_result_columns: bool,
) -> Result<OdbcStatement<'a>, 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?;

Expand All @@ -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 {
Expand Down
22 changes: 22 additions & 0 deletions tests/any/odbc.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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::<i32, _>("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(())
}
75 changes: 75 additions & 0 deletions tests/odbc/odbc.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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::<Odbc>().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::<i64>(), 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::<Odbc>().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::<i64>(), 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::<Odbc>().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::<Odbc>().await?;

conn.execute("CREATE TEMPORARY TABLE sqlx_odbc_partial_reads (id INTEGER NOT NULL)")
.await?;
let rows = (1..=200)
.map(|id| format!("({id})"))
.collect::<Vec<_>>()
.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::<i64>(), 1);
}

assert_eq!(conn.cached_statements_size(), 1);

Ok(())
}

#[tokio::test]
async fn it_rolls_back_dropped_transaction() -> anyhow::Result<()> {
let mut conn = new::<Odbc>().await?;
Expand Down
Loading