diff --git a/README.md b/README.md index 1b33bbe1cb..15e57ca1b9 100644 --- a/README.md +++ b/README.md @@ -308,6 +308,20 @@ while let Some(row) = rows.try_next().await? { } ``` +Databases with native named parameter support can bind using `bind_named()`. Pass the exact +parameter token used in the SQL statement, including its marker. SQLite supports `:name`, `@name`, +and `$name`; MSSQL uses `@name`. + +```rust +let row = query_scalar::<_, i64>("SELECT :value + :value") + .bind_named(":value", 150_i64) + .fetch_one(&mut conn) + .await?; +``` + +Named and positional parameters cannot be mixed in one query. Databases without native named +parameter support do not provide `bind_named()`. + To assist with mapping the row into a domain type, there are two idioms that may be used: ```rust diff --git a/sqlx-core/src/arguments.rs b/sqlx-core/src/arguments.rs index 2261be7fa7..65ca5b695a 100644 --- a/sqlx-core/src/arguments.rs +++ b/sqlx-core/src/arguments.rs @@ -23,6 +23,19 @@ pub trait Arguments<'q>: Send + Sized + Default { } } +/// Arguments bound by their database-native parameter names. +/// +/// The name must include the parameter marker used in the SQL statement, e.g. `:id` for SQLite +/// or `@id` for MSSQL. This trait is implemented only by databases whose native parameter +/// protocol supports named parameters. Named and positional parameters must not be mixed in one +/// query. +pub trait NamedArguments<'q>: Arguments<'q> { + /// Add a value for the named parameter. + fn add_named(&mut self, name: &'q str, value: T) + where + T: 'q + Send + Encode<'q, Self::Database> + Type; +} + pub trait IntoArguments<'q, DB: HasArguments<'q>>: Sized + Send { fn into_arguments(self) -> >::Arguments; } diff --git a/sqlx-core/src/logger.rs b/sqlx-core/src/logger.rs index 6cfd80d215..ec52f04664 100644 --- a/sqlx-core/src/logger.rs +++ b/sqlx-core/src/logger.rs @@ -51,7 +51,7 @@ impl<'q> QueryLogger<'q> { let sql = if summary != self.sql { summary.push_str(" …"); - format!("\n\n{}\n", &self.sql) + format!("\n\n{}\n", self.sql) } else { String::new() }; @@ -125,7 +125,7 @@ impl<'q, O: Debug + Hash + Eq, R: Debug, P: Debug> QueryPlanLogger<'q, O, R, P> let sql = if summary != self.sql { summary.push_str(" …"); - format!("\n\n{}\n", &self.sql) + format!("\n\n{}\n", self.sql) } else { String::new() }; diff --git a/sqlx-core/src/mssql/arguments.rs b/sqlx-core/src/mssql/arguments.rs index 98e6de84be..967742c3f6 100644 --- a/sqlx-core/src/mssql/arguments.rs +++ b/sqlx-core/src/mssql/arguments.rs @@ -1,4 +1,4 @@ -use crate::arguments::Arguments; +use crate::arguments::{Arguments, NamedArguments}; use crate::encode::Encode; use crate::mssql::database::Mssql; use crate::mssql::io::MssqlBufMutExt; @@ -14,14 +14,12 @@ pub struct MssqlArguments { name: String, pub(crate) data: Vec, pub(crate) declarations: String, + positional: bool, + named: bool, } impl MssqlArguments { - pub(crate) fn add_named<'q, T: Encode<'q, Mssql> + Type>( - &mut self, - name: &str, - value: T, - ) { + fn add_rpc_named<'q, T: Encode<'q, Mssql> + Type>(&mut self, name: &str, value: T) { let ty = value.produces().unwrap_or_else(T::type_info); let mut ty_name = String::new(); @@ -35,7 +33,7 @@ impl MssqlArguments { } pub(crate) fn add_unnamed<'q, T: Encode<'q, Mssql> + Type>(&mut self, value: T) { - self.add_named("", value); + self.add_rpc_named("", value); } pub(crate) fn declare<'q, T: Encode<'q, Mssql> + Type>( @@ -64,7 +62,7 @@ impl MssqlArguments { where T: Encode<'q, Mssql> + Type, { - let ty = value.produces().unwrap_or_else(T::type_info); + self.positional = true; // produce an ordinal parameter name // @p1, @p2, ... @pN @@ -75,31 +73,33 @@ impl MssqlArguments { self.ordinal += 1; self.name.push_str(itoa::Buffer::new().format(self.ordinal)); - let MssqlArguments { - ref name, - ref mut declarations, - ref mut data, - .. - } = self; + let name = std::mem::take(&mut self.name); + self.add_query_named(&name, value); + self.name = name; + } - // add this to our variable declaration list - // @p1 int, @p2 nvarchar(10), ... + fn add_query_named<'q, T>(&mut self, name: &str, value: T) + where + T: Encode<'q, Mssql> + Type, + { + let ty = value.produces().unwrap_or_else(T::type_info); - if !declarations.is_empty() { - declarations.push(','); + if !self.declarations.is_empty() { + self.declarations.push(','); } - declarations.push_str(name); - declarations.push(' '); - ty.0.fmt(declarations); + self.declarations.push_str(name); + self.declarations.push(' '); + ty.0.fmt(&mut self.declarations); - // write out the parameter - - data.put_b_varchar(name); // [ParamName] - data.push(0); // [StatusFlags] + self.data.put_b_varchar(name); // [ParamName] + self.data.push(0); // [StatusFlags] + ty.0.put(&mut self.data); // [TYPE_INFO] + ty.0.put_value(&mut self.data, value); // [ParamLenData] + } - ty.0.put(data); // [TYPE_INFO] - ty.0.put_value(data, value); // [ParamLenData] + pub(crate) fn has_mixed_binding(&self) -> bool { + self.positional && self.named } } @@ -126,6 +126,16 @@ impl<'q> Arguments<'q> for MssqlArguments { } } +impl<'q> NamedArguments<'q> for MssqlArguments { + fn add_named(&mut self, name: &'q str, value: T) + where + T: 'q + Send + Encode<'q, Self::Database> + Type, + { + self.named = true; + self.add_query_named(name, value); + } +} + #[cfg(test)] mod tests { use super::*; @@ -167,4 +177,22 @@ mod tests { assert_eq!(sql, "SELECT * FROM table WHERE id=@p1 AND name=@p2"); } + + #[test] + fn test_named_query_parameter() { + let mut args = MssqlArguments::default(); + args.add_query_named("@id", 42_i32); + + assert_eq!(args.declarations, "@id int"); + assert!(!args.data.is_empty()); + } + + #[test] + fn test_mixed_query_parameters_are_detected() { + let mut args = MssqlArguments::default(); + args.add(42_i32); + args.add_named("@id", 42_i32); + + assert!(args.has_mixed_binding()); + } } diff --git a/sqlx-core/src/mssql/connection/executor.rs b/sqlx-core/src/mssql/connection/executor.rs index 3358e94ea6..89ceda33e6 100644 --- a/sqlx-core/src/mssql/connection/executor.rs +++ b/sqlx-core/src/mssql/connection/executor.rs @@ -22,6 +22,14 @@ use std::sync::Arc; impl MssqlConnection { async fn run(&mut self, query: &str, arguments: Option) -> Result<(), Error> { + if let Some(arguments) = arguments.as_ref() { + if arguments.has_mixed_binding() { + return Err(err_protocol!( + "cannot mix named and positional MSSQL parameters" + )); + } + } + self.stream.wait_until_ready().await?; self.stream.pending_done_count += 1; diff --git a/sqlx-core/src/mssql/connection/stream.rs b/sqlx-core/src/mssql/connection/stream.rs index 0911d4d235..8a35beae30 100644 --- a/sqlx-core/src/mssql/connection/stream.rs +++ b/sqlx-core/src/mssql/connection/stream.rs @@ -342,7 +342,7 @@ pub(crate) fn write_packets<'en, T: Encode<'en>>( ); } - packet_header.truncate(0); + packet_header.clear(); PacketHeader { r#type: ty, status: if is_last { diff --git a/sqlx-core/src/postgres/connection/sasl.rs b/sqlx-core/src/postgres/connection/sasl.rs index 6e6ee68339..665590d313 100644 --- a/sqlx-core/src/postgres/connection/sasl.rs +++ b/sqlx-core/src/postgres/connection/sasl.rs @@ -113,7 +113,7 @@ pub(crate) async fn authenticate( let client_final_message_wo_proof = format!( "{channel_binding},r={nonce}", channel_binding = channel_binding, - nonce = &cont.nonce + nonce = cont.nonce ); // AuthMessage := client-first-message-bare + "," + server-first-message + "," + client-final-message-without-proof diff --git a/sqlx-core/src/query.rs b/sqlx-core/src/query.rs index 1a19359644..2b0fedfbe9 100644 --- a/sqlx-core/src/query.rs +++ b/sqlx-core/src/query.rs @@ -4,7 +4,7 @@ use either::Either; use futures_core::stream::BoxStream; use futures_util::{future, StreamExt, TryFutureExt, TryStreamExt}; -use crate::arguments::{Arguments, IntoArguments}; +use crate::arguments::{Arguments, IntoArguments, NamedArguments}; use crate::database::{Database, HasArguments, HasStatement, HasStatementCache}; use crate::encode::Encode; use crate::error::Error; @@ -82,6 +82,21 @@ impl<'q, DB: Database> Query<'q, DB, >::Arguments> { .add(value); self } + + /// Bind a value for use with a database-native named SQL parameter. + /// + /// `name` must include the parameter marker used in the query, such as `:id` for SQLite or + /// `@id` for MSSQL. Named and positional parameters must not be mixed in one query. + pub fn bind_named(mut self, name: &'q str, value: T) -> Self + where + >::Arguments: NamedArguments<'q, Database = DB>, + T: 'q + Send + Encode<'q, DB> + Type, + { + self.arguments + .get_or_insert_with(Default::default) + .add_named(name, value); + self + } } impl<'q, DB, A> Query<'q, DB, A> diff --git a/sqlx-core/src/query_as.rs b/sqlx-core/src/query_as.rs index c13f687e7c..79b06e0aa0 100644 --- a/sqlx-core/src/query_as.rs +++ b/sqlx-core/src/query_as.rs @@ -4,7 +4,7 @@ use either::Either; use futures_core::stream::BoxStream; use futures_util::{StreamExt, TryStreamExt}; -use crate::arguments::IntoArguments; +use crate::arguments::{IntoArguments, NamedArguments}; use crate::database::{Database, HasArguments, HasStatement, HasStatementCache}; use crate::encode::Encode; use crate::error::Error; @@ -55,6 +55,16 @@ impl<'q, DB: Database, O> QueryAs<'q, DB, O, >::Arguments self.inner = self.inner.bind(value); self } + + /// Bind a value for use with a database-native named SQL parameter. + pub fn bind_named(mut self, name: &'q str, value: T) -> Self + where + >::Arguments: NamedArguments<'q, Database = DB>, + T: 'q + Send + Encode<'q, DB> + Type, + { + self.inner = self.inner.bind_named(name, value); + self + } } impl<'q, DB, O, A> QueryAs<'q, DB, O, A> diff --git a/sqlx-core/src/query_builder.rs b/sqlx-core/src/query_builder.rs index 8f8bbf8d87..10a9f86305 100644 --- a/sqlx-core/src/query_builder.rs +++ b/sqlx-core/src/query_builder.rs @@ -4,7 +4,7 @@ use std::fmt::Display; use std::fmt::Write; use std::marker::PhantomData; -use crate::arguments::Arguments; +use crate::arguments::{Arguments, NamedArguments}; use crate::database::{Database, HasArguments}; use crate::encode::Encode; use crate::from_row::FromRow; @@ -135,6 +135,26 @@ where self } + /// Push an exact database-native named parameter token and bind a value to it. + /// + /// For example, use `:id` for SQLite or `@id` for MSSQL. The token is appended verbatim to the + /// SQL query. Named and positional parameters must not be mixed in one query. + pub fn push_bind_named(&mut self, name: &'args str, value: T) -> &mut Self + where + >::Arguments: NamedArguments<'args, Database = DB>, + T: 'args + Encode<'args, DB> + Send + Type, + { + self.sanity_check(); + + self.arguments + .as_mut() + .expect("BUG: Arguments taken already") + .add_named(name, value); + self.query.push_str(name); + + self + } + /// Start a list separated by `separator`. /// /// The returned type exposes identical [`.push()`][Separated::push] and @@ -527,6 +547,23 @@ where self } + /// Push the separator if applicable, then append an exact database-native named parameter + /// token and bind a value to it. + pub fn push_bind_named(&mut self, name: &'args str, value: T) -> &mut Self + where + >::Arguments: NamedArguments<'args, Database = DB>, + T: 'args + Encode<'args, DB> + Send + Type, + { + if self.push_separator { + self.query_builder.push(&self.separator); + } + + self.query_builder.push_bind_named(name, value); + self.push_separator = true; + + self + } + /// Push a bind argument placeholder (`?` or `$N` for Postgres) and bind a value to it /// without a separator. /// @@ -544,6 +581,9 @@ where mod test { use crate::postgres::Postgres; + #[cfg(feature = "sqlite")] + use crate::sqlite::Sqlite; + use super::*; #[test] @@ -588,6 +628,16 @@ mod test { ); } + #[cfg(feature = "sqlite")] + #[test] + fn test_push_bind_named() { + let mut qb: QueryBuilder<'_, Sqlite> = QueryBuilder::new("SELECT * FROM users WHERE id = "); + + qb.push_bind_named(":id", 42_i32); + + assert_eq!(qb.sql(), "SELECT * FROM users WHERE id = :id"); + } + #[test] fn test_build() { let mut qb: QueryBuilder<'_, Postgres> = QueryBuilder::new("SELECT * FROM users"); diff --git a/sqlx-core/src/query_scalar.rs b/sqlx-core/src/query_scalar.rs index b6faca064e..9d2cb957d5 100644 --- a/sqlx-core/src/query_scalar.rs +++ b/sqlx-core/src/query_scalar.rs @@ -2,7 +2,7 @@ use either::Either; use futures_core::stream::BoxStream; use futures_util::{StreamExt, TryFutureExt, TryStreamExt}; -use crate::arguments::IntoArguments; +use crate::arguments::{IntoArguments, NamedArguments}; use crate::database::{Database, HasArguments, HasStatement, HasStatementCache}; use crate::encode::Encode; use crate::error::Error; @@ -52,6 +52,16 @@ impl<'q, DB: Database, O> QueryScalar<'q, DB, O, >::Argum self.inner = self.inner.bind(value); self } + + /// Bind a value for use with a database-native named SQL parameter. + pub fn bind_named(mut self, name: &'q str, value: T) -> Self + where + >::Arguments: NamedArguments<'q, Database = DB>, + T: 'q + Send + Encode<'q, DB> + Type, + { + self.inner = self.inner.bind_named(name, value); + self + } } impl<'q, DB, O, A> QueryScalar<'q, DB, O, A> diff --git a/sqlx-core/src/sqlite/arguments.rs b/sqlx-core/src/sqlite/arguments.rs index 17b3b90f54..97e5c14233 100644 --- a/sqlx-core/src/sqlite/arguments.rs +++ b/sqlx-core/src/sqlite/arguments.rs @@ -1,4 +1,4 @@ -use crate::arguments::Arguments; +use crate::arguments::{Arguments, NamedArguments}; use crate::encode::{Encode, IsNull}; use crate::error::Error; use crate::sqlite::statement::StatementHandle; @@ -20,6 +20,8 @@ pub enum SqliteArgumentValue<'q> { #[derive(Default, Debug, Clone)] pub struct SqliteArguments<'q> { pub(crate) values: Vec>, + // (index into values, exact SQLite parameter token) + pub(crate) named: Vec<(usize, Cow<'q, str>)>, } impl<'q> SqliteArguments<'q> { @@ -39,6 +41,11 @@ impl<'q> SqliteArguments<'q> { .into_iter() .map(SqliteArgumentValue::into_static) .collect(), + named: self + .named + .into_iter() + .map(|(index, name)| (index, Cow::Owned(name.into_owned()))) + .collect(), } } } @@ -58,29 +65,63 @@ impl<'q> Arguments<'q> for SqliteArguments<'q> { } } +impl<'q> NamedArguments<'q> for SqliteArguments<'q> { + fn add_named(&mut self, name: &'q str, value: T) + where + T: 'q + Send + Encode<'q, Self::Database> + crate::types::Type, + { + self.add(value); + self.named + .push((self.values.len() - 1, Cow::Borrowed(name))); + } +} + impl SqliteArguments<'_> { pub(super) fn bind(&self, handle: &mut StatementHandle, offset: usize) -> Result { let mut arg_i = offset; - // for handle in &statement.handles { + + if !self.named.is_empty() && self.named.len() != self.values.len() { + return Err(err_protocol!( + "cannot mix named and positional SQLite parameters" + )); + } let cnt = handle.bind_parameter_count(); + // SQLite resolves exact parameter tokens natively. Keep only the name metadata here; + // no Rust hashmap or owned NUL-terminated name is needed. + for (value_i, name) in &self.named { + let param_i = handle + .bind_parameter_index(name) + .ok_or_else(|| err_protocol!("unknown SQLite parameter: {}", name))?; + + self.values[*value_i].bind(handle, param_i)?; + } + for param_i in 1..=cnt { + let parameter_name = handle.bind_parameter_name(param_i); + // figure out the index of this bind parameter into our argument tuple - let n: usize = if let Some(name) = handle.bind_parameter_name(param_i) { - if let Some(name) = name.strip_prefix('?') { + let n: usize = if let Some(name) = parameter_name { + if self + .named + .iter() + .any(|(_, bound_name)| bound_name.as_ref() == name) + { + continue; + } else if let Some(name) = name.strip_prefix('?') { // parameter should have the form ?NNN atoi(name.as_bytes()).expect("parameter of the form ?NNN") } else if let Some(name) = name.strip_prefix('$') { // parameter should have the form $NNN atoi(name.as_bytes()).ok_or_else(|| { - err_protocol!( - "parameters with non-integer names are not currently supported: {}", - name - ) + err_protocol!("named SQLite parameter was not bound: {}", name) })? } else { - return Err(err_protocol!("unsupported SQL parameter format: {}", name)); + return Err(err_protocol!( + "named SQLite parameter was not bound: {}", + name + )); } } else { arg_i += 1; diff --git a/sqlx-core/src/sqlite/statement/handle.rs b/sqlx-core/src/sqlite/statement/handle.rs index e3dd9e4787..9789441ea8 100644 --- a/sqlx-core/src/sqlite/statement/handle.rs +++ b/sqlx-core/src/sqlite/statement/handle.rs @@ -192,6 +192,12 @@ impl StatementHandle { } } + #[inline] + pub(crate) fn bind_parameter_index(&self, name: &str) -> Option { + (1..=self.bind_parameter_count()) + .find(|&index| self.bind_parameter_name(index) == Some(name)) + } + // Binding Values To Prepared Statements // https://www.sqlite.org/c3ref/bind_blob.html diff --git a/src/lib.rs b/src/lib.rs index e2a9426ab4..08190b3339 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,7 +1,7 @@ #![cfg_attr(docsrs, feature(doc_cfg))] pub use sqlx_core::acquire::Acquire; -pub use sqlx_core::arguments::{Arguments, IntoArguments}; +pub use sqlx_core::arguments::{Arguments, IntoArguments, NamedArguments}; pub use sqlx_core::column::Column; pub use sqlx_core::column::ColumnIndex; pub use sqlx_core::connection::{ConnectOptions, Connection}; @@ -151,6 +151,7 @@ pub mod prelude { pub use super::Executor; pub use super::FromRow; pub use super::IntoArguments; + pub use super::NamedArguments; pub use super::Row; pub use super::Statement; pub use super::Type; diff --git a/tests/sqlite/sqlite.rs b/tests/sqlite/sqlite.rs index 33453e176c..e5c6c3bbc7 100644 --- a/tests/sqlite/sqlite.rs +++ b/tests/sqlite/sqlite.rs @@ -94,6 +94,48 @@ async fn it_maths() -> anyhow::Result<()> { Ok(()) } +#[sqlx_macros::test] +async fn test_named_parameters() -> anyhow::Result<()> { + let mut conn = new::().await?; + + let value: i32 = sqlx_oldapi::query_scalar("select :value + :value") + .bind_named(":value", 5_i32) + .fetch_one(&mut conn) + .await?; + + assert_eq!(value, 10); + + let value: i32 = sqlx_oldapi::query_scalar("select @value + $value") + .bind_named("@value", 2_i32) + .bind_named("$value", 3_i32) + .fetch_one(&mut conn) + .await?; + + assert_eq!(value, 5); + + let error = sqlx_oldapi::query_scalar::<_, i32>("select :named + ?") + .bind_named(":named", 5_i32) + .bind(7_i32) + .fetch_one(&mut conn) + .await + .expect_err("mixed named and positional parameters should be rejected"); + assert!(error + .to_string() + .contains("cannot mix named and positional")); + + let error = sqlx_oldapi::query_scalar::<_, i32>("select $1 + :named") + .bind_named(":named", 5_i32) + .bind(7_i32) + .fetch_one(&mut conn) + .await + .expect_err("mixed named and positional parameters should be rejected"); + assert!(error + .to_string() + .contains("cannot mix named and positional")); + + Ok(()) +} + #[sqlx_macros::test] async fn test_bind_multiple_statements_multiple_values() -> anyhow::Result<()> { let mut conn = new::().await?;